diff --git a/.gitignore b/.gitignore index 91a15d8..54321a2 100644 --- a/.gitignore +++ b/.gitignore @@ -11,3 +11,4 @@ devenv.local.yaml # pre-commit .pre-commit-config.yaml /stream.db +/db.sqlite diff --git a/Cargo.lock b/Cargo.lock index 7f54fbb..8cebc65 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2866,7 +2866,7 @@ dependencies = [ [[package]] name = "server" -version = "0.1.0" +version = "0.1.1" dependencies = [ "argon2", "async-broadcast", diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..6786fe5 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,46 @@ +FROM rust:1-bookworm AS builder + +RUN apt-get update && apt-get install -y \ + libopus-dev \ + libsqlite3-dev \ + clang \ + && rm -rf /var/lib/apt/lists/* + +WORKDIR /app + +COPY Cargo.toml Cargo.lock ./ +COPY crates/server/Cargo.toml crates/server/Cargo.toml +COPY crates/entity/Cargo.toml crates/entity/Cargo.toml +COPY crates/migration/Cargo.toml crates/migration/Cargo.toml + +RUN mkdir -p crates/server/src crates/entity/src crates/migration/src \ + && echo "fn main() {}" > crates/server/src/main.rs \ + && echo "" > crates/entity/src/lib.rs \ + && echo "" > crates/migration/src/lib.rs + +ENV RUSTFLAGS="" + +RUN cargo build --release --bin rtmp-to-whip + +COPY crates crates +RUN touch crates/server/src/main.rs crates/entity/src/lib.rs crates/migration/src/lib.rs \ + && cargo build --release --bin rtmp-to-whip \ + && strip target/release/rtmp-to-whip + +RUN mkdir /deps \ + && ldd target/release/rtmp-to-whip \ + | awk 'NF==4{print $3} NF==2{print $1}' \ + | grep -v vdso \ + | xargs -I{} cp --parents {} /deps + +FROM scratch + +COPY --from=builder /deps / +COPY --from=builder /etc/ssl/certs /etc/ssl/certs +COPY --from=builder /app/target/release/rtmp-to-whip /rtmp-to-whip + +EXPOSE 1935 +EXPOSE 3000 +EXPOSE 6969/udp + +CMD ["/rtmp-to-whip"] diff --git a/Dockerfile.aarch64 b/Dockerfile.aarch64 new file mode 100644 index 0000000..d623ec8 --- /dev/null +++ b/Dockerfile.aarch64 @@ -0,0 +1,58 @@ +FROM --platform=linux/amd64 rust:1-bookworm AS builder + +RUN dpkg --add-architecture arm64 \ + && apt-get update && apt-get install -y \ + gcc-aarch64-linux-gnu \ + libopus-dev:arm64 \ + libsqlite3-dev:arm64 \ + libc6:arm64 \ + ca-certificates \ + && rm -rf /var/lib/apt/lists/* + +RUN rustup target add aarch64-unknown-linux-gnu + +ENV CARGO_TARGET_AARCH64_UNKNOWN_LINUX_GNU_LINKER=aarch64-linux-gnu-gcc \ + PKG_CONFIG_ALLOW_CROSS=1 \ + PKG_CONFIG_PATH=/usr/lib/aarch64-linux-gnu/pkgconfig \ + RUSTFLAGS="" + +WORKDIR /app + +COPY Cargo.toml Cargo.lock ./ +COPY crates/server/Cargo.toml crates/server/Cargo.toml +COPY crates/entity/Cargo.toml crates/entity/Cargo.toml +COPY crates/migration/Cargo.toml crates/migration/Cargo.toml + +RUN mkdir -p crates/server/src crates/entity/src crates/migration/src \ + && echo "fn main() {}" > crates/server/src/main.rs \ + && echo "" > crates/entity/src/lib.rs \ + && echo "" > crates/migration/src/lib.rs + +RUN cargo build --release --target aarch64-unknown-linux-gnu --bin rtmp-to-whip + +COPY crates crates +RUN touch crates/server/src/main.rs crates/entity/src/lib.rs crates/migration/src/lib.rs \ + && cargo build --release --target aarch64-unknown-linux-gnu --bin rtmp-to-whip \ + && aarch64-linux-gnu-strip target/aarch64-unknown-linux-gnu/release/rtmp-to-whip + +RUN mkdir /deps \ + && find /usr/lib/aarch64-linux-gnu -name 'libopus.so*' -exec cp --parents {} /deps \; \ + && find /usr/lib/aarch64-linux-gnu -name 'libsqlite3.so*' -exec cp --parents {} /deps \; \ + && find /lib/aarch64-linux-gnu -name 'libc*' -exec cp --parents {} /deps \; \ + && find /lib/aarch64-linux-gnu -name 'libm*' -exec cp --parents {} /deps \; \ + && find /lib/aarch64-linux-gnu -name 'libpthread*' -exec cp --parents {} /deps \; \ + && find /lib/aarch64-linux-gnu -name 'libdl*' -exec cp --parents {} /deps \; \ + && find /lib/aarch64-linux-gnu -name 'libgcc_s*' -exec cp --parents {} /deps \; \ + && cp --parents /lib/ld-linux-aarch64.so.1 /deps + +FROM --platform=linux/arm64/v8 scratch + +COPY --from=builder /deps / +COPY --from=builder /etc/ssl/certs /etc/ssl/certs +COPY --from=builder /app/target/aarch64-unknown-linux-gnu/release/rtmp-to-whip /rtmp-to-whip + +EXPOSE 1935 +EXPOSE 3000 +EXPOSE 6969/udp + +CMD ["/rtmp-to-whip"] diff --git a/Justfile b/Justfile new file mode 100644 index 0000000..1eece5a --- /dev/null +++ b/Justfile @@ -0,0 +1,17 @@ +build-dev: + docker build ./ --tag reg.h.doloro.co.uk/doloro/rtmp-to-whip:dev + +push-dev: + docker push reg.h.doloro.co.uk/doloro/rtmp-to-whip:dev + +build-latest: + docker build ./ --tag reg.h.doloro.co.uk/doloro/rtmp-to-whip:latest + +push-latest: build-latest + docker push reg.h.doloro.co.uk/doloro/rtmp-to-whip:latest + +build-aarch64: + docker build ./ --file Dockerfile.aarch64 --tag reg.h.doloro.co.uk/doloro/rtmp-to-whip:latest-aarch64 + +push-aarch64: build-aarch64 + docker push reg.h.doloro.co.uk/doloro/rtmp-to-whip:latest-aarch64 diff --git a/compose.yaml b/compose.yaml new file mode 100644 index 0000000..7a2f304 --- /dev/null +++ b/compose.yaml @@ -0,0 +1,17 @@ +services: + server: + build: . + ports: + - "1935:1935" # RTMP + - "3000:3000" # HTTP API + - "6969:6969/udp" # WebRTC UDP proxy + environment: + - PUBLIC_DOMAIN= # set to your domain, e.g. rtmp.example.com; falls back to STUN if unset + - RTC_PORT=6969 + - SIGNUP_CODE= # required to create accounts; leave empty to disable signup + volumes: + - db:/app/db.sqlite + restart: unless-stopped + +volumes: + db: diff --git a/crates/server/Cargo.toml b/crates/server/Cargo.toml index 5e83237..430292c 100644 --- a/crates/server/Cargo.toml +++ b/crates/server/Cargo.toml @@ -1,7 +1,8 @@ [package] name = "server" -version = "0.1.0" +version = "0.1.1" edition = "2024" + [target.x86_64-unknown-linux-gnu] linker = "clang" rustflags = ["-Clink-arg=-fuse-ld=/usr/local/bin/mold", "-Clink-arg=-Wl,--no-rosegment"] diff --git a/crates/server/src/http.rs b/crates/server/src/http.rs index 3ecb5fc..e6dd92e 100644 --- a/crates/server/src/http.rs +++ b/crates/server/src/http.rs @@ -22,7 +22,7 @@ use tokio::{ use tower_http::cors::CorsLayer; use uuid::Uuid; -use tracing::{debug, info}; +use tracing::{debug, info, warn}; use crate::{ AppState, @@ -30,12 +30,20 @@ use crate::{ webrtc_ingest::handle_whip_injest, }; +const MAX_USERNAME_LEN: usize = 32; +const MAX_LABEL_LEN: usize = 64; + +pub struct HttpServerConfig { + pub signup_code: String, +} + pub struct HttpServer { pub offer_tx: Sender<(i32, i32, String)>, pub accept_rx: async_broadcast::InactiveReceiver<(i32, Option)>, pub appstate: Arc>, pub request_count: Mutex, pub db: DatabaseConnection, + pub config: HttpServerConfig, } impl HttpServer { @@ -73,7 +81,7 @@ impl HttpServer { .layer(cors) .with_state(state); - let addr = SocketAddr::from(([127, 0, 0, 1], 3000)); + let addr = SocketAddr::from(([0, 0, 0, 0], 3000)); let listener = TcpListener::bind(addr).await.unwrap(); axum::serve(listener, app).await.unwrap(); }); @@ -175,12 +183,18 @@ async fn create_stream_key_handler( .await .unwrap() .len(); + if payload.label.is_empty() || payload.label.len() > MAX_LABEL_LEN { + warn!(user_id = auth.0.id, label_len = payload.label.len(), max = MAX_LABEL_LEN, "stream key creation rejected: label length invalid"); + return StatusCode::UNPROCESSABLE_ENTITY; + } if key_amount >= auth.0.stream_key_limit.try_into().unwrap() { + warn!(user_id = auth.0.id, limit = auth.0.stream_key_limit, "stream key limit reached"); 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 { + if let Ok(ref k) = key { + info!(user_id = auth.0.id, stream_key_id = k.id, label = %k.label, "stream key created"); StatusCode::CREATED } else { StatusCode::INTERNAL_SERVER_ERROR @@ -226,15 +240,17 @@ async fn login_handler( State(state): State>, Json(payload): Json, ) -> impl IntoResponse { - if let Ok(x) = users::Entity::find_by_username(&state.db, payload.username).await { + if let Ok(x) = users::Entity::find_by_username(&state.db, payload.username.clone()).await { if let Some(x) = x { let pass = verify_password(&payload.password, &x.hashed_password); if !pass { + warn!(username = %payload.username, "login failed: wrong password"); 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) + info!(user_id = x.id, username = %x.username, "login successful"); + let auth = auth_session::Entity::create(&state.db, x.id).await.unwrap(); let token = auth.value; let mut meow = Response::new("".to_string()); meow.headers_mut().insert( @@ -246,11 +262,13 @@ async fn login_handler( *meow.status_mut() = StatusCode::OK; return meow; } else { + warn!(username = %payload.username, "login failed: user not found"); let mut meow = Response::new("".to_string()); *meow.status_mut() = StatusCode::UNAUTHORIZED; return meow; } } else { + warn!(username = %payload.username, "login failed: DB error"); let mut meow = Response::new("".to_string()); *meow.status_mut() = StatusCode::UNAUTHORIZED; return meow; @@ -268,23 +286,30 @@ async fn create_user_handler( State(state): State>, Json(payload): Json, ) -> (HeaderMap, StatusCode) { - if payload.ref_token != "TEST" { + if state.config.signup_code.is_empty() || payload.ref_token != state.config.signup_code { + warn!(username = %payload.username, "signup rejected: invalid signup code"); return (HeaderMap::new(), StatusCode::UNAUTHORIZED); } + if payload.username.is_empty() || payload.username.len() > MAX_USERNAME_LEN { + warn!(username = %payload.username, max = MAX_USERNAME_LEN, "signup rejected: username length invalid"); + return (HeaderMap::new(), StatusCode::UNPROCESSABLE_ENTITY); + } let meow = users::Entity::create( &state.db, - payload.username, + payload.username.clone(), hash_password(&payload.password).unwrap(), ) .await; // Create session - let session: auth_session::Model = if let Ok(meow) = meow { - auth_session::Entity::create(&state.db, meow.id) + let session: auth_session::Model = if let Ok(ref user) = meow { + info!(user_id = user.id, username = %user.username, "user created"); + auth_session::Entity::create(&state.db, user.id) .await - .unwrap() // This should be ok (hopefully) + .unwrap() } else { + warn!(username = %payload.username, "user creation failed (likely username conflict)"); return (HeaderMap::new(), StatusCode::CONFLICT); }; let token = session.value; @@ -319,6 +344,7 @@ async fn stream_handler( let stream_key_id = if let Some(id) = stream_key_id { id } else { + warn!(slug = %slug, "WHEP request for unknown or inactive stream"); return Response::builder() .status(StatusCode::NOT_FOUND) .header("content-type", "application/text") @@ -326,6 +352,7 @@ async fn stream_handler( .unwrap(); }; + info!(request_id = request_id_clone, slug = %slug, stream_key_id, "WHEP offer received"); let mut accept_rx = state.accept_rx.activate_cloned(); let _ = state .offer_tx diff --git a/crates/server/src/main.rs b/crates/server/src/main.rs index 9a8768a..f57bf9f 100644 --- a/crates/server/src/main.rs +++ b/crates/server/src/main.rs @@ -1,5 +1,5 @@ -use std::{error::Error, sync::Arc}; -use tracing::{level_filters::LevelFilter, warn}; +use std::{env, error::Error, sync::Arc}; +use tracing::{info, level_filters::LevelFilter, warn}; use async_broadcast::broadcast; use dashmap::DashMap; @@ -23,8 +23,9 @@ use tracing_subscriber::EnvFilter; use crate::{ audio::OpusAudioFrame, - http::HttpServer, + http::{HttpServer, HttpServerConfig}, media::{H264Parser, VideoFrame}, + webrtc_proxy::WebRtcProxyConfig, }; mod audio; @@ -34,6 +35,7 @@ mod media; mod rtmp; mod webrtc; mod webrtc_ingest; +mod webrtc_proxy; // #[derive(Debug)] pub struct AppState { @@ -51,17 +53,22 @@ pub struct StreamSession { async fn main() -> Result<(), Box> { let env_filter = EnvFilter::builder() .with_default_directive(LevelFilter::DEBUG.into()) - .parse("") + .from_env() + // .parse("") .unwrap(); tracing_subscriber::fmt().with_env_filter(env_filter).init(); let listener = TcpListener::bind("0.0.0.0:1935").await?; + info!("RTMP listening on 0.0.0.0:1935"); + info!("HTTP API listening on 0.0.0.0:3000"); - let db = Database::connect("sqlite://./db.sqlite?mode=rwc") + let db = Database::connect("sqlite://./db/db.sqlite?mode=rwc") .await .unwrap(); + info!("database connected"); Migrator::up(&db, None).await.unwrap(); + info!("migrations complete"); stream_session::Model::clean_unended_streams(&db) .await @@ -85,15 +92,31 @@ async fn main() -> Result<(), Box> { appstate: appstate.clone(), request_count: Mutex::new(0), db: db.clone(), + config: HttpServerConfig { + signup_code: env::var("SIGNUP_CODE").unwrap_or_else(|_| { + warn!("SIGNUP_CODE not set; signup will be disabled"); + String::new() + }), + }, }; http.start()?; + let proxyconfig = WebRtcProxyConfig { + proxy_port: env::var("RTC_PORT") + .unwrap_or("6969".into()) + .parse() + .expect("RTC_PORT needs to be a number (i32)"), + }; + let proxy = webrtc_proxy::WebrtcProxy::new(proxyconfig).await.unwrap(); + proxy.start().unwrap(); + let app = appstate.lock().await; let webrtc = webrtc::Webrtc { offer_rx, accept_tx: answer_tx, sessions_ref: app.stream_sessions.clone(), db: db.clone(), + proxy: proxy.into(), }; webrtc.start()?; diff --git a/crates/server/src/rtmp.rs b/crates/server/src/rtmp.rs index e747640..e57d172 100644 --- a/crates/server/src/rtmp.rs +++ b/crates/server/src/rtmp.rs @@ -12,7 +12,7 @@ use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::{TcpListener, TcpStream}, }; -use tracing::warn; +use tracing::{debug, info, warn}; use crate::{ StreamSession, @@ -72,8 +72,10 @@ impl Rtmp { } = self; tokio::spawn(async move { loop { - let (socket, _) = listener.accept().await.unwrap(); + let (socket, peer_addr) = listener.accept().await.unwrap(); + info!(%peer_addr, "RTMP connection accepted"); let (mut session, mut socket) = Rtmp::handshake(socket).await.unwrap(); + info!(%peer_addr, "RTMP handshake complete"); let (mut video_tx, mut video_rx) = broadcast::>(32); let (mut audio_tx, mut audio_rx) = broadcast::>(32); // video_rx.cycle @@ -87,10 +89,13 @@ impl Rtmp { 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(); + if n == 0 { + debug!("RTMP connection closed by peer"); + return; + } let events = session.handle_input(&buf[..n]).unwrap(); // Blankly using it, so it doesnt drop video_rx.is_closed(); @@ -105,6 +110,7 @@ impl Rtmp { ServerSessionEvent::ConnectionRequested { request_id, .. } => { + debug!("RTMP ConnectionRequested, accepting"); let reply = session.accept_request(request_id).unwrap(); write_outbound(&mut socket, reply).await; } @@ -113,6 +119,7 @@ impl Rtmp { stream_key, .. } => { + info!(stream_key = %stream_key, "publish stream requested"); let key = entity::stream_key::Entity::find_by_key( &db, &stream_key, @@ -122,6 +129,7 @@ impl Rtmp { let key = if let Ok(Some(key)) = key { key } else { + warn!(stream_key = %stream_key, "stream key not found, rejecting"); let reply = session .reject_request( request_id, @@ -140,6 +148,7 @@ impl Rtmp { .await .unwrap(); if already_live.is_some() { + warn!(stream_key_id = key.id, label = %key.label, "stream key already live, rejecting duplicate publish"); let reply = session .reject_request( request_id, @@ -151,6 +160,7 @@ impl Rtmp { break; } + info!(stream_key_id = key.id, label = %key.label, "stream started"); stream_sessions.insert( key.id, StreamSession { @@ -175,6 +185,7 @@ impl Rtmp { stream_key, .. } => { + info!(stream_key = %stream_key, "publish stream finished"); let key = entity::stream_key::Entity::find_by_key( &db, &stream_key, @@ -214,12 +225,10 @@ impl Rtmp { .. } => { // 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(); + 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(); } diff --git a/crates/server/src/webrtc.rs b/crates/server/src/webrtc.rs index c1fc81e..316de81 100644 --- a/crates/server/src/webrtc.rs +++ b/crates/server/src/webrtc.rs @@ -1,8 +1,10 @@ +use bytes::Bytes; use dashmap::DashMap; use entity::stream_key; use sea_orm::{DatabaseConnection, EntityTrait}; use std::{ error::Error, + net::SocketAddr, sync::Arc, time::{Duration, Instant}, }; @@ -19,12 +21,13 @@ use str0m::{ net::{Protocol, Receive}, }; -use crate::{StreamSession, audio::OpusAudioFrame, media::VideoFrame}; +use crate::{StreamSession, audio::OpusAudioFrame, media::VideoFrame, webrtc_proxy::WebrtcProxy}; pub struct Webrtc { pub offer_rx: Receiver<(i32, i32, String)>, pub accept_tx: async_broadcast::Sender<(i32, Option)>, pub sessions_ref: Arc>, + pub proxy: Arc, pub db: DatabaseConnection, } @@ -42,12 +45,20 @@ impl Webrtc { .await; if let Ok(None) = stream_key { + warn!(request_id, stream_id, "stream key not found in DB, rejecting offer"); + self.accept_tx.broadcast((request_id, None)).await.unwrap(); + continue; + } + if let Err(ref e) = stream_key { + warn!(request_id, stream_id, "DB error looking up stream key: {:?}", e); self.accept_tx.broadcast((request_id, None)).await.unwrap(); continue; } - let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap(); - let local_addr = socket.local_addr().unwrap(); + // let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + // let local_addr = socket.local_addr().unwrap(); + + let local_addr = self.proxy.public_addr(); let mut builder = Rtc::builder(); { @@ -78,6 +89,7 @@ impl Webrtc { Some("0".to_string()), None, ); + changes.add_channel("meow".into()); let offer_answer = match changes.accept_offer(offer_sdp) { Ok(a) => a, Err(e) => { @@ -92,6 +104,7 @@ impl Webrtc { .and_then(|l| l.strip_prefix("a=ice-ufrag:")) .map(|s| s.trim().to_string()); info!("Serving webrtc ufrag: {:?}", ufrag); + let (socket, rx) = self.proxy.add_client(ufrag.unwrap()); debug!(request_id, "sending answer back"); self.accept_tx @@ -100,8 +113,9 @@ impl Webrtc { .unwrap(); let sessions_ref = self.sessions_ref.clone(); + let public_addr = self.proxy.public_addr(); tokio::spawn(async move { - Webrtc::detach_connection(socket, rtc, sessions_ref, stream_id, mid).await; + Webrtc::detach_connection(socket, rx, rtc, sessions_ref, stream_id, mid, public_addr).await; }); } }); @@ -109,11 +123,13 @@ impl Webrtc { } async fn detach_connection( - socket: UdpSocket, + socket: Arc, + mut rx: Receiver<(Bytes, SocketAddr)>, mut rtc: Rtc, sessions_ref: Arc>, stream_id: i32, _hint_mid: Mid, + local_addr: SocketAddr, ) { let mut video_mid: Option = None; let mut video_pt = None; @@ -124,9 +140,6 @@ impl Webrtc { let mut audio_stream: Option>> = None; let mut saw_keyframe = false; - let mut recv_buf = vec![0u8; PER_CLIENT_CONNECTION_BUF]; - let local_addr = socket.local_addr().unwrap(); - loop { let deadline = loop { match rtc.poll_output() { @@ -251,11 +264,11 @@ impl Webrtc { } Err(async_broadcast::TryRecvError::Empty) => break, Err(async_broadcast::TryRecvError::Closed) => { - info!("video channel closed, stream ended"); + warn!("video channel closed, stream ended"); break; } Err(async_broadcast::TryRecvError::Overflowed(_)) => { - saw_keyframe = false; + // saw_keyframe = false; continue; } } @@ -317,10 +330,9 @@ impl Webrtc { 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() { + result = rx.recv() => { + if let Some((data, from)) = result { + if let Ok(contents) = (&data[..]).try_into() { if let Err(e) = rtc.handle_input(Input::Receive( Instant::now(), Receive { diff --git a/crates/server/src/webrtc_proxy.rs b/crates/server/src/webrtc_proxy.rs new file mode 100644 index 0000000..b067b44 --- /dev/null +++ b/crates/server/src/webrtc_proxy.rs @@ -0,0 +1,218 @@ +use std::{env, error::Error, net::SocketAddr, sync::Arc}; + +use bytes::Bytes; +use dashmap::DashMap; +use str0m::net::DatagramRecv; +use tokio::{ + net::UdpSocket, + sync::mpsc::{self, Receiver}, +}; +use tracing::{debug, error, info, warn}; + +pub struct WebRtcProxyConfig { + pub proxy_port: i32, +} + +pub struct WebrtcProxy { + clients_ufrag: Arc>>, + clients_addr: Arc>>, + socket: Arc, + public_addr: SocketAddr, +} + +const STUN_MAGIC: u32 = 0x2112A442; + +impl WebrtcProxy { + pub async fn new(config: WebRtcProxyConfig) -> Result> { + let sock = UdpSocket::bind(format!("0.0.0.0:{}", config.proxy_port)).await?; + let port = sock.local_addr()?.port(); + + let public_ip = match env::var("PUBLIC_DOMAIN") { + Ok(domain) => { + let ip = resolve_domain(&domain).await?; + info!(%domain, %ip, "resolved PUBLIC_DOMAIN for WebRTC candidates"); + ip + } + Err(_) => { + let ip = stun_public_ip().await?; + info!(%ip, "discovered public IP via STUN for WebRTC candidates"); + ip + } + }; + let public_addr = SocketAddr::new(public_ip, port); + info!(%public_addr, "WebRTC UDP proxy listening"); + + Ok(WebrtcProxy { + socket: Arc::new(sock), + clients_ufrag: Arc::new(DashMap::new()), + clients_addr: Arc::new(DashMap::new()), + public_addr, + }) + } + pub fn start(&self) -> Result<(), Box> { + // let self_arc = Arc::new(self); + let by_ufrag = self.clients_ufrag.clone(); + let by_addr = self.clients_addr.clone(); + let socket = self.socket.clone(); + + tokio::spawn(async move { + let mut buf = vec![0u8; 65535]; + loop { + let (b, from) = match socket.recv_from(&mut buf).await { + Ok(data) => data, + Err(err) => { + warn!("proxy couldnt eat data from socket, (({}))", err); + continue; + } + }; + let data = Bytes::copy_from_slice(&buf[..b]); + + // By addr + if let Some(tx) = by_addr.get(&from) { + match tx.try_send((data, from)) { + Ok(_) => continue, + Err(e) => { + match e { + mpsc::error::TrySendError::Full(_) => continue, + mpsc::error::TrySendError::Closed(_) => { + // the rv is ded + drop(tx); + by_addr.remove(&from); + continue; + } + }; + } + }; + }; + + let Some(ufrag) = self::WebrtcProxy::ufrag(&data) else { + debug!("huh, packet isnt stun or added as client."); + continue; + }; + + let Some((_, tx)) = by_ufrag.remove(&ufrag) else { + warn!("STUN packet ({}), isnt registored", ufrag); + continue; + }; + + by_addr.insert(from, tx.clone()); + debug!("got ufrag {}", ufrag); + debug!("sending data"); + if let Err(e) = tx.try_send((data, from)) { + match e { + mpsc::error::TrySendError::Full(_) => { + error!("Channel full") + } + mpsc::error::TrySendError::Closed(_) => { + error!("Channel is closed"); + } + } + }; + } + }); + Ok(()) + } + pub fn add_client(&self, ufrag: String) -> (Arc, Receiver<(Bytes, SocketAddr)>) { + debug!("Added client {}", ufrag); + let (tx, rx) = mpsc::channel(256); + self.clients_ufrag.insert(ufrag, tx); + (self.socket.clone(), rx) + } + pub fn local_addr(&self) -> SocketAddr { + self.socket.local_addr().unwrap() + } + + pub fn public_addr(&self) -> SocketAddr { + self.public_addr + } + pub fn ufrag(b: &Bytes) -> Option { + if b.len() <= 20 { + return None; + } + let magic = u32::from_be_bytes(b[4..8].try_into().ok()?); + if magic != STUN_MAGIC { + return None; + } + // attribies start at 20 + let mut pos = 20usize; + while (pos + 4) <= b.len() { + let attr_type: u16 = u16::from_be_bytes(b[pos..pos + 2].try_into().ok()?); + let attr_len: u16 = u16::from_be_bytes(b[pos + 2..pos + 4].try_into().ok()?); + pos = pos + 4; + if attr_type == 0x0006 { + let value = std::str::from_utf8(b[pos..pos + (attr_len as usize)].try_into().ok()?); + let local = value.unwrap().split(":").next(); + return Some(local.unwrap().to_string()); + } + pos += (attr_len as usize + 3) & !3; + } + None + } +} + +async fn resolve_domain(domain: &str) -> Result> { + let addr = tokio::net::lookup_host(format!("{}:0", domain)) + .await? + .find(|a| a.is_ipv4()) + .ok_or_else(|| format!("no IPv4 address found for {}", domain))?; + Ok(addr.ip()) +} + +// Send a STUN Binding Request to a public STUN server and extract our public IP +// from the XOR-MAPPED-ADDRESS attribute in the response. +async fn stun_public_ip() -> Result> { + let sock = UdpSocket::bind("0.0.0.0:0").await?; + sock.connect("stun.l.google.com:19302").await?; + + // Build a minimal STUN Binding Request (RFC 5389). + // Header: type(2) | length(2) | magic(4) | transaction-id(12) + let mut req = [0u8; 20]; + req[0..2].copy_from_slice(&0x0001u16.to_be_bytes()); // Binding Request + req[2..4].copy_from_slice(&0u16.to_be_bytes()); // no attributes + req[4..8].copy_from_slice(&STUN_MAGIC.to_be_bytes()); + req[8..20].copy_from_slice(b"rtmp2whip_tx"); // transaction ID (12 bytes) + + sock.send(&req).await?; + + let mut buf = [0u8; 512]; + let n = tokio::time::timeout(std::time::Duration::from_secs(5), sock.recv(&mut buf)).await??; + + let data = &buf[..n]; + parse_xor_mapped_address(data).ok_or("no XOR-MAPPED-ADDRESS in STUN response".into()) +} + +// Parse XOR-MAPPED-ADDRESS (0x0020) from a STUN response. +// The IP is XOR'd with the magic cookie (IPv4) or magic+transaction-id (IPv6). +fn parse_xor_mapped_address(data: &[u8]) -> Option { + if data.len() < 20 { + return None; + } + let magic = u32::from_be_bytes(data[4..8].try_into().ok()?); + if magic != STUN_MAGIC { + return None; + } + + let mut pos = 20usize; + while pos + 4 <= data.len() { + let attr_type = u16::from_be_bytes(data[pos..pos + 2].try_into().ok()?); + let attr_len = u16::from_be_bytes(data[pos + 2..pos + 4].try_into().ok()?) as usize; + pos += 4; + if pos + attr_len > data.len() { + break; + } + if attr_type == 0x0020 && attr_len >= 8 { + // byte 0: reserved, byte 1: family (0x01=IPv4, 0x02=IPv6) + let family = data[pos + 1]; + let x_port = u16::from_be_bytes(data[pos + 2..pos + 4].try_into().ok()?); + let _ = x_port ^ (STUN_MAGIC >> 16) as u16; // port (unused here) + + if family == 0x01 { + let x_addr = u32::from_be_bytes(data[pos + 4..pos + 8].try_into().ok()?); + let addr = x_addr ^ STUN_MAGIC; + return Some(std::net::IpAddr::V4(std::net::Ipv4Addr::from(addr))); + } + } + pos += (attr_len + 3) & !3; + } + None +} diff --git a/flake.nix b/flake.nix index 074a8fb..edbf5f6 100644 --- a/flake.nix +++ b/flake.nix @@ -31,13 +31,11 @@ src = craneLib.cleanCargoSource ./.; strictDeps = true; - buildInputs = [ - # Add additional build inputs here - ] - ++ pkgs.lib.optionals pkgs.stdenv.isDarwin [ - # Additional darwin specific inputs can be set here - pkgs.libiconv - ]; + buildInputs = + [ ] + ++ pkgs.lib.optionals pkgs.stdenv.isDarwin [ + pkgs.libiconv + ]; }; fileSetForCrate = @@ -64,20 +62,7 @@ src = fileSetForCrate ./crates/server; } ); - # migration = craneLib.buildPackage ( - # commonArgs - # // { - # pname = "migration"; - # version = "0.1.0"; - # cargoArtifacts = craneLib.buildDepsOnly commonArgs; - # cargoExtraArgs = "-p migration"; - # src = fileSetForCrate ./crates/migration; - # - # # Additional environment variables or build phases/hooks can be set - # # here *without* rebuilding all dependency crates - # # MY_CUSTOM_VAR = "some value"; - # } - # ); + in { checks = { @@ -91,15 +76,10 @@ }; devShells.default = craneLib.devShell { - # Inherit inputs from checks. checks = self.checks.${system}; ADMIN_REF_CODE = "meowmeowpurrrmeow"; - # Additional dev-shell environment variables can be set directly - # MY_CUSTOM_DEVELOPMENT_VAR = "something else"; - - # Extra inputs can be added here; cargo and rustc are provided by default. packages = [ pkgs.clang pkgs.mold @@ -107,6 +87,7 @@ pkgs.cmake pkgs.opus pkgs.pkgconf + pkgs.just ]; }; }