use std::{error::Error, sync::Arc, time::Duration}; use async_broadcast::broadcast; use chrono::Utc; use dashmap::DashMap; use entity::{stream_session, users}; use rml_rtmp::{ handshake::{Handshake, HandshakeProcessResult, PeerType}, sessions::{ServerSession, ServerSessionConfig, ServerSessionEvent, ServerSessionResult}, }; use sea_orm::{DatabaseConnection, EntityTrait, IntoActiveModel, sqlx::types::chrono::Local}; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::{TcpListener, TcpStream}, time::timeout, }; use tracing::{debug, info, warn}; use uuid::Uuid; use crate::{ StreamCodec, StreamSession, audio::{AACParser, AudioProcesser, OpusAudioFrame}, codec::{ CodecParser, VideoFrame, av1::Av1CodecParser, h264::H264CodecParser, h265::H265CodecParser, }, }; const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(30); pub struct Rtmp { pub listener: TcpListener, pub stream_sessions: Arc>, pub db: DatabaseConnection, } async fn write_outbound(socket: &mut TcpStream, results: Vec) { for r in results { if let ServerSessionResult::OutboundResponse(p) = r && let Err(e) = socket.write_all(&p.bytes).await { warn!("RTMP write error: {e}"); return; } } } impl Rtmp { fn parse_video_codec(payload: &[u8]) -> Result> { if payload.is_empty() { return Err("empty video payload".into()); } let is_ex = payload[0] & 0x80 != 0; if is_ex { if payload.len() < 5 { return Err("enhanced RTMP payload too short".into()); } match &payload[1..5] { b"hvc1" => Ok(StreamCodec::H265), b"avc1" => Ok(StreamCodec::H264), b"av01" => Ok(StreamCodec::AV1), _ => Err("unsupported FourCC".into()), } } else { match payload[0] & 0x0F { 7 => Ok(StreamCodec::H264), _ => Err("unsupported legacy codec ID".into()), } } } async fn handshake( mut socket: TcpStream, ) -> Result<(ServerSession, TcpStream), Box> { let mut server = Handshake::new(PeerType::Server); let mut c0_c1 = [0u8; 1537]; 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) { Ok(HandshakeProcessResult::InProgress { response_bytes }) => response_bytes, x => return Err(format!("unexpected handshake state: {:?}", x).into()), }; socket.write_all(&s0_s1_s2).await?; let mut c2 = [0u8; 1536]; timeout(HANDSHAKE_TIMEOUT, socket.read_exact(&mut c2)) .await .map_err(|_| "handshake C2 timeout")??; match server.process_bytes(&c2) { Ok(HandshakeProcessResult::Completed { .. }) => {} Ok(HandshakeProcessResult::InProgress { response_bytes }) => { socket.write_all(&response_bytes).await? } x => return Err(format!("unexpected handshake state: {:?}", x).into()), } let (rtmp_session, init_bytes) = ServerSession::new(ServerSessionConfig::new()) .map_err(|e| format!("ServerSession::new failed: {e}"))?; write_outbound(&mut socket, init_bytes).await; Ok((rtmp_session, socket)) } pub async fn run(self) { let Self { listener, stream_sessions, db, } = self; loop { let connection = listener.accept().await; let (socket, peer_addr) = match connection { Ok(e) => e, Err(e) => { tracing::error!("Error on connection accept: {}", e); continue; } }; info!(%peer_addr, "RTMP connection accepted"); let (mut session, mut socket) = match Rtmp::handshake(socket).await { Ok(v) => v, Err(e) => { warn!(%peer_addr, "RTMP handshake failed: {e}"); continue; } }; info!(%peer_addr, "RTMP handshake complete"); // Ring sized by TIME, not frame count: 512 frames ≈ 5.7s @ 90fps, // 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::>(512); let (mut audio_tx, audio_rx) = broadcast::>(512); video_tx.set_overflow(true); audio_tx.set_overflow(true); let mut parser: Option> = None; let mut stream_id: Option = None; let mut aac_parser = AACParser::new(); let mut audio_proc = AudioProcesser::new(); let db = db.clone(); let stream_sessions = stream_sessions.clone(); tokio::spawn(async move { let mut current_stream_key_id: Option = None; let mut codec_stamped = false; let cleanup = |id: i32| { let stream_sessions = stream_sessions.clone(); let db = db.clone(); async move { stream_sessions.remove(&id); if let Ok(Some(s)) = stream_session::Model::get_active_by_stream_key_id(&db, id).await { s.into_active_model() .finish_stream_session(&db, Local::now().into()) .await .ok(); } } }; loop { let mut buf = [0u8; 4096]; const IDLE_TIMEOUT: Duration = Duration::from_secs(15); let n = match timeout(IDLE_TIMEOUT, socket.read(&mut buf)).await { Ok(Ok(0)) => { debug!("RTMP connection closed by peer"); if let Some(id) = current_stream_key_id { cleanup(id).await; } return; } Ok(Ok(n)) => n, Ok(Err(e)) => { warn!("RTMP read error: {e}"); if let Some(id) = current_stream_key_id { cleanup(id).await; } return; } Err(_) => { warn!("RTMP timed out (15 sec no packets) id: {:?}", stream_id); if let Some(id) = current_stream_key_id { cleanup(id).await; } return; } }; let events = match session.handle_input(&buf[..n]) { Ok(e) => e, Err(e) => { warn!("RTMP session error: {e}"); if let Some(id) = current_stream_key_id { cleanup(id).await; } return; } }; // Blankly using it, so it doesnt drop video_rx.is_closed(); audio_rx.is_closed(); for event in events { match event { ServerSessionResult::OutboundResponse(p) => { if let Err(e) = socket.write_all(&p.bytes).await { warn!("RTMP outbound write error: {e}"); return; } } ServerSessionResult::RaisedEvent(e) => match e { ServerSessionEvent::ConnectionRequested { request_id, .. } => { debug!("RTMP ConnectionRequested, accepting"); 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; } ServerSessionEvent::PublishStreamRequested { request_id, stream_key, .. } => { info!(stream_key = %stream_key, "publish stream requested"); let key = entity::stream_key::Entity::find_by_key(&db, &stream_key) .await; let key = if let Ok(Some(key)) = key { key } else { warn!(stream_key = %stream_key, "stream key not found, rejecting"); let reply = match session.reject_request( request_id, "", "Stream key invalid", ) { Ok(r) => r, Err(e) => { warn!("Failed to reject stream key request: {e}"); return; } }; write_outbound(&mut socket, reply).await; break; }; let already_live = match stream_session::Model::get_active_by_stream_key_id( &db, key.id, ) .await { 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() { warn!(stream_key_id = key.id, label = %key.label, "stream key already live, rejecting duplicate publish"); let reply = match session.reject_request( request_id, "", "You're already streaming...", ) { Ok(r) => r, Err(e) => { warn!("Failed to reject duplicate publish: {e}"); return; } }; write_outbound(&mut socket, reply).await; break; } 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); stream_sessions.insert( key.id, StreamSession { stream_key_id: key.id, stream_key_label: key.label, stream_key_user: user.username, custom_id: key.custom_id, is_unlisted: key.is_unlisted, password: key.password, frame_channel: video_tx.clone(), audio_channel: audio_tx.clone(), codec: None, started_at: Utc::now(), active_clients: 0.into(), session_id: Uuid::new_v4(), video_pt: None, video_profile_level_id: None, }, ); let reply = match session.accept_request(request_id) { Ok(r) => r, Err(e) => { warn!("Failed to accept publish request: {e}"); return; } }; if let Err(e) = stream_session::Model::create_stream_session( &db, key.id, Local::now().into(), ) .await { warn!("Failed to create stream session record: {e}"); let _ = session.reject_request( request_id, "", "Internal error", ); return; } stream_id = Some(key.id); write_outbound(&mut socket, reply).await; } ServerSessionEvent::PublishStreamFinished { stream_key, .. } => { info!(stream_key = %stream_key, "publish stream finished"); match entity::stream_key::Entity::find_by_key(&db, &stream_key) .await { Ok(Some(key)) => cleanup(key.id).await, Ok(None) => { warn!(stream_key = %stream_key, "publish finished for unknown stream key"); } Err(e) => { warn!(stream_key = %stream_key, "failed to look up stream key on finish: {e}"); } } current_stream_key_id = None; } ServerSessionEvent::VideoDataReceived { data, timestamp, .. } => match Self::parse_video_codec(&data) { Ok(codec) => { if !codec_stamped { if let Some(id) = current_stream_key_id && let Some(mut session) = stream_sessions.get_mut(&id) { session.codec = Some(codec.clone()); } codec_stamped = true; } let p = parser.get_or_insert_with(|| match codec { StreamCodec::H264 => Box::new(H264CodecParser::new()), StreamCodec::H265 => Box::new(H265CodecParser::new()), StreamCodec::AV1 => Box::new(Av1CodecParser::new()), }); if let Some(frame) = p.parse(&data, timestamp.value) { video_tx.broadcast(Arc::new(frame)).await.ok(); } } Err(_err) => { warn!(""); } }, ServerSessionEvent::AudioDataReceived { data, timestamp, .. } => { // Consume the non-Send error before any await point. let opus_frames: Vec<_> = match aac_parser.parse(&data, timestamp.value) { Err(e) => { warn!("AAC parse error: {}", e); vec![] } Ok(frame) => frame .map(|f| audio_proc.encode(f)) .unwrap_or_default(), }; for frame in opus_frames { audio_tx.broadcast(Arc::new(frame)).await.ok(); } } _ => {} }, _ => {} } } } }); } } }