Compare commits

...
3 Commits
5 changed files with 348 additions and 385 deletions
Generated
+197 -350
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -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
View File
@@ -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 -1
View File
@@ -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
View File
@@ -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;
} }