Compare commits
3
Commits
ad997d3b65
...
4bb103e9b2
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4bb103e9b2
|
||
|
|
fd1e59267d
|
||
|
|
a2c74a6e81
|
Generated
+197
-350
File diff suppressed because it is too large
Load Diff
@@ -19,7 +19,7 @@ rand = "0.10.1"
|
|||||||
rml_rtmp = "0.8.0"
|
rml_rtmp = "0.8.0"
|
||||||
str0m = "0.20.0"
|
str0m = "0.20.0"
|
||||||
tokio = { version = "1", features = ["full"] }
|
tokio = { version = "1", features = ["full"] }
|
||||||
axum = "0.8"
|
axum = { version = "0.8", features = ["macros"] }
|
||||||
serde = { version = "1.0.228", features = ["serde_derive"] }
|
serde = { version = "1.0.228", features = ["serde_derive"] }
|
||||||
serde_json = "1.0.150"
|
serde_json = "1.0.150"
|
||||||
sea-orm = { version = "1", features = [ "sqlx-sqlite", "runtime-tokio-rustls", "macros" ] }
|
sea-orm = { version = "1", features = [ "sqlx-sqlite", "runtime-tokio-rustls", "macros" ] }
|
||||||
|
|||||||
+75
-11
@@ -207,6 +207,22 @@ struct CreateStreamKeyBody {
|
|||||||
label: String,
|
label: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct SessionCookie(Option<String>);
|
||||||
|
|
||||||
|
impl<S> FromRequestParts<S> for SessionCookie
|
||||||
|
where
|
||||||
|
S: Send + Sync,
|
||||||
|
{
|
||||||
|
type Rejection = HttpError;
|
||||||
|
|
||||||
|
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
|
||||||
|
let jar = CookieJar::from_request_parts(parts, _state).await.unwrap();
|
||||||
|
Ok(SessionCookie(
|
||||||
|
jar.get("session").map(|c| c.value().to_string()),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
struct AuthUser(entity::users::Model);
|
struct AuthUser(entity::users::Model);
|
||||||
|
|
||||||
impl FromRequestParts<Arc<HttpServer>> for AuthUser {
|
impl FromRequestParts<Arc<HttpServer>> for AuthUser {
|
||||||
@@ -216,12 +232,18 @@ impl FromRequestParts<Arc<HttpServer>> for AuthUser {
|
|||||||
parts: &mut Parts,
|
parts: &mut Parts,
|
||||||
state: &Arc<HttpServer>,
|
state: &Arc<HttpServer>,
|
||||||
) -> Result<Self, HttpError> {
|
) -> Result<Self, HttpError> {
|
||||||
let jar = CookieJar::from_request_parts(parts, state).await.unwrap();
|
let jar = CookieJar::from_headers(&parts.headers);
|
||||||
let session = jar.get("session").ok_or(HttpError::Unauthorized)?;
|
let session = jar.get("session").ok_or(HttpError::Unauthorized)?;
|
||||||
|
let session_val = session.value().to_string();
|
||||||
|
tracing::debug!(?session_val, "AuthUser: extracted session cookie");
|
||||||
|
|
||||||
let user = users::Entity::find_by_auth_session(&state.db, session.to_string())
|
let user = match users::Entity::find_by_auth_session(&state.db, session_val.clone()).await {
|
||||||
.await?
|
Ok(u) => u.ok_or(HttpError::Unauthorized)?,
|
||||||
.ok_or(HttpError::Unauthorized)?;
|
Err(e) => {
|
||||||
|
tracing::error!(?session_val, error = %e, "AuthUser: db query failed");
|
||||||
|
return Err(HttpError::DbErr(e));
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
Ok(AuthUser(user))
|
Ok(AuthUser(user))
|
||||||
}
|
}
|
||||||
@@ -289,11 +311,45 @@ struct LoginResponse {
|
|||||||
session_token: String,
|
session_token: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
struct DevFlag(bool);
|
||||||
|
|
||||||
|
impl<S> FromRequestParts<S> for DevFlag
|
||||||
|
where
|
||||||
|
S: Send + Sync,
|
||||||
|
{
|
||||||
|
type Rejection = std::convert::Infallible;
|
||||||
|
|
||||||
|
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
|
||||||
|
let query = parts.uri.query().unwrap_or("");
|
||||||
|
Ok(DevFlag(query.contains("dev=1")))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn cookie_for_token(token: &str, dev: bool) -> Cookie<'static> {
|
||||||
|
let token = token.to_owned();
|
||||||
|
match dev {
|
||||||
|
true => Cookie::build(("session", token))
|
||||||
|
.path("/")
|
||||||
|
.http_only(true)
|
||||||
|
.same_site(axum_extra::extract::cookie::SameSite::None)
|
||||||
|
.build(),
|
||||||
|
false => Cookie::build(("session", token))
|
||||||
|
.path("/")
|
||||||
|
.http_only(true)
|
||||||
|
.same_site(axum_extra::extract::cookie::SameSite::Lax)
|
||||||
|
.build(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[axum::debug_handler]
|
||||||
async fn login_handler(
|
async fn login_handler(
|
||||||
|
session: SessionCookie,
|
||||||
State(state): State<Arc<HttpServer>>,
|
State(state): State<Arc<HttpServer>>,
|
||||||
|
DevFlag(dev): DevFlag,
|
||||||
Json(payload): Json<LoginForm>,
|
Json(payload): Json<LoginForm>,
|
||||||
jar: CookieJar,
|
|
||||||
) -> Result<impl IntoResponse, HttpError> {
|
) -> Result<impl IntoResponse, HttpError> {
|
||||||
|
tracing::debug!(session = ?session.0, "login: existing session cookie");
|
||||||
|
|
||||||
let user = users::Entity::find_by_username(&state.db, payload.username.clone())
|
let user = users::Entity::find_by_username(&state.db, payload.username.clone())
|
||||||
.await?
|
.await?
|
||||||
.ok_or_else(|| {
|
.ok_or_else(|| {
|
||||||
@@ -310,8 +366,11 @@ async fn login_handler(
|
|||||||
let auth = auth_session::Entity::create(&state.db, user.id).await?;
|
let auth = auth_session::Entity::create(&state.db, user.id).await?;
|
||||||
let token = auth.value;
|
let token = auth.value;
|
||||||
|
|
||||||
let jar = jar.add(Cookie::new("session", token));
|
let cookie = cookie_for_token(&token, dev);
|
||||||
Ok((StatusCode::OK, jar))
|
Ok((
|
||||||
|
StatusCode::OK,
|
||||||
|
[(axum::http::header::SET_COOKIE, cookie.to_string())],
|
||||||
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Deserialize)]
|
#[derive(Deserialize)]
|
||||||
@@ -322,10 +381,12 @@ struct CreateUserForm {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn create_user_handler(
|
async fn create_user_handler(
|
||||||
|
session: SessionCookie,
|
||||||
State(state): State<Arc<HttpServer>>,
|
State(state): State<Arc<HttpServer>>,
|
||||||
|
DevFlag(dev): DevFlag,
|
||||||
Json(payload): Json<CreateUserForm>,
|
Json(payload): Json<CreateUserForm>,
|
||||||
jar: CookieJar,
|
|
||||||
) -> Result<impl IntoResponse, HttpError> {
|
) -> Result<impl IntoResponse, HttpError> {
|
||||||
|
tracing::debug!(session = ?session.0, "create_user: existing session cookie");
|
||||||
if state.config.signup_code.is_empty() || payload.ref_token != state.config.signup_code {
|
if state.config.signup_code.is_empty() || payload.ref_token != state.config.signup_code {
|
||||||
warn!(username = %payload.username, "signup rejected: invalid signup code");
|
warn!(username = %payload.username, "signup rejected: invalid signup code");
|
||||||
return Err(HttpError::Unauthorized);
|
return Err(HttpError::Unauthorized);
|
||||||
@@ -352,8 +413,11 @@ async fn create_user_handler(
|
|||||||
let session = auth_session::Entity::create(&state.db, user.id).await?;
|
let session = auth_session::Entity::create(&state.db, user.id).await?;
|
||||||
let token = session.value;
|
let token = session.value;
|
||||||
|
|
||||||
let jar = jar.add(Cookie::new("session", token));
|
let cookie = cookie_for_token(&token, dev);
|
||||||
Ok((jar, StatusCode::OK))
|
Ok((
|
||||||
|
StatusCode::OK,
|
||||||
|
[(axum::http::header::SET_COOKIE, cookie.to_string())],
|
||||||
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn stream_handler(
|
async fn stream_handler(
|
||||||
@@ -372,7 +436,7 @@ async fn stream_handler(
|
|||||||
};
|
};
|
||||||
|
|
||||||
let stream_key_id = stream_key_id.ok_or_else(|| {
|
let stream_key_id = stream_key_id.ok_or_else(|| {
|
||||||
warn!(slug = %slug, "WHEP request for unknown or inactive stream");
|
// warn!(slug = %slug, "WHEP request for unknown or inactive stream");
|
||||||
HttpError::NotFound
|
HttpError::NotFound
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
use ::chrono::{DateTime, Utc};
|
use ::chrono::{DateTime, Utc};
|
||||||
use std::{env, error::Error, sync::Arc};
|
use std::{env, error::Error, sync::Arc};
|
||||||
use tracing::{info, level_filters::LevelFilter, warn};
|
use tracing::{info, warn};
|
||||||
|
|
||||||
use async_broadcast::broadcast;
|
use async_broadcast::broadcast;
|
||||||
use dashmap::DashMap;
|
use dashmap::DashMap;
|
||||||
|
|||||||
+74
-22
@@ -12,7 +12,7 @@ use sea_orm::{DatabaseConnection, IntoActiveModel, sqlx::types::chrono::Local};
|
|||||||
use tokio::{
|
use tokio::{
|
||||||
io::{AsyncReadExt, AsyncWriteExt},
|
io::{AsyncReadExt, AsyncWriteExt},
|
||||||
net::{TcpListener, TcpStream},
|
net::{TcpListener, TcpStream},
|
||||||
time::{Instant, timeout},
|
time::timeout,
|
||||||
};
|
};
|
||||||
use tracing::{debug, info, warn};
|
use tracing::{debug, info, warn};
|
||||||
|
|
||||||
@@ -24,6 +24,8 @@ use crate::{
|
|||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
|
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(30);
|
||||||
|
|
||||||
pub struct Rtmp {
|
pub struct Rtmp {
|
||||||
pub listener: TcpListener,
|
pub listener: TcpListener,
|
||||||
pub stream_sessions: Arc<DashMap<i32, StreamSession>>,
|
pub stream_sessions: Arc<DashMap<i32, StreamSession>>,
|
||||||
@@ -33,7 +35,10 @@ pub struct Rtmp {
|
|||||||
async fn write_outbound(socket: &mut TcpStream, results: Vec<ServerSessionResult>) {
|
async fn write_outbound(socket: &mut TcpStream, results: Vec<ServerSessionResult>) {
|
||||||
for r in results {
|
for r in results {
|
||||||
if let ServerSessionResult::OutboundResponse(p) = r {
|
if let ServerSessionResult::OutboundResponse(p) = r {
|
||||||
socket.write_all(&p.bytes).await.unwrap();
|
if let Err(e) = socket.write_all(&p.bytes).await {
|
||||||
|
warn!("RTMP write error: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -69,7 +74,10 @@ impl Rtmp {
|
|||||||
let mut server = Handshake::new(PeerType::Server);
|
let mut server = Handshake::new(PeerType::Server);
|
||||||
|
|
||||||
let mut c0_c1 = [0u8; 1537];
|
let mut c0_c1 = [0u8; 1537];
|
||||||
socket.read_exact(&mut c0_c1).await?;
|
let n = timeout(HANDSHAKE_TIMEOUT, socket.read_exact(&mut c0_c1))
|
||||||
|
.await
|
||||||
|
.map_err(|_| "handshake C0+C1 timeout")??;
|
||||||
|
debug!("handshake C0+C1 read {} bytes", n);
|
||||||
let s0_s1_s2 = match server.process_bytes(&c0_c1) {
|
let s0_s1_s2 = match server.process_bytes(&c0_c1) {
|
||||||
Ok(HandshakeProcessResult::InProgress { response_bytes }) => response_bytes,
|
Ok(HandshakeProcessResult::InProgress { response_bytes }) => response_bytes,
|
||||||
x => return Err(format!("unexpected handshake state: {:?}", x).into()),
|
x => return Err(format!("unexpected handshake state: {:?}", x).into()),
|
||||||
@@ -77,7 +85,9 @@ impl Rtmp {
|
|||||||
socket.write_all(&s0_s1_s2).await?;
|
socket.write_all(&s0_s1_s2).await?;
|
||||||
|
|
||||||
let mut c2 = [0u8; 1536];
|
let mut c2 = [0u8; 1536];
|
||||||
socket.read_exact(&mut c2).await?;
|
timeout(HANDSHAKE_TIMEOUT, socket.read_exact(&mut c2))
|
||||||
|
.await
|
||||||
|
.map_err(|_| "handshake C2 timeout")??;
|
||||||
match server.process_bytes(&c2) {
|
match server.process_bytes(&c2) {
|
||||||
Ok(HandshakeProcessResult::Completed { .. }) => {}
|
Ok(HandshakeProcessResult::Completed { .. }) => {}
|
||||||
Ok(HandshakeProcessResult::InProgress { response_bytes }) => {
|
Ok(HandshakeProcessResult::InProgress { response_bytes }) => {
|
||||||
@@ -86,7 +96,8 @@ impl Rtmp {
|
|||||||
x => return Err(format!("unexpected handshake state: {:?}", x).into()),
|
x => return Err(format!("unexpected handshake state: {:?}", x).into()),
|
||||||
}
|
}
|
||||||
|
|
||||||
let (rtmp_session, init_bytes) = ServerSession::new(ServerSessionConfig::new()).unwrap();
|
let (rtmp_session, init_bytes) =
|
||||||
|
ServerSession::new(ServerSessionConfig::new()).map_err(|e| format!("ServerSession::new failed: {e}"))?;
|
||||||
write_outbound(&mut socket, init_bytes).await;
|
write_outbound(&mut socket, init_bytes).await;
|
||||||
|
|
||||||
Ok((rtmp_session, socket))
|
Ok((rtmp_session, socket))
|
||||||
@@ -193,12 +204,21 @@ impl Rtmp {
|
|||||||
for event in events {
|
for event in events {
|
||||||
match event {
|
match event {
|
||||||
ServerSessionResult::OutboundResponse(p) => {
|
ServerSessionResult::OutboundResponse(p) => {
|
||||||
socket.write_all(&p.bytes).await.unwrap();
|
if let Err(e) = socket.write_all(&p.bytes).await {
|
||||||
|
warn!("RTMP outbound write error: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
ServerSessionResult::RaisedEvent(e) => match e {
|
ServerSessionResult::RaisedEvent(e) => match e {
|
||||||
ServerSessionEvent::ConnectionRequested { request_id, .. } => {
|
ServerSessionEvent::ConnectionRequested { request_id, .. } => {
|
||||||
debug!("RTMP ConnectionRequested, accepting");
|
debug!("RTMP ConnectionRequested, accepting");
|
||||||
let reply = session.accept_request(request_id).unwrap();
|
let reply = match session.accept_request(request_id) {
|
||||||
|
Ok(r) => r,
|
||||||
|
Err(e) => {
|
||||||
|
warn!("Failed to accept connection request {request_id}: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
write_outbound(&mut socket, reply).await;
|
write_outbound(&mut socket, reply).await;
|
||||||
}
|
}
|
||||||
ServerSessionEvent::PublishStreamRequested {
|
ServerSessionEvent::PublishStreamRequested {
|
||||||
@@ -215,28 +235,50 @@ impl Rtmp {
|
|||||||
key
|
key
|
||||||
} else {
|
} else {
|
||||||
warn!(stream_key = %stream_key, "stream key not found, rejecting");
|
warn!(stream_key = %stream_key, "stream key not found, rejecting");
|
||||||
let reply = session
|
let reply = match session.reject_request(request_id, "", "Stream key invalid") {
|
||||||
.reject_request(request_id, "", "Stream key invalid")
|
Ok(r) => r,
|
||||||
.unwrap();
|
Err(e) => {
|
||||||
|
warn!("Failed to reject stream key request: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
write_outbound(&mut socket, reply).await;
|
write_outbound(&mut socket, reply).await;
|
||||||
break;
|
break;
|
||||||
};
|
};
|
||||||
|
|
||||||
let already_live =
|
let already_live =
|
||||||
stream_session::Model::get_active_by_stream_key_id(
|
match stream_session::Model::get_active_by_stream_key_id(
|
||||||
&db, key.id,
|
&db, key.id,
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
.unwrap();
|
{
|
||||||
|
Ok(s) => s,
|
||||||
|
Err(e) => {
|
||||||
|
warn!("DB error checking active stream: {e}");
|
||||||
|
let reply = match session.reject_request(request_id, "", "Internal error") {
|
||||||
|
Ok(r) => r,
|
||||||
|
Err(e2) => {
|
||||||
|
warn!("Failed to reject request on DB error: {e2}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
write_outbound(&mut socket, reply).await;
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
if already_live.is_some() {
|
if already_live.is_some() {
|
||||||
warn!(stream_key_id = key.id, label = %key.label, "stream key already live, rejecting duplicate publish");
|
warn!(stream_key_id = key.id, label = %key.label, "stream key already live, rejecting duplicate publish");
|
||||||
let reply = session
|
let reply = match session.reject_request(
|
||||||
.reject_request(
|
request_id,
|
||||||
request_id,
|
"",
|
||||||
"",
|
"You're already streaming...",
|
||||||
"You're already streaming...",
|
) {
|
||||||
)
|
Ok(r) => r,
|
||||||
.unwrap();
|
Err(e) => {
|
||||||
|
warn!("Failed to reject duplicate publish: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
write_outbound(&mut socket, reply).await;
|
write_outbound(&mut socket, reply).await;
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
@@ -255,14 +297,24 @@ impl Rtmp {
|
|||||||
},
|
},
|
||||||
);
|
);
|
||||||
|
|
||||||
let reply = session.accept_request(request_id).unwrap();
|
let reply = match session.accept_request(request_id) {
|
||||||
stream_session::Model::create_stream_session(
|
Ok(r) => r,
|
||||||
|
Err(e) => {
|
||||||
|
warn!("Failed to accept publish request: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
if let Err(e) = stream_session::Model::create_stream_session(
|
||||||
&db,
|
&db,
|
||||||
key.id,
|
key.id,
|
||||||
Local::now().into(),
|
Local::now().into(),
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
.unwrap();
|
{
|
||||||
|
warn!("Failed to create stream session record: {e}");
|
||||||
|
let _ = session.reject_request(request_id, "", "Internal error");
|
||||||
|
return;
|
||||||
|
}
|
||||||
stream_id = Some(key.id);
|
stream_id = Some(key.id);
|
||||||
write_outbound(&mut socket, reply).await;
|
write_outbound(&mut socket, reply).await;
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user