chore: clean up and optimisation

This commit is contained in:
2026-08-07 16:06:01 +01:00
parent 5f152e9db7
commit 3287a0ed7c
6 changed files with 132 additions and 170 deletions
+3 -5
View File
@@ -1,10 +1,10 @@
use std::{error::Error, fmt::Display}; use std::{error::Error, fmt::Display};
use bytes::Bytes; use bytes::Bytes;
use rubato::{Fft, Resampler, audioadapter_buffers::direct::InterleavedSlice}; use rubato::{audioadapter_buffers::direct::InterleavedSlice, Fft, Resampler};
use symphonia::core::{ use symphonia::core::{
audio::SampleBuffer, audio::SampleBuffer,
codecs::{CODEC_TYPE_AAC, CodecParameters, Decoder, DecoderOptions}, codecs::{CodecParameters, Decoder, DecoderOptions, CODEC_TYPE_AAC},
formats::Packet, formats::Packet,
}; };
@@ -83,7 +83,6 @@ impl AudioProcesser {
pub struct AudioFrame { pub struct AudioFrame {
pub data: Bytes, // interleaved f32 PCM, little-endian pub data: Bytes, // interleaved f32 PCM, little-endian
pub timestamp_ms: u32,
pub sample_rate: u32, pub sample_rate: u32,
} }
@@ -123,7 +122,7 @@ impl AACParser {
pub fn parse( pub fn parse(
&mut self, &mut self,
bytes: &[u8], bytes: &[u8],
timestamp_ms: u32, _timestamp_ms: u32,
) -> Result<Option<AudioFrame>, Box<dyn Error>> { ) -> Result<Option<AudioFrame>, Box<dyn Error>> {
if bytes.len() < 2 { if bytes.len() < 2 {
return Ok(None); return Ok(None);
@@ -169,7 +168,6 @@ impl AACParser {
Ok(Some(AudioFrame { Ok(Some(AudioFrame {
data: Bytes::from(pcm_bytes), data: Bytes::from(pcm_bytes),
timestamp_ms,
sample_rate: spec.rate, sample_rate: spec.rate,
})) }))
} }
+54 -65
View File
@@ -4,6 +4,7 @@ use std::{
Arc, LazyLock, Arc, LazyLock,
atomic::{AtomicI32, Ordering}, atomic::{AtomicI32, Ordering},
}, },
time::Duration,
}; };
use axum_extra::extract::{CookieJar, cookie::Cookie}; use axum_extra::extract::{CookieJar, cookie::Cookie};
@@ -30,7 +31,7 @@ use axum::{
}; };
use chrono::{DateTime, Utc}; use chrono::{DateTime, Utc};
use entity::{auth_session, stream_key, stream_session, users}; use entity::{auth_session, stream_key, stream_session, users};
use sea_orm::{DatabaseConnection, EntityTrait, IntoActiveModel, QueryFilter}; use sea_orm::{DatabaseConnection, EntityTrait, IntoActiveModel};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use sysinfo::System; use sysinfo::System;
use tokio::{ use tokio::{
@@ -80,7 +81,13 @@ impl HttpServer {
]; ];
let cors = CorsLayer::new() let cors = CorsLayer::new()
.allow_methods([Method::GET, Method::POST, Method::PATCH, Method::DELETE, Method::OPTIONS]) .allow_methods([
Method::GET,
Method::POST,
Method::PATCH,
Method::DELETE,
Method::OPTIONS,
])
.allow_headers([ .allow_headers([
AUTHORIZATION, AUTHORIZATION,
ACCEPT, ACCEPT,
@@ -103,9 +110,15 @@ impl HttpServer {
// .route("/api/whip", post(handle_whip_injest)) // .route("/api/whip", post(handle_whip_injest))
.route("/api/login", post(login_handler)) .route("/api/login", post(login_handler))
.route("/api/stream/{slug}", post(stream_handler)) .route("/api/stream/{slug}", post(stream_handler))
.route("/stream", post(handle_whip_injest)) .route("/api/whip", post(handle_whip_injest))
.route("/stream/{slug}", delete(handle_whip_injest_delete).patch(handle_whip_injest_patch)) .route(
.route("/api/whip/{slug}", delete(handle_whip_injest_delete).patch(handle_whip_injest_patch)) "/api/whip/{slug}",
delete(handle_whip_injest_delete).patch(handle_whip_injest_patch),
)
// .route(
// "/api/whip/{slug}",
// delete(handle_whip_injest_delete).patch(handle_whip_injest_patch),
// )
.route("/api/meow", get(meow_handler)) .route("/api/meow", get(meow_handler))
.route("/api/health", get(health_handler)) .route("/api/health", get(health_handler))
.route("/api/uptime", get(uptime_handler)) .route("/api/uptime", get(uptime_handler))
@@ -125,11 +138,6 @@ impl HttpServer {
} }
} }
#[derive(Serialize)]
struct StreamCatalog {
active_streams: Vec<StreamListing>,
}
#[derive(Serialize)] #[derive(Serialize)]
struct StreamListing { struct StreamListing {
label: String, label: String,
@@ -181,39 +189,20 @@ async fn catalog_handler(
) -> Result<impl IntoResponse, HttpError> { ) -> Result<impl IntoResponse, HttpError> {
let streams = stream_session::Model::get_all_active_sessions(&state.db).await?; let streams = stream_session::Model::get_all_active_sessions(&state.db).await?;
debug!("{:#?}", streams); debug!("{:#?}", streams);
let catalog: Vec<StreamListing> = futures::future::join_all(streams.iter().map(|listing| { let catalog: Vec<StreamListing> = state
let db = state.db.clone();
let stream_key_id = listing.stream_key_id;
let state2 = state.clone();
async move {
let key_info = stream_key::Entity::find_by_id(stream_key_id)
.one(&db)
.await?
.ok_or(HttpError::NotFound)?;
let user = users::Entity::find_by_id(key_info.user_id)
.one(&db)
.await?
.ok_or(HttpError::NotFound)?;
let meow = state2
.appstate .appstate
.clone()
.lock() .lock()
.await .await
.stream_sessions .stream_sessions
.get(&key_info.id) .iter()
.ok_or(HttpError::NotFound)? .map(|x| StreamListing {
.started_at; id: x.stream_key_id,
Ok::<_, HttpError>(StreamListing { label: x.stream_key_label.clone(),
id: stream_key_id, user: x.stream_key_user.clone(),
label: key_info.label, started_at: x.started_at,
user: user.username,
started_at: meow,
}) })
}
}))
.await
.into_iter() .into_iter()
.collect::<Result<Vec<_>, HttpError>>()?; .collect::<Vec<_>>();
Ok(Json(catalog)) Ok(Json(catalog))
} }
@@ -297,16 +286,6 @@ async fn create_stream_key_handler(
Ok(StatusCode::CREATED) Ok(StatusCode::CREATED)
} }
struct StreamKeys {
keys: Vec<StreamKey>,
}
struct StreamKey {
id: i32,
label: String,
value: String,
}
async fn get_all_stream_keys( async fn get_all_stream_keys(
State(state): State<Arc<HttpServer>>, State(state): State<Arc<HttpServer>>,
auth: AuthUser, auth: AuthUser,
@@ -321,11 +300,6 @@ struct LoginForm {
password: String, password: String,
} }
#[derive(Serialize)]
struct LoginResponse {
session_token: String,
}
struct DevFlag(bool); struct DevFlag(bool);
impl<S> FromRequestParts<S> for DevFlag impl<S> FromRequestParts<S> for DevFlag
@@ -457,7 +431,7 @@ async fn stream_handler(
})?; })?;
info!(request_id = request_id_clone, slug = %slug, stream_key_id, "WHEP offer received"); info!(request_id = request_id_clone, slug = %slug, stream_key_id, "WHEP offer received");
let mut accept_rx = state.accept_rx.activate_cloned(); let accept_rx = state.accept_rx.activate_cloned();
// The webrtc worker owning the receiver died if this fails. // The webrtc worker owning the receiver died if this fails.
state state
.offer_tx .offer_tx
@@ -469,35 +443,50 @@ async fn stream_handler(
"offer sent, waiting for answer" "offer sent, waiting for answer"
); );
let reply_body = String::new(); // Bound the wait: a WebRTC setup failure (malformed SDP, codec mismatch)
while let Ok(answer) = accept_rx.recv().await { // would otherwise leave this HTTP request hanging forever.
debug!(request_id = request_id_clone, "received answer candidate"); match tokio::time::timeout(Duration::from_secs(10), async {
if let Some(reply) = answer.1 { let mut accept_rx = accept_rx;
if answer.0 == request_id_clone { loop {
match accept_rx.recv().await {
Ok(answer) if answer.0 == request_id_clone => return Some(answer.1),
Ok(_) => continue,
Err(_) => return None,
}
}
})
.await
{
Ok(Some(Some(reply))) => {
return Ok(Response::builder() return Ok(Response::builder()
.status(StatusCode::CREATED) .status(StatusCode::CREATED)
.header("content-type", "application/sdp") .header("content-type", "application/sdp")
.body(reply) .body(reply)
.unwrap()); .unwrap());
} else {
continue;
} }
} else if answer.0 == request_id_clone { Ok(Some(None)) => {
return Ok(Response::builder() return Ok(Response::builder()
.status(StatusCode::UNSUPPORTED_MEDIA_TYPE) .status(StatusCode::UNSUPPORTED_MEDIA_TYPE)
.body(String::new()) .body(String::new())
.unwrap()); .unwrap());
} }
} Ok(None) => {
info!( info!(
request_id = request_id_clone, request_id = request_id_clone,
"answer channel closed without a match" "answer channel closed without a match"
); );
}
Err(_) => {
warn!(
request_id = request_id_clone,
"timed out waiting for WHEP answer"
);
}
};
Ok(Response::builder() Ok(Response::builder()
.status(StatusCode::CREATED) .status(StatusCode::GATEWAY_TIMEOUT)
.header("content-type", "application/sdp") .body(String::new())
.body(reply_body)
.unwrap()) .unwrap())
} }
+2 -7
View File
@@ -11,11 +11,7 @@ use dashmap::DashMap;
use entity::stream_session; use entity::stream_session;
use migration::{Migrator, MigratorTrait}; use migration::{Migrator, MigratorTrait};
use sea_orm::Database; use sea_orm::Database;
use tokio::{ use tokio::{net::TcpListener, sync::Mutex, task::JoinSet};
net::TcpListener,
sync::Mutex,
task::JoinSet,
};
use tracing_subscriber::EnvFilter; use tracing_subscriber::EnvFilter;
use crate::{ use crate::{
@@ -50,11 +46,10 @@ pub enum StreamCodec {
AV1, AV1,
} }
pub struct StreamSession { pub struct StreamSession {
pub stream_key_id: i32, pub stream_key_id: i32,
pub stream_key_label: String, pub stream_key_label: String,
pub stream_key_user: String,
pub frame_channel: async_broadcast::Sender<Arc<VideoFrame>>, pub frame_channel: async_broadcast::Sender<Arc<VideoFrame>>,
pub audio_channel: async_broadcast::Sender<Arc<OpusAudioFrame>>, pub audio_channel: async_broadcast::Sender<Arc<OpusAudioFrame>>,
pub codec: Option<StreamCodec>, pub codec: Option<StreamCodec>,
+16 -5
View File
@@ -3,12 +3,12 @@ use std::{error::Error, sync::Arc, time::Duration};
use async_broadcast::broadcast; use async_broadcast::broadcast;
use chrono::Utc; use chrono::Utc;
use dashmap::DashMap; use dashmap::DashMap;
use entity::stream_session; use entity::{stream_session, users};
use rml_rtmp::{ use rml_rtmp::{
handshake::{Handshake, HandshakeProcessResult, PeerType}, handshake::{Handshake, HandshakeProcessResult, PeerType},
sessions::{ServerSession, ServerSessionConfig, ServerSessionEvent, ServerSessionResult}, sessions::{ServerSession, ServerSessionConfig, ServerSessionEvent, ServerSessionResult},
}; };
use sea_orm::{DatabaseConnection, IntoActiveModel, sqlx::types::chrono::Local}; use sea_orm::{DatabaseConnection, EntityTrait, IntoActiveModel, sqlx::types::chrono::Local};
use tokio::{ use tokio::{
io::{AsyncReadExt, AsyncWriteExt}, io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream}, net::{TcpListener, TcpStream},
@@ -35,7 +35,8 @@ 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
&& let Err(e) = socket.write_all(&p.bytes).await { && let Err(e) = socket.write_all(&p.bytes).await
{
warn!("RTMP write error: {e}"); warn!("RTMP write error: {e}");
return; return;
} }
@@ -126,8 +127,12 @@ impl Rtmp {
} }
}; };
info!(%peer_addr, "RTMP handshake complete"); info!(%peer_addr, "RTMP handshake complete");
let (mut video_tx, video_rx) = broadcast::<Arc<VideoFrame>>(32); // Ring sized by TIME, not frame count: 512 frames ≈ 5.7s @ 90fps,
let (mut audio_tx, audio_rx) = broadcast::<Arc<OpusAudioFrame>>(32); // so transient ingest/send stalls can't overflow it (overflow drops
// the oldest frames and forces a freeze until the next keyframe).
// ponytail: fixed 512; per-stream dynamic sizing if memory ever matters.
let (mut video_tx, video_rx) = broadcast::<Arc<VideoFrame>>(512);
let (mut audio_tx, audio_rx) = broadcast::<Arc<OpusAudioFrame>>(512);
// video_rx.cycle // video_rx.cycle
video_tx.set_overflow(true); video_tx.set_overflow(true);
@@ -295,12 +300,18 @@ impl Rtmp {
} }
info!(stream_key_id = key.id, label = %key.label, "stream started"); info!(stream_key_id = key.id, label = %key.label, "stream started");
let user = users::Entity::find_by_id(key.user_id)
.one(&db)
.await
.unwrap()
.unwrap();
current_stream_key_id = Some(key.id); current_stream_key_id = Some(key.id);
stream_sessions.insert( stream_sessions.insert(
key.id, key.id,
StreamSession { StreamSession {
stream_key_id: key.id, stream_key_id: key.id,
stream_key_label: key.label, stream_key_label: key.label,
stream_key_user: user.username,
frame_channel: video_tx.clone(), frame_channel: video_tx.clone(),
audio_channel: audio_tx.clone(), audio_channel: audio_tx.clone(),
codec: None, codec: None,
+28 -66
View File
@@ -18,15 +18,13 @@ use crate::{
}; };
pub struct Webrtc { pub struct Webrtc {
pub offer_rx: Receiver<(i32, i32, String)>, pub offer_rx: Receiver<(i32, i32, String)>,
pub accept_tx: async_broadcast::Sender<(i32, Option<String>)>, pub accept_tx: async_broadcast::Sender<(i32, Option<String>)>,
pub sessions_ref: Arc<DashMap<i32, StreamSession>>, pub sessions_ref: Arc<DashMap<i32, StreamSession>>,
pub proxy: Arc<WebrtcProxy>, pub proxy: Arc<WebrtcProxy>,
pub db: DatabaseConnection, pub db: DatabaseConnection,
} }
const PER_CLIENT_CONNECTION_BUF: usize = 65535;
impl Webrtc { impl Webrtc {
pub async fn run(mut self) { pub async fn run(mut self) {
while let Some(offer) = self.offer_rx.recv().await { while let Some(offer) = self.offer_rx.recv().await {
@@ -129,7 +127,17 @@ impl Webrtc {
let candidate = Candidate::host(local_addr, Protocol::Udp).unwrap(); let candidate = Candidate::host(local_addr, Protocol::Udp).unwrap();
rtc.add_local_candidate(candidate); rtc.add_local_candidate(candidate);
let offer_sdp = SdpOffer::from_sdp_string(&sdp_body).unwrap(); // Malformed SDP from an unauthenticated client must not panic here:
// a panic kills the webrtc worker and main() aborts every other
// worker, taking the whole server down.
let offer_sdp = match SdpOffer::from_sdp_string(&sdp_body) {
Ok(sdp) => sdp,
Err(e) => {
warn!(request_id, stream_id, "malformed SDP offer: {:?}", e);
self.accept_tx.broadcast((request_id, None)).await.unwrap();
continue;
}
};
let mut changes = rtc.sdp_api(); let mut changes = rtc.sdp_api();
let mid = changes.add_media( let mid = changes.add_media(
MediaKind::Video, MediaKind::Video,
@@ -378,15 +386,21 @@ impl Webrtc {
res = Webrtc::recv_video(&mut video_stream), if video_stream.is_some() => { res = Webrtc::recv_video(&mut video_stream), if video_stream.is_some() => {
match res { match res {
Ok(frame) => { Ok(frame) => {
// One frame per wake: the select! re-arms immediately while
// more frames are pending, so backlog catch-up is paced
// frame-by-frame instead of an 8-frame burst (which overflows
// the browser jitter buffer and self-reinforces the backlog).
Webrtc::write_video_frame( Webrtc::write_video_frame(
frame, &mut saw_keyframe, video_pt, video_mid, &mut rtc, stream_id, frame, &mut saw_keyframe, video_pt, video_mid, &mut rtc, stream_id,
); );
Webrtc::drain_video(
&mut video_stream, &mut saw_keyframe, video_pt, video_mid, &mut rtc, stream_id,
);
} }
Err(async_broadcast::RecvError::Closed) => { Err(async_broadcast::RecvError::Closed) => {
debug!(stream_id, "video channel closed, stream ended"); debug!(stream_id, "video channel closed, stream ended");
// Drop the dead receiver so the select! branch disarms
// instead of spinning on Closed every iteration. The
// subscribe block above re-arms it if the same stream
// key is re-published.
video_stream = None;
} }
Err(async_broadcast::RecvError::Overflowed(_)) => { Err(async_broadcast::RecvError::Overflowed(_)) => {
// Frames were dropped from the ring buffer; resuming mid-GOP // Frames were dropped from the ring buffer; resuming mid-GOP
@@ -400,10 +414,10 @@ impl Webrtc {
match res { match res {
Ok(frame) => { Ok(frame) => {
Webrtc::write_audio_frame(frame, audio_pt, audio_mid, &mut rtc); Webrtc::write_audio_frame(frame, audio_pt, audio_mid, &mut rtc);
Webrtc::drain_audio(&mut audio_stream, audio_pt, audio_mid, &mut rtc);
} }
Err(async_broadcast::RecvError::Closed) => { Err(async_broadcast::RecvError::Closed) => {
debug!("audio channel closed, stream ended"); debug!("audio channel closed, stream ended");
audio_stream = None;
} }
Err(async_broadcast::RecvError::Overflowed(_)) => {} Err(async_broadcast::RecvError::Overflowed(_)) => {}
} }
@@ -460,38 +474,6 @@ impl Webrtc {
} }
} }
fn drain_video(
stream: &mut Option<async_broadcast::Receiver<Arc<VideoFrame>>>,
saw_keyframe: &mut bool,
video_pt: Option<Pt>,
video_mid: Option<Mid>,
rtc: &mut Rtc,
stream_id: i32,
) {
let Some(s) = stream.as_mut() else { return };
for _ in 0..7 {
match s.try_recv() {
Ok(frame) => Webrtc::write_video_frame(
frame,
saw_keyframe,
video_pt,
video_mid,
rtc,
stream_id,
),
Err(async_broadcast::TryRecvError::Empty) => break,
Err(async_broadcast::TryRecvError::Closed) => {
warn!(stream_id, "video channel closed, stream ended");
*stream = None;
break;
}
Err(async_broadcast::TryRecvError::Overflowed(_)) => {
*saw_keyframe = false;
}
}
}
}
fn write_audio_frame( fn write_audio_frame(
frame: Arc<OpusAudioFrame>, frame: Arc<OpusAudioFrame>,
audio_pt: Option<Pt>, audio_pt: Option<Pt>,
@@ -505,26 +487,6 @@ impl Webrtc {
{ {
warn!("RTP write error: {:?}", e); warn!("RTP write error: {:?}", e);
} }
}
fn drain_audio(
stream: &mut Option<async_broadcast::Receiver<Arc<OpusAudioFrame>>>,
audio_pt: Option<Pt>,
audio_mid: Option<Mid>,
rtc: &mut Rtc,
) {
let Some(s) = stream.as_mut() else { return };
for _ in 0..7 {
match s.try_recv() {
Ok(frame) => Webrtc::write_audio_frame(frame, audio_pt, audio_mid, rtc),
Err(async_broadcast::TryRecvError::Empty) => break,
Err(async_broadcast::TryRecvError::Closed) => {
info!("audio channel closed, stream ended");
*stream = None;
break;
}
Err(async_broadcast::TryRecvError::Overflowed(_)) => {}
}
}
}
} }
}
+8 -1
View File
@@ -236,7 +236,8 @@ pub async fn handle_whip_injest(
// Insert the candidate line after the video m= line. // Insert the candidate line after the video m= line.
let cand_replacement = format!( let cand_replacement = format!(
"m=video 9 UDP/TLS/RTP/SAVPF 96\r\nc=IN IP4 {}\r\n{}\r\n", "m=video 9 UDP/TLS/RTP/SAVPF 96\r\nc=IN IP4 {}\r\n{}\r\n",
public_addr.ip(), cand_line public_addr.ip(),
cand_line
); );
answer_sdp.replace( answer_sdp.replace(
"m=video 9 UDP/TLS/RTP/SAVPF 96\r\nc=IN IP4 0.0.0.0\r\n", "m=video 9 UDP/TLS/RTP/SAVPF 96\r\nc=IN IP4 0.0.0.0\r\n",
@@ -277,11 +278,17 @@ pub async fn handle_whip_injest(
video_tx.set_overflow(true); video_tx.set_overflow(true);
audio_tx.set_overflow(true); audio_tx.set_overflow(true);
let user = users::Entity::find_by_id(key.user_id)
.one(&db)
.await
.unwrap()
.unwrap();
stream_sessions.insert( stream_sessions.insert(
key.id, key.id,
StreamSession { StreamSession {
stream_key_id: key.id, stream_key_id: key.id,
stream_key_label: key.label, stream_key_label: key.label,
stream_key_user: user.username,
frame_channel: video_tx, frame_channel: video_tx,
audio_channel: audio_tx, audio_channel: audio_tx,
codec: negotiated_codec, codec: negotiated_codec,