This commit is contained in:
2026-07-01 01:49:51 +01:00
parent 808d6ac5db
commit c1da3abe7d
12 changed files with 1234 additions and 266 deletions
+27
View File
@@ -1,6 +1,7 @@
use sea_orm::{
ActiveValue::{NotSet, Set},
entity::prelude::*,
sea_query::Expr,
sqlx::types::chrono::{self, Utc},
};
use serde::{Deserialize, Serialize};
@@ -43,6 +44,7 @@ impl Model {
id: NotSet,
stream_key_id: Set(stream_key_id),
started_at: Set(started_at),
ended_at: NotSet,
..Default::default()
}
.insert(db)
@@ -59,6 +61,31 @@ impl Model {
"stream_session {stream_session_id}"
)))
}
pub async fn get_all_active_sessions(db: &DatabaseConnection) -> Result<Vec<Model>, DbErr> {
Ok(Entity::find()
.filter(Column::EndedAt.is_null())
.all(db)
.await?)
}
pub async fn get_active_by_stream_key_id(
db: &DatabaseConnection,
stream_key_id: i32,
) -> Result<Option<Model>, DbErr> {
Ok(Entity::find()
.filter(Column::StreamKeyId.eq(stream_key_id))
.filter(Column::EndedAt.is_null())
.one(db)
.await?)
}
pub async fn clean_unended_streams(db: &DatabaseConnection) -> Result<u64, DbErr> {
let result = Entity::update_many()
.filter(Column::EndedAt.is_null())
.col_expr(Column::EndedAt, Expr::col(Column::StartedAt).into())
.exec(db)
.await?;
Ok(result.rows_affected)
}
}
impl ActiveModel {
+10
View File
@@ -61,6 +61,16 @@ impl Entity {
return Ok(None);
}
}
pub async fn find_by_username(
db: &DatabaseConnection,
username: String,
) -> Result<Option<Model>, DbErr> {
let user = Entity::find()
.filter(Column::Username.eq(username))
.one(db)
.await;
user
}
}
impl ActiveModel {
+10
View File
@@ -2,6 +2,9 @@
name = "server"
version = "0.1.0"
edition = "2024"
[target.x86_64-unknown-linux-gnu]
linker = "clang"
rustflags = ["-Clink-arg=-fuse-ld=/usr/local/bin/mold", "-Clink-arg=-Wl,--no-rosegment"]
[[bin]]
name = "rtmp-to-whip"
@@ -23,3 +26,10 @@ entity = {path = "../entity"}
migration = {path = "../migration"}
argon2 = "0.5.3"
uuid = { version = "1.23.3", features = ["v4"] }
tower-http = { version = "0.6", features = ["cors"] }
futures = "0.3.32"
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
symphonia = { version = "0.5", features = ["aac"] }
opus = "0.3.1"
rubato = "3.0.0"
+176
View File
@@ -0,0 +1,176 @@
use std::{error::Error, fmt::Display};
use bytes::Bytes;
use rubato::{Fft, Resampler, audioadapter_buffers::direct::InterleavedSlice};
use symphonia::core::{
audio::SampleBuffer,
codecs::{CODEC_TYPE_AAC, CodecParameters, Decoder, DecoderOptions},
formats::Packet,
};
// Opus at 48kHz: valid frame sizes are 120/240/480/960/1920/2880 samples per channel.
// We use 960 (20ms), the standard VoIP size.
const OPUS_FRAME_SAMPLES: usize = 960;
const OPUS_FRAME_INTERLEAVED: usize = OPUS_FRAME_SAMPLES * 2; // stereo
// Assumes 44100 Hz stereo input from AAC-LC.
pub struct AudioProcesser {
resampler: Fft<f32>,
resample_out: Vec<f32>,
// Accumulates resampled stereo f32 PCM until we have a full Opus frame.
pcm_buf: Vec<f32>,
encoder: opus::Encoder,
// Monotonic 48kHz sample counter; drives RTP timestamps independent of RTMP timestamps.
samples_emitted: u64,
}
impl AudioProcesser {
pub fn new() -> AudioProcesser {
let resampler =
Fft::<f32>::new(44100, 48000, 1024, 2, 2, rubato::FixedSync::Input).unwrap();
let resample_out = vec![0.0f32; resampler.output_frames_max() * 2];
AudioProcesser {
resampler,
resample_out,
pcm_buf: Vec::with_capacity(OPUS_FRAME_INTERLEAVED * 2),
encoder: opus::Encoder::new(48000, opus::Channels::Stereo, opus::Application::LowDelay)
.unwrap(),
samples_emitted: 0,
}
}
// Takes a decoded PCM AudioFrame, resamples it, buffers it, and returns however
// many complete 20ms Opus frames could be produced.
pub fn encode(&mut self, frame: AudioFrame) -> Vec<OpusAudioFrame> {
let mut samples: Vec<f32> = frame
.data
.chunks_exact(4)
.map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.collect();
if frame.sample_rate == 48000 {
self.pcm_buf.extend_from_slice(&samples);
} else {
// Resample from 44100 Hz to 48000 Hz.
let input_frames = samples.len() / 2;
let input_buf = InterleavedSlice::new_mut(&mut samples, 2, input_frames).unwrap();
let max_out = self.resample_out.len() / 2;
let mut output_buf =
InterleavedSlice::new_mut(&mut self.resample_out, 2, max_out).unwrap();
let (_, output_frames) = self
.resampler
.process_into_buffer(&input_buf, &mut output_buf, None)
.unwrap();
self.pcm_buf
.extend_from_slice(&self.resample_out[..output_frames * 2]);
}
let mut result = Vec::new();
while self.pcm_buf.len() >= OPUS_FRAME_INTERLEAVED {
let chunk: Vec<f32> = self.pcm_buf.drain(..OPUS_FRAME_INTERLEAVED).collect();
let encoded = self.encoder.encode_vec_float(&chunk, 4096).unwrap();
result.push(OpusAudioFrame {
data: Bytes::copy_from_slice(&encoded),
// timestamp_ms here represents the 48kHz sample clock ÷ 48,
// so webrtc.rs can multiply by 48 to get the RTP clock tick.
timestamp_ms: (self.samples_emitted / 48) as u32,
});
self.samples_emitted += OPUS_FRAME_SAMPLES as u64;
}
result
}
}
pub struct AudioFrame {
pub data: Bytes, // interleaved f32 PCM, little-endian
pub timestamp_ms: u32,
pub sample_rate: u32,
}
pub struct OpusAudioFrame {
pub data: Bytes,
pub timestamp_ms: u32,
}
pub struct AACParser {
decoder: Option<Box<dyn Decoder>>,
}
#[derive(Debug)]
enum AudioParseError {
InvalidCodec,
NoConfigPacket,
}
impl Display for AudioParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::InvalidCodec => write!(f, "Not the right codec provided"),
Self::NoConfigPacket => write!(f, "No config packet cached"),
}
}
}
impl Error for AudioParseError {}
// Byte 0: upper nibble = sound format (10 = AAC)
// Byte 1: 0 = AudioSpecificConfig, 1 = raw AAC frame
impl AACParser {
pub fn new() -> Self {
Self { decoder: None }
}
pub fn parse(
&mut self,
bytes: &[u8],
timestamp_ms: u32,
) -> Result<Option<AudioFrame>, Box<dyn Error>> {
if bytes.len() < 2 {
return Ok(None);
}
if (bytes[0] >> 4) != 10 {
return Err(Box::new(AudioParseError::InvalidCodec));
}
if bytes[1] == 0 {
let mut params = CodecParameters::new();
params
.for_codec(CODEC_TYPE_AAC)
.with_extra_data(bytes[2..].to_vec().into_boxed_slice());
self.decoder =
Some(symphonia::default::get_codecs().make(&params, &DecoderOptions::default())?);
return Ok(None);
}
let decoder = self
.decoder
.as_mut()
.ok_or(AudioParseError::NoConfigPacket)?;
let packet = Packet::new_from_boxed_slice(
0,
0,
1024, // AAC-LC frame is always 1024 samples
bytes[2..].to_vec().into_boxed_slice(),
);
let audio_buf = decoder.decode(&packet)?;
let spec = *audio_buf.spec();
let mut sample_buf = SampleBuffer::<f32>::new(audio_buf.capacity() as u64, spec);
sample_buf.copy_interleaved_ref(audio_buf);
let pcm_bytes: Vec<u8> = sample_buf
.samples()
.iter()
.flat_map(|s| s.to_le_bytes())
.collect();
Ok(Some(AudioFrame {
data: Bytes::from(pcm_bytes),
timestamp_ms,
sample_rate: spec.rate,
}))
}
}
+150 -48
View File
@@ -4,24 +4,35 @@ use axum::{
Json, Router,
body::{Body, Bytes},
extract::{Form, FromRequestParts, Path, Query, State},
http::{HeaderMap, StatusCode, header::SET_COOKIE, request::Parts},
http::{
HeaderMap, HeaderName, Method, StatusCode,
header::{ACCEPT, AUTHORIZATION, CONTENT_TYPE, SET_COOKIE},
request::Parts,
},
response::{IntoResponse, Response},
routing::{get, post},
};
use entity::{auth_session, stream_key, users};
use sea_orm::DatabaseConnection;
use entity::{auth_session, stream_key, stream_session, users};
use sea_orm::{DatabaseConnection, EntityTrait, QueryFilter};
use serde::{Deserialize, Serialize};
use tokio::{
net::TcpListener,
sync::{Mutex, mpsc::Sender},
};
use tower_http::cors::CorsLayer;
use uuid::Uuid;
use crate::{AppState, hash::hash_password};
use tracing::{debug, info};
use crate::{
AppState,
hash::{hash_password, verify_password},
webrtc_ingest::handle_whip_injest,
};
pub struct HttpServer {
pub offer_tx: Sender<(i32, i32, String)>,
pub accept_rx: async_broadcast::Receiver<(i32, Option<String>)>,
pub accept_rx: async_broadcast::InactiveReceiver<(i32, Option<String>)>,
pub appstate: Arc<Mutex<AppState>>,
pub request_count: Mutex<i32>,
pub db: DatabaseConnection,
@@ -31,6 +42,22 @@ impl HttpServer {
pub fn start(self) -> Result<(), Box<dyn std::error::Error>> {
let state = Arc::new(self);
let origins = [
"http://localhost:5173".parse().unwrap(),
"https://stream.h.doloro.co.uk".parse().unwrap(),
];
let cors = CorsLayer::new()
.allow_methods([Method::GET, Method::POST, Method::OPTIONS])
.allow_headers([
AUTHORIZATION,
ACCEPT,
CONTENT_TYPE,
HeaderName::from_static("session"),
])
.allow_credentials(true)
.allow_origin(origins);
tokio::spawn(async move {
let app = Router::new()
.route("/api/catalog", get(catalog_handler))
@@ -39,9 +66,11 @@ impl HttpServer {
"/api/stream-key",
post(create_stream_key_handler).get(get_all_stream_keys),
)
.route("/api/whip", post(handle_whip_injest))
.route("/api/login", post(login_handler))
.route("/api/stream/{slug}", post(stream_handler))
.route("/api/meow", get(meow_handler))
.layer(cors)
.with_state(state);
let addr = SocketAddr::from(([127, 0, 0, 1], 3000));
@@ -52,8 +81,45 @@ impl HttpServer {
}
}
#[derive(Serialize)]
struct StreamCatalog {
active_streams: Vec<StreamListing>,
}
#[derive(Serialize)]
struct StreamListing {
label: String,
id: i32,
user: String,
}
async fn catalog_handler(State(state): State<Arc<HttpServer>>) -> impl IntoResponse {
let catalog = catalog_from_state(&state.appstate).await;
let streams = stream_session::Model::get_all_active_sessions(&state.db)
.await
.unwrap();
debug!("{:#?}", streams);
let catalog: Vec<StreamListing> = futures::future::join_all(streams.iter().map(|listing| {
let db = state.db.clone();
let stream_key_id = listing.stream_key_id;
async move {
let key_info = stream_key::Entity::find_by_id(stream_key_id)
.one(&db)
.await
.unwrap()
.unwrap();
let user = users::Entity::find_by_id(key_info.user_id)
.one(&db)
.await
.unwrap()
.unwrap();
StreamListing {
id: stream_key_id,
label: key_info.label,
user: user.username,
}
}
}))
.await;
Json(catalog)
}
@@ -62,6 +128,22 @@ struct CreateStreamKeyBody {
label: String,
}
fn extract_session_token(headers: &HeaderMap) -> Option<String> {
// Check bare `session` header first (curl / API clients)
if let Some(v) = headers.get("session") {
return Some(v.to_str().ok()?.to_string());
}
// Fall back to Cookie header (browsers)
let cookie_header = headers.get("cookie")?.to_str().ok()?;
cookie_header
.split(';')
.find_map(|pair| {
let pair = pair.trim();
pair.strip_prefix("session=")
})
.map(|v| v.to_string())
}
struct AuthUser(entity::users::Model);
impl FromRequestParts<Arc<HttpServer>> for AuthUser {
@@ -71,17 +153,12 @@ impl FromRequestParts<Arc<HttpServer>> for AuthUser {
parts: &mut Parts,
state: &Arc<HttpServer>,
) -> Result<Self, StatusCode> {
let token = parts
.headers
.get("session")
// parse token...
.ok_or(StatusCode::UNAUTHORIZED)?;
let token = extract_session_token(&parts.headers).ok_or(StatusCode::UNAUTHORIZED)?;
let user =
users::Entity::find_by_auth_session(&state.db, token.to_str().unwrap().to_string())
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
.ok_or(StatusCode::UNAUTHORIZED)?;
let user = users::Entity::find_by_auth_session(&state.db, token)
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?
.ok_or(StatusCode::UNAUTHORIZED)?;
Ok(AuthUser(user))
}
@@ -94,6 +171,13 @@ async fn create_stream_key_handler(
) -> impl IntoResponse {
let uuid = uuid::Uuid::new_v4();
let value = format!("stream-key-{uuid}");
let key_amount = stream_key::Entity::find_by_user(&state.db, auth.0.id)
.await
.unwrap()
.len();
if key_amount >= auth.0.stream_key_limit.try_into().unwrap() {
return StatusCode::NOT_ACCEPTABLE;
}
let key = stream_key::Entity::create(&state.db, auth.0.id, value, payload.label, false).await;
if let Ok(_key) = key {
@@ -117,9 +201,9 @@ async fn get_all_stream_keys(
State(state): State<Arc<HttpServer>>,
auth: AuthUser,
) -> impl IntoResponse {
let user_keys = stream_key::Entity::find_by_user(&state.db, auth.0.id)
.await
.unwrap();
// let user_keys = stream_key::Entity::find_by_user(&state.db, auth.0.id)
// .await
// .unwrap();
match stream_key::Entity::find_by_user(&state.db, auth.0.id).await {
Ok(keys) => Json(keys).into_response(),
@@ -141,12 +225,36 @@ struct LoginResponse {
async fn login_handler(
State(state): State<Arc<HttpServer>>,
Json(payload): Json<LoginForm>,
) -> StatusCode {
// TODO: hash payload.password with a real password hasher (argon2/bcrypt), look up user by username
// TODO: verify hash matches stored hashed_password
// TODO: call auth_session::Entity::create, return session token
// TODO: return 401 on bad credentials
todo!("login not yet implemented")
) -> impl IntoResponse {
if let Ok(x) = users::Entity::find_by_username(&state.db, payload.username).await {
if let Some(x) = x {
let pass = verify_password(&payload.password, &x.hashed_password);
if !pass {
let mut meow = Response::new("".to_string());
*meow.status_mut() = StatusCode::UNAUTHORIZED;
return meow;
};
let auth = auth_session::Entity::create(&state.db, x.id).await.unwrap(); // This should be ok (hopefully)
let token = auth.value;
let mut meow = Response::new("".to_string());
meow.headers_mut().insert(
SET_COOKIE,
format!("session={token}; HttpOnly; SameSite=Strict; Path=/")
.parse()
.unwrap(),
);
*meow.status_mut() = StatusCode::OK;
return meow;
} else {
let mut meow = Response::new("".to_string());
*meow.status_mut() = StatusCode::UNAUTHORIZED;
return meow;
}
} else {
let mut meow = Response::new("".to_string());
*meow.status_mut() = StatusCode::UNAUTHORIZED;
return meow;
}
}
#[derive(Deserialize)]
@@ -204,7 +312,7 @@ async fn stream_handler(
let app = state.appstate.lock().await;
app.stream_sessions
.iter()
.find(|e| e.value().stream_key_label == slug)
.find(|e| e.value().stream_key_id.to_string() == slug)
.map(|e| *e.key())
};
@@ -218,31 +326,41 @@ async fn stream_handler(
.unwrap();
};
let mut accept_rx = state.accept_rx.new_receiver();
let mut accept_rx = state.accept_rx.activate_cloned();
let _ = state
.offer_tx
.send((request_id_clone, stream_key_id, body))
.await
.unwrap();
debug!(
request_id = request_id_clone,
"offer sent, waiting for answer"
);
let mut reply_body = String::new();
while let Ok(answer) = accept_rx.recv().await {
debug!(request_id = request_id_clone, "received answer candidate");
if let Some(reply) = answer.1 {
if answer.0 == request_id_clone {
Response::builder()
return Response::builder()
.status(StatusCode::CREATED)
.header("content-type", "application/sdp")
.body(reply)
.unwrap();
} else {
continue;
}
} else {
Response::builder()
} else if answer.0 == request_id_clone {
return Response::builder()
.status(StatusCode::NOT_FOUND)
// .header("content-type", "application/sdp")
.body("")
.body(String::new())
.unwrap();
}
}
info!(
request_id = request_id_clone,
"answer channel closed without a match"
);
Response::builder()
.status(StatusCode::CREATED)
@@ -254,19 +372,3 @@ async fn stream_handler(
async fn meow_handler() -> &'static str {
"meow"
}
#[derive(Serialize)]
struct StreamCatalog {
active_streams: Vec<i32>,
}
// TODO: Make this return streak labels and their id instead of the whole stream key omfg
async fn catalog_from_state(state: &Arc<Mutex<AppState>>) -> StreamCatalog {
let app = state.lock().await;
let active_streams = app
.stream_sessions
.iter()
.map(|e| e.key().clone())
.collect();
StreamCatalog { active_streams }
}
+41 -135
View File
@@ -1,30 +1,41 @@
use std::{error::Error, sync::Arc};
use tracing::{level_filters::LevelFilter, warn};
use async_broadcast::broadcast;
use dashmap::DashMap;
use entity::stream_session;
use migration::{Migrator, MigratorTrait};
use rml_rtmp::{
handshake::{Handshake, HandshakeProcessResult, PeerType},
sessions::{ServerSession, ServerSessionConfig, ServerSessionEvent, ServerSessionResult},
};
use sea_orm::{Database, sqlx::types::chrono};
use sea_orm::{
Database, IntoActiveModel,
sqlx::types::chrono::{self, Local},
};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
sync::Mutex,
time::Instant,
};
use tracing_subscriber::EnvFilter;
use crate::{
audio::OpusAudioFrame,
http::HttpServer,
media::{H264Parser, VideoFrame},
};
mod audio;
mod hash;
mod http;
mod media;
mod rtmp;
mod webrtc;
mod webrtc_ingest;
// #[derive(Debug)]
pub struct AppState {
pub stream_sessions: Arc<DashMap<i32, StreamSession>>,
}
@@ -33,32 +44,44 @@ pub struct StreamSession {
pub stream_key_id: i32,
pub stream_key_label: String,
pub frame_channel: async_broadcast::Sender<Arc<VideoFrame>>,
pub audio_channel: async_broadcast::Sender<Arc<OpusAudioFrame>>,
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn Error>> {
let listener = TcpListener::bind("0.0.0.0:8123").await?;
let env_filter = EnvFilter::builder()
.with_default_directive(LevelFilter::DEBUG.into())
.parse("")
.unwrap();
tracing_subscriber::fmt().with_env_filter(env_filter).init();
let listener = TcpListener::bind("0.0.0.0:1935").await?;
let db = Database::connect("sqlite://./db.sqlite?mode=rwc")
.await
.unwrap();
Migrator::up(&db, None).await.unwrap();
stream_session::Model::clean_unended_streams(&db)
.await
.unwrap();
let appstate = Arc::new(Mutex::new(AppState {
stream_sessions: Arc::new(DashMap::new()),
}));
let (offer_tx, offer_rx) = tokio::sync::mpsc::channel::<(i32, i32, String)>(4);
let (offer_tx, offer_rx) = tokio::sync::mpsc::channel::<(i32, i32, String)>(64);
// Request_Id,
// String_Label,
// Offer_body
// --
// We are using broadcast instead of a regular channel because the way answers are processed are
// mulithreaded and we could easily process the answer of a completely different offer request
let (answer_tx, answer_rx) = broadcast::<(i32, Option<String>)>(4);
let (answer_tx, answer_rx) = broadcast::<(i32, Option<String>)>(64);
// Request_Id,
// Answer_body
let http = HttpServer {
offer_tx,
accept_rx: answer_rx,
accept_rx: answer_rx.deactivate(),
appstate: appstate.clone(),
request_count: Mutex::new(0),
db: db.clone(),
@@ -73,133 +96,16 @@ async fn main() -> Result<(), Box<dyn Error>> {
db: db.clone(),
};
webrtc.start()?;
let rtmp = rtmp::Rtmp {
db: db.clone(),
stream_sessions: app.stream_sessions.clone(),
listener: listener,
};
rtmp.start()?;
drop(app);
loop {
let (stream, _) = listener.accept().await?;
let appstate = appstate.clone();
let db = db.clone();
tokio::spawn(async move {
let mut stream = stream;
let mut server = Handshake::new(PeerType::Server);
let mut c0_c1: [u8; 1537] = [0; 1537];
stream.read_exact(&mut c0_c1).await.unwrap();
let s0_s1_s2 = server.process_bytes(&c0_c1);
let s0_s1_s2 = match s0_s1_s2 {
Ok(HandshakeProcessResult::InProgress {
response_bytes: bytes,
}) => bytes,
_ => panic!("handshake failed"),
};
stream.write_all(&s0_s1_s2).await.unwrap();
let mut c2 = vec![0u8; 1536];
stream.read_exact(&mut c2).await.unwrap();
match server.process_bytes(&c2[..]) {
Ok(HandshakeProcessResult::Completed { .. }) => {}
Ok(HandshakeProcessResult::InProgress {
response_bytes: meow,
}) => stream.write_all(&meow).await.unwrap(),
x => panic!("Unexpected process_bytes response: {:?}", x),
}
let config = ServerSessionConfig::new();
let (mut rtmp_session, bytes) = ServerSession::new(config).unwrap();
for x in bytes {
if let ServerSessionResult::OutboundResponse(packet) = x {
stream.write_all(&packet.bytes).await.unwrap();
}
}
let (mut video_channel, _video_rx) = broadcast::<Arc<VideoFrame>>(32);
video_channel.set_overflow(true);
let mut parser = H264Parser::new();
loop {
let mut buf: [u8; 4096] = [0; 4096];
let n = stream.read(&mut buf).await.unwrap();
let obs_events = rtmp_session.handle_input(&buf[..n]).unwrap();
for event in obs_events {
match event {
ServerSessionResult::OutboundResponse(packet) => {
stream.write_all(&packet.bytes).await.unwrap();
}
ServerSessionResult::RaisedEvent(x) => match x {
ServerSessionEvent::PublishStreamFinished { stream_key, .. } => {
let key = entity::stream_key::Entity::find_by_key(&db, &stream_key)
.await
.unwrap()
.unwrap();
appstate.lock().await.stream_sessions.remove(&key.id);
}
ServerSessionEvent::PublishStreamRequested {
request_id,
stream_key,
..
} => {
let key =
entity::stream_key::Entity::find_by_key(&db, &stream_key).await;
let key = if let Ok(Some(key)) = key {
key
} else {
// Early return, gives error to client
let reply = rtmp_session
.reject_request(request_id, "", "Stream key invalid")
.unwrap();
for x in reply {
if let ServerSessionResult::OutboundResponse(y) = x {
stream.write_all(&y.bytes).await.unwrap();
}
}
break;
};
let session = StreamSession {
stream_key_id: key.id,
stream_key_label: key.label,
frame_channel: video_channel.clone(),
};
appstate
.lock()
.await
.stream_sessions
.insert(key.id, session);
let reply = rtmp_session.accept_request(request_id).unwrap();
for x in reply {
if let ServerSessionResult::OutboundResponse(y) = x {
stream.write_all(&y.bytes).await.unwrap();
}
}
}
ServerSessionEvent::ConnectionRequested { request_id, .. } => {
let reply = rtmp_session.accept_request(request_id).unwrap();
for x in reply {
if let ServerSessionResult::OutboundResponse(y) = x {
stream.write_all(&y.bytes).await.unwrap();
}
}
}
ServerSessionEvent::VideoDataReceived {
data, timestamp, ..
} => {
if data.len() >= 5 && &data[1..5] == b"hvc1" {
println!("HEVC/H.265 not supported, closing connection");
return;
}
if let Some(parsed_frame) = parser.parse(&data, timestamp.value) {
video_channel.broadcast(Arc::new(parsed_frame)).await.ok();
}
}
_ => {}
},
_ => {}
}
}
}
});
}
tokio::signal::ctrl_c().await?;
Ok(())
}
+225 -39
View File
@@ -1,52 +1,238 @@
#[derive(Default, Debug)]
pub struct RtmpSession {
client: RtmpClient,
use std::{error::Error, sync::Arc};
use async_broadcast::broadcast;
use dashmap::DashMap;
use entity::stream_session;
use rml_rtmp::{
handshake::{Handshake, HandshakeProcessResult, PeerType},
sessions::{ServerSession, ServerSessionConfig, ServerSessionEvent, ServerSessionResult},
};
use sea_orm::{DatabaseConnection, IntoActiveModel, sqlx::types::chrono::Local};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use tracing::warn;
use crate::{
StreamSession,
audio::{AACParser, AudioProcesser, OpusAudioFrame},
media::{H264Parser, VideoFrame},
};
pub struct Rtmp {
pub listener: TcpListener,
pub stream_sessions: Arc<DashMap<i32, StreamSession>>,
pub db: DatabaseConnection,
}
#[derive(Debug)]
struct RtmpClient {
version: u8,
timestamp: u32,
magic_bytes: [u8; 1536],
}
impl Default for RtmpClient {
fn default() -> Self {
Self {
magic_bytes: [0; 1536],
version: 0,
timestamp: 0,
async fn write_outbound(socket: &mut TcpStream, results: Vec<ServerSessionResult>) {
for r in results {
if let ServerSessionResult::OutboundResponse(p) = r {
socket.write_all(&p.bytes).await.unwrap();
}
}
}
impl RtmpSession {
pub fn consume_handshake(&mut self, bytes: &[u8]) -> Result<(), ()> {
let version = bytes[0];
let timestamp = u32::from_be_bytes(bytes[1..5].try_into().unwrap());
let mut random: [u8; 1536] = [0; 1536];
random.copy_from_slice(&bytes[1..1537]);
impl Rtmp {
async fn handshake(
mut socket: TcpStream,
) -> Result<(ServerSession, TcpStream), Box<dyn Error>> {
let mut server = Handshake::new(PeerType::Server);
println!("size of rand: {}", random.len());
let client = RtmpClient {
version,
timestamp,
magic_bytes: random,
let mut c0_c1 = [0u8; 1537];
socket.read_exact(&mut c0_c1).await.unwrap();
let s0_s1_s2 = match server.process_bytes(&c0_c1) {
Ok(HandshakeProcessResult::InProgress { response_bytes }) => response_bytes,
_ => panic!("handshake failed"),
};
socket.write_all(&s0_s1_s2).await.unwrap();
self.client = client;
let mut c2 = [0u8; 1536];
socket.read_exact(&mut c2).await.unwrap();
match server.process_bytes(&c2) {
Ok(HandshakeProcessResult::Completed { .. }) => {}
Ok(HandshakeProcessResult::InProgress { response_bytes }) => {
socket.write_all(&response_bytes).await.unwrap()
}
x => panic!("Unexpected process_bytes response: {:?}", x),
}
let (rtmp_session, init_bytes) = ServerSession::new(ServerSessionConfig::new()).unwrap();
write_outbound(&mut socket, init_bytes).await;
Ok((rtmp_session, socket))
}
pub fn start(self) -> Result<(), Box<dyn Error>> {
let Self {
listener,
stream_sessions,
db,
} = self;
tokio::spawn(async move {
loop {
let (socket, _) = listener.accept().await.unwrap();
let (mut session, mut socket) = Rtmp::handshake(socket).await.unwrap();
let (mut video_tx, mut video_rx) = broadcast::<Arc<VideoFrame>>(32);
let (mut audio_tx, mut audio_rx) = broadcast::<Arc<OpusAudioFrame>>(32);
// video_rx.cycle
video_tx.set_overflow(true);
audio_tx.set_overflow(true);
let mut parser = H264Parser::new();
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 video_rx = video_rx;
loop {
let mut buf = [0u8; 4096];
let n = socket.read(&mut buf).await.unwrap();
let events = session.handle_input(&buf[..n]).unwrap();
// Blankly using it, so it doesnt drop
video_rx.is_closed();
audio_rx.is_closed();
for event in events {
match event {
ServerSessionResult::OutboundResponse(p) => {
socket.write_all(&p.bytes).await.unwrap();
}
ServerSessionResult::RaisedEvent(e) => match e {
ServerSessionEvent::ConnectionRequested {
request_id, ..
} => {
let reply = session.accept_request(request_id).unwrap();
write_outbound(&mut socket, reply).await;
}
ServerSessionEvent::PublishStreamRequested {
request_id,
stream_key,
..
} => {
let key = entity::stream_key::Entity::find_by_key(
&db,
&stream_key,
)
.await;
let key = if let Ok(Some(key)) = key {
key
} else {
let reply = session
.reject_request(
request_id,
"",
"Stream key invalid",
)
.unwrap();
write_outbound(&mut socket, reply).await;
break;
};
let already_live =
stream_session::Model::get_active_by_stream_key_id(
&db, key.id,
)
.await
.unwrap();
if already_live.is_some() {
let reply = session
.reject_request(
request_id,
"",
"You're already streaming...",
)
.unwrap();
write_outbound(&mut socket, reply).await;
break;
}
stream_sessions.insert(
key.id,
StreamSession {
stream_key_id: key.id,
stream_key_label: key.label,
frame_channel: video_tx.clone(),
audio_channel: audio_tx.clone(),
},
);
let reply = session.accept_request(request_id).unwrap();
stream_session::Model::create_stream_session(
&db,
key.id,
Local::now().into(),
)
.await
.unwrap();
write_outbound(&mut socket, reply).await;
}
ServerSessionEvent::PublishStreamFinished {
stream_key,
..
} => {
let key = entity::stream_key::Entity::find_by_key(
&db,
&stream_key,
)
.await
.unwrap()
.unwrap();
stream_sessions.remove(&key.id);
stream_session::Model::get_active_by_stream_key_id(
&db, key.id,
)
.await
.unwrap()
.unwrap()
.into_active_model()
.finish_stream_session(&db, Local::now().into())
.await
.unwrap();
stream_sessions.get(&key.id).unwrap().frame_channel.close();
}
ServerSessionEvent::VideoDataReceived {
data,
timestamp,
..
} => {
if data.len() >= 5 && &data[1..5] == b"hvc1" {
warn!("HEVC/H.265 not supported, closing connection");
return;
}
if let Some(frame) = parser.parse(&data, timestamp.value) {
video_tx.broadcast(Arc::new(frame)).await.ok();
}
}
ServerSessionEvent::AudioDataReceived {
data,
timestamp,
..
} => {
// Consume the non-Send error before any await point.
let opus_frames: Vec<_> = aac_parser
.parse(&data, timestamp.value)
.ok()
.flatten()
.map(|f| audio_proc.encode(f))
.unwrap_or_default();
for frame in opus_frames {
audio_tx.broadcast(Arc::new(frame)).await.ok();
}
}
_ => {}
},
_ => {}
}
}
}
});
}
});
Ok(())
}
pub fn response(&self) -> [u8; 1537] {
let mut reply: [u8; 1537] = [0; 1537];
reply[0] = 3;
// reply[1..1537].copy_from_slice(&self.client.magic_bytes);
let mut rand: [u8; 1536] = [0; 1536];
rand.fill(1);
reply[1..1537].copy_from_slice(&rand);
reply
}
}
+172 -36
View File
@@ -1,6 +1,6 @@
use dashmap::DashMap;
use entity::stream_key;
use sea_orm::DatabaseConnection;
use sea_orm::{DatabaseConnection, EntityTrait};
use std::{
error::Error,
sync::Arc,
@@ -10,15 +10,16 @@ use tokio::{
net::UdpSocket,
sync::mpsc::{Receiver, Sender},
};
use tracing::{debug, error, info, warn};
use str0m::{
Candidate, Event, Input, Output, Rtc,
change::SdpOffer,
media::{MediaKind, MediaTime, Mid},
media::{Frequency, MediaKind, MediaTime, Mid},
net::{Protocol, Receive},
};
use crate::{StreamSession, media::VideoFrame};
use crate::{StreamSession, audio::OpusAudioFrame, media::VideoFrame};
pub struct Webrtc {
pub offer_rx: Receiver<(i32, i32, String)>,
@@ -34,9 +35,11 @@ impl Webrtc {
tokio::spawn(async move {
while let Some(offer) = self.offer_rx.recv().await {
let (request_id, stream_id, sdp_body) = offer;
info!(request_id, stream_id, "processing offer");
let stream_key =
stream_key::Entity::find_by_key(&self.db, &stream_id.to_string()).await;
let stream_key = stream_key::Entity::find_by_id(stream_id)
.one(&self.db)
.await;
if let Ok(None) = stream_key {
self.accept_tx.broadcast((request_id, None)).await.unwrap();
@@ -65,18 +68,32 @@ impl Webrtc {
MediaKind::Video,
str0m::media::Direction::SendOnly,
Some(stream_id.to_string()),
Some("video0".to_string()),
Some("0".to_string()),
None,
);
changes.add_media(
MediaKind::Audio,
str0m::media::Direction::SendOnly,
Some(stream_id.to_string()),
Some("0".to_string()),
None,
);
let offer_answer = match changes.accept_offer(offer_sdp) {
Ok(a) => a,
Err(e) => {
println!("accept_offer failed: {:?}", e);
error!("accept_offer failed: {:?}", e);
continue;
}
};
let answer_sdp = offer_answer.to_sdp_string();
let ufrag = answer_sdp
.lines()
.find(|l| l.starts_with("a=ice-ufrag:"))
.and_then(|l| l.strip_prefix("a=ice-ufrag:"))
.map(|s| s.trim().to_string());
info!("Serving webrtc ufrag: {:?}", ufrag);
debug!(request_id, "sending answer back");
self.accept_tx
.broadcast((request_id, Some(answer_sdp)))
.await
@@ -100,8 +117,12 @@ impl Webrtc {
) {
let mut video_mid: Option<Mid> = None;
let mut video_pt = None;
let mut audio_mid: Option<Mid> = None;
let mut audio_pt = None;
let mut connected = false;
let mut video_stream: Option<async_broadcast::Receiver<Arc<VideoFrame>>> = None;
let mut audio_stream: Option<async_broadcast::Receiver<Arc<OpusAudioFrame>>> = None;
let mut saw_keyframe = false;
let mut recv_buf = vec![0u8; PER_CLIENT_CONNECTION_BUF];
let local_addr = socket.local_addr().unwrap();
@@ -111,7 +132,8 @@ impl Webrtc {
match rtc.poll_output() {
Ok(Output::Timeout(t)) => break t,
Ok(Output::Transmit(t)) => {
if socket.send_to(&t.contents, t.destination).await.is_err() {
if let Err(e) = socket.send_to(&t.contents, t.destination).await {
warn!("UDP send error: {:?}", e);
return;
}
}
@@ -123,23 +145,43 @@ impl Webrtc {
p.spec().format.profile_level_id.unwrap_or(0)
});
if let Some(params) = best {
println!("Selected PT {:?}", params.pt());
info!("selected PT {:?}", params.pt());
video_pt = Some(params.pt());
video_mid = Some(ma.mid);
}
}
}
if ma.kind == MediaKind::Audio {
if let Some(writer) = rtc.writer(ma.mid) {
let best = writer.payload_params().max_by_key(|p| {
p.spec().format.profile_level_id.unwrap_or(0)
});
if let Some(params) = best {
info!("selected PT {:?}", params.pt());
audio_pt = Some(params.pt());
audio_mid = Some(ma.mid);
}
}
}
}
Event::IceConnectionStateChange(state) => {
println!("ICE state: {:?}", state);
use str0m::IceConnectionState;
info!("ICE state: {:?}", state);
if matches!(state, IceConnectionState::Disconnected) {
info!("ICE disconnected, closing connection");
return;
}
}
Event::Connected => {
println!("DTLS+ICE connected, ready for media");
info!("DTLS+ICE connected, ready for media");
connected = true;
}
_ => {}
},
Err(_) => return,
Err(e) => {
error!("poll_output error (connection closing): {:?}", e);
return;
}
}
};
@@ -147,33 +189,122 @@ impl Webrtc {
if video_stream.is_none() {
if let Some(session) = sessions_ref.get(&stream_id) {
video_stream = Some(session.frame_channel.new_receiver());
debug!(stream_id, "subscribed to video channel");
} else {
warn!(
stream_id,
"stream session not found in map, cannot subscribe"
);
}
}
if let Some(ref mut stream) = video_stream {
for _ in 0..8 {
match stream.try_recv() {
Ok(frame) => {
let now = Instant::now();
let rtp_time =
MediaTime::from_90khz(frame.timestamp_ms as u64 * 90);
if let (Some(pt), Some(writer)) =
(video_pt, video_mid.and_then(|m| rtc.writer(m)))
{
if let Err(e) =
writer.write(pt, now, rtp_time, frame.data.to_vec())
{
println!("write error: {:?}", e);
}
}
}
Err(async_broadcast::TryRecvError::Empty) => break,
Err(async_broadcast::TryRecvError::Closed) => return,
Err(async_broadcast::TryRecvError::Overflowed(_)) => continue,
if let Some(video) = &video_stream {
if video.is_closed() {
if let Some(session) = sessions_ref.get(&stream_id) {
video_stream = Some(session.frame_channel.new_receiver());
debug!(stream_id, "subscribed to video channel");
}
}
}
if audio_stream.is_none() {
if let Some(session) = sessions_ref.get(&stream_id) {
audio_stream = Some(session.audio_channel.new_receiver());
debug!(stream_id, "subscribed to audio channel");
}
}
if let Some(audio) = &audio_stream {
if audio.is_closed() {
if let Some(session) = sessions_ref.get(&stream_id) {
audio_stream = Some(session.audio_channel.new_receiver());
debug!(stream_id, "subscribed to audio channel");
}
}
}
if let Some(ref mut stream) = video_stream {
let mut wrote_any = false;
for _ in 0..8 {
match stream.try_recv() {
Ok(frame) => {
if !saw_keyframe {
if !frame.is_keyframe {
continue;
}
saw_keyframe = true;
debug!(stream_id, "first keyframe, starting RTP send");
}
let now = Instant::now();
let rtp_time =
MediaTime::from_90khz(frame.timestamp_ms as u64 * 90);
match (video_pt, video_mid.and_then(|m| rtc.writer(m))) {
(Some(pt), Some(writer)) => {
match writer.write(pt, now, rtp_time, frame.data.to_vec()) {
Ok(_) => wrote_any = true,
Err(e) => warn!("video RTP write error: {:?}", e),
}
}
_ => warn!(
stream_id,
"video_pt or writer not ready, dropping frame (video_pt={:?} video_mid={:?})",
video_pt,
video_mid
),
}
}
Err(async_broadcast::TryRecvError::Empty) => break,
Err(async_broadcast::TryRecvError::Closed) => {
info!("video channel closed, stream ended");
break;
}
Err(async_broadcast::TryRecvError::Overflowed(_)) => {
saw_keyframe = false;
continue;
}
}
}
if let Some(ref mut audio) = audio_stream {
for _ in 0..8 {
match audio.try_recv() {
Ok(frame) => {
let now = Instant::now();
let rtp_time = MediaTime::new(
frame.timestamp_ms as u64 * 48,
Frequency::FORTY_EIGHT_KHZ,
);
if let (Some(pt), Some(writer)) =
(audio_pt, audio_mid.and_then(|m| rtc.writer(m)))
{
match writer.write(pt, now, rtp_time, frame.data.to_vec()) {
Ok(_) => wrote_any = true,
Err(e) => warn!("RTP write error: {:?}", e),
}
}
}
Err(async_broadcast::TryRecvError::Empty) => break,
Err(async_broadcast::TryRecvError::Closed) => {
info!("audio channel closed, stream ended");
break;
}
Err(async_broadcast::TryRecvError::Overflowed(_)) => continue,
}
}
}
// If we queued RTP data, drain poll_output immediately so packets are
// transmitted in this iteration rather than 20ms later. But still drive
// str0m's timeout if the deadline has passed — ICE consent refresh depends on it.
if wrote_any {
let now = Instant::now();
if now >= deadline {
if let Err(e) = rtc.handle_input(Input::Timeout(now)) {
error!("handle_input(Timeout) error: {:?}", e);
return;
}
}
continue;
}
}
}
// debug!("waiting");
let wait_until = deadline
.min(Instant::now() + Duration::from_millis(20))
.max(Instant::now());
@@ -181,13 +312,16 @@ impl Webrtc {
tokio::select! {
_ = sleep => {
rtc.handle_input(Input::Timeout(Instant::now())).ok();
if let Err(e) = rtc.handle_input(Input::Timeout(Instant::now())) {
error!("handle_input(Timeout) error: {:?}", e);
return;
}
}
result = socket.recv_from(&mut recv_buf) => {
if let Ok((n, from)) = result {
let data = recv_buf[..n].to_vec();
if let Ok(contents) = data.as_slice().try_into() {
rtc.handle_input(Input::Receive(
if let Err(e) = rtc.handle_input(Input::Receive(
Instant::now(),
Receive {
proto: Protocol::Udp,
@@ -195,8 +329,10 @@ impl Webrtc {
destination: local_addr,
contents,
},
))
.ok();
)) {
error!("handle_input(Receive) error: {:?}", e);
return;
}
}
}
}
+17
View File
@@ -0,0 +1,17 @@
use std::sync::Arc;
use axum::{extract::State, response::IntoResponse};
use str0m::change::SdpOffer;
use tracing::info;
use crate::http::HttpServer;
pub async fn handle_whip_injest(
State(state): State<Arc<HttpServer>>,
offer: String,
) -> impl IntoResponse {
let sdp_offer = SdpOffer::from_sdp_string(&offer).unwrap();
for x in &sdp_offer.media_lines {
info!("{}", x);
}
}