use std::collections::VecDeque;
use std::net::{SocketAddr, UdpSocket};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::mpsc::{Receiver, Sender, TryRecvError, channel};
use std::thread::JoinHandle;
use std::time::{Duration, Instant};
use parking_lot::Mutex;
use str0m::change::{SdpAnswer, SdpOffer, SdpPendingOffer};
use str0m::channel::ChannelId;
use str0m::net::{Protocol, Receive};
use str0m::{Candidate, Event, Input, Output, Rtc};
use crate::webrtc_transport::DataChannel;
#[derive(Debug)]
pub enum Str0mNetError {
Io(std::io::Error),
Rtc(str0m::RtcError),
Sdp(String),
Closed,
Backpressure,
NotOfferer,
}
impl std::fmt::Display for Str0mNetError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Io(e) => write!(f, "str0m net io error: {e}"),
Self::Rtc(e) => write!(f, "str0m net rtc error: {e}"),
Self::Sdp(e) => write!(f, "str0m net sdp/ice error: {e}"),
Self::Closed => write!(f, "str0m net driver closed"),
Self::Backpressure => {
write!(
f,
"str0m net outbound queue full (apply flow control and retry)"
)
}
Self::NotOfferer => write!(f, "accept_answer called on a non-offerer peer"),
}
}
}
impl std::error::Error for Str0mNetError {}
impl From<std::io::Error> for Str0mNetError {
fn from(e: std::io::Error) -> Self {
Self::Io(e)
}
}
impl From<str0m::RtcError> for Str0mNetError {
fn from(e: str0m::RtcError) -> Self {
Self::Rtc(e)
}
}
enum DriverCmd {
AcceptAnswer(String),
AddRemoteCandidate(String),
Send(Vec<u8>),
Shutdown,
}
const MAX_PENDING_FRAMES: usize = 1024;
const MAX_INBOX_FRAMES: usize = 1024;
struct Shared {
inbox: Mutex<VecDeque<Vec<u8>>>,
open: AtomicBool,
closed: AtomicBool,
pending_frames: AtomicUsize,
dropped_inbox_frames: AtomicUsize,
last_error: Mutex<Option<String>>,
}
pub struct Str0mNet {
cmd_tx: Sender<DriverCmd>,
shared: Arc<Shared>,
local_candidate: String,
driver: Option<JoinHandle<()>>,
is_offerer: bool,
}
impl Str0mNet {
pub fn offer(bind: SocketAddr) -> Result<(Self, String), Str0mNetError> {
let socket = UdpSocket::bind(bind)?;
let local_addr = socket.local_addr()?;
let now = Instant::now();
let mut rtc = Rtc::new(now);
let mut api = rtc.sdp_api();
let cid = api.add_channel("lazily-ipc".to_string());
let (offer, pending) = api
.apply()
.ok_or_else(|| Str0mNetError::Sdp("str0m produced no offer".into()))?;
let local = host_candidate(local_addr)?;
rtc.add_local_candidate(local.clone());
let net = Self::spawn(
rtc,
Some(pending),
Some(cid),
socket,
local_addr,
local,
true,
);
Ok((net, offer.to_sdp_string()))
}
pub fn answer(bind: SocketAddr, offer_sdp: &str) -> Result<(Self, String), Str0mNetError> {
let socket = UdpSocket::bind(bind)?;
let local_addr = socket.local_addr()?;
let now = Instant::now();
let mut rtc = Rtc::new(now);
let offer =
SdpOffer::from_sdp_string(offer_sdp).map_err(|e| Str0mNetError::Sdp(e.to_string()))?;
let answer = rtc
.sdp_api()
.accept_offer(offer)
.map_err(|e| Str0mNetError::Sdp(e.to_string()))?;
let local = host_candidate(local_addr)?;
rtc.add_local_candidate(local.clone());
let net = Self::spawn(rtc, None, None, socket, local_addr, local, false);
Ok((net, answer.to_sdp_string()))
}
fn spawn(
rtc: Rtc,
pending: Option<SdpPendingOffer>,
cid: Option<ChannelId>,
socket: UdpSocket,
local_addr: SocketAddr,
local_candidate: Candidate,
is_offerer: bool,
) -> Self {
let shared = Arc::new(Shared {
inbox: Mutex::new(VecDeque::new()),
open: AtomicBool::new(false),
closed: AtomicBool::new(false),
pending_frames: AtomicUsize::new(0),
dropped_inbox_frames: AtomicUsize::new(0),
last_error: Mutex::new(None),
});
let (cmd_tx, cmd_rx) = channel();
let candidate_string = local_candidate.to_sdp_string();
let driver_shared = shared.clone();
let driver = std::thread::Builder::new()
.name("str0m-net-driver".to_string())
.spawn(move || {
run_driver(rtc, pending, cid, socket, local_addr, cmd_rx, driver_shared);
})
.expect("spawn str0m-net driver thread");
Self {
cmd_tx,
shared,
local_candidate: candidate_string,
driver: Some(driver),
is_offerer,
}
}
pub fn local_candidate(&self) -> &str {
&self.local_candidate
}
pub fn accept_answer(&self, answer_sdp: &str) -> Result<(), Str0mNetError> {
if !self.is_offerer {
return Err(Str0mNetError::NotOfferer);
}
SdpAnswer::from_sdp_string(answer_sdp).map_err(|e| Str0mNetError::Sdp(e.to_string()))?;
self.cmd_tx
.send(DriverCmd::AcceptAnswer(answer_sdp.to_string()))
.map_err(|_| Str0mNetError::Closed)
}
pub fn add_remote_candidate(&self, candidate_sdp: &str) -> Result<(), Str0mNetError> {
Candidate::from_sdp_string(candidate_sdp).map_err(|e| Str0mNetError::Sdp(e.to_string()))?;
self.cmd_tx
.send(DriverCmd::AddRemoteCandidate(candidate_sdp.to_string()))
.map_err(|_| Str0mNetError::Closed)
}
pub fn is_open(&self) -> bool {
self.shared.open.load(Ordering::SeqCst) && !self.shared.closed.load(Ordering::SeqCst)
}
pub fn dropped_inbox_frames(&self) -> usize {
self.shared.dropped_inbox_frames.load(Ordering::Relaxed)
}
pub fn last_error(&self) -> Option<String> {
self.shared.last_error.lock().clone()
}
pub fn wait_open(&self, timeout: Duration) -> bool {
let deadline = Instant::now() + timeout;
loop {
if self.is_open() {
return true;
}
if self.shared.closed.load(Ordering::SeqCst) {
return false;
}
if Instant::now() >= deadline {
return self.is_open();
}
std::thread::sleep(Duration::from_millis(5));
}
}
pub fn channel(&self) -> Str0mNetChannel {
Str0mNetChannel {
cmd_tx: self.cmd_tx.clone(),
shared: self.shared.clone(),
}
}
}
impl Drop for Str0mNet {
fn drop(&mut self) {
let _ = self.cmd_tx.send(DriverCmd::Shutdown);
if let Some(driver) = self.driver.take() {
let _ = driver.join();
}
}
}
#[derive(Clone)]
pub struct Str0mNetChannel {
cmd_tx: Sender<DriverCmd>,
shared: Arc<Shared>,
}
impl DataChannel for Str0mNetChannel {
type Error = Str0mNetError;
fn send_frame(&self, frame: Vec<u8>) -> Result<(), Self::Error> {
if self.shared.closed.load(Ordering::SeqCst) {
return Err(Str0mNetError::Closed);
}
if self.shared.pending_frames.load(Ordering::Relaxed) >= MAX_PENDING_FRAMES {
return Err(Str0mNetError::Backpressure);
}
self.cmd_tx
.send(DriverCmd::Send(frame))
.map_err(|_| Str0mNetError::Closed)
}
fn try_recv_frame(&self) -> Result<Option<Vec<u8>>, Self::Error> {
if let Some(frame) = self.shared.inbox.lock().pop_front() {
return Ok(Some(frame));
}
if self.shared.closed.load(Ordering::SeqCst) && !self.is_open() {
return Err(Str0mNetError::Closed);
}
Ok(None)
}
fn is_open(&self) -> bool {
self.shared.open.load(Ordering::SeqCst) && !self.shared.closed.load(Ordering::SeqCst)
}
}
fn host_candidate(addr: SocketAddr) -> Result<Candidate, Str0mNetError> {
Candidate::host(addr, "udp").map_err(|e| Str0mNetError::Sdp(e.to_string()))
}
fn run_driver(
mut rtc: Rtc,
mut pending: Option<SdpPendingOffer>,
mut cid: Option<ChannelId>,
socket: UdpSocket,
local_addr: SocketAddr,
cmd_rx: Receiver<DriverCmd>,
shared: Arc<Shared>,
) {
let mut buf = [0u8; 2048];
let mut out_pending: VecDeque<Vec<u8>> = VecDeque::new();
'outer: loop {
loop {
match cmd_rx.try_recv() {
Ok(DriverCmd::Shutdown) => break 'outer,
Ok(DriverCmd::AcceptAnswer(sdp)) => {
if let (Some(p), Ok(answer)) =
(pending.take(), SdpAnswer::from_sdp_string(&sdp))
&& let Err(e) = rtc.sdp_api().accept_answer(p, answer)
{
*shared.last_error.lock() =
Some(format!("accept_answer apply failed: {e}"));
}
}
Ok(DriverCmd::AddRemoteCandidate(s)) => {
if let Ok(c) = Candidate::from_sdp_string(&s) {
rtc.add_remote_candidate(c);
}
}
Ok(DriverCmd::Send(bytes)) => {
out_pending.push_back(bytes);
shared.pending_frames.fetch_add(1, Ordering::Relaxed);
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => break 'outer,
}
}
if shared.open.load(Ordering::SeqCst)
&& let Some(id) = cid
{
while let Some(frame) = out_pending.pop_front() {
match rtc.channel(id) {
Some(mut ch) => match ch.write(true, &frame) {
Ok(true) => {
shared.pending_frames.fetch_sub(1, Ordering::Relaxed);
}
Ok(false) | Err(_) => {
out_pending.push_front(frame);
break;
}
},
None => {
out_pending.push_front(frame);
break;
}
}
}
}
if !rtc.is_alive() {
break;
}
let timeout = loop {
match rtc.poll_output() {
Ok(Output::Timeout(t)) => break t,
Ok(Output::Transmit(t)) => {
match socket.send_to(&t.contents, t.destination) {
Ok(_) => {}
Err(e)
if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::Interrupted =>
{
continue;
}
Err(_) => break 'outer,
}
}
Ok(Output::Event(e)) => handle_event(e, &mut cid, &shared),
Err(_) => break 'outer,
}
};
const COMMAND_POLL_INTERVAL: Duration = Duration::from_millis(15);
let now = Instant::now();
let wait = timeout
.checked_duration_since(now)
.unwrap_or(Duration::ZERO)
.min(COMMAND_POLL_INTERVAL)
.max(Duration::from_millis(1));
if socket.set_read_timeout(Some(wait)).is_err() {
break;
}
match socket.recv_from(&mut buf) {
Ok((n, src)) => {
if let Ok(contents) = buf[..n].try_into() {
let _ = rtc.handle_input(Input::Receive(
Instant::now(),
Receive {
proto: Protocol::Udp,
source: src,
destination: local_addr,
contents,
},
));
}
}
Err(e)
if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::TimedOut =>
{
let _ = rtc.handle_input(Input::Timeout(Instant::now()));
}
Err(_) => break,
}
}
shared.open.store(false, Ordering::SeqCst);
shared.closed.store(true, Ordering::SeqCst);
shared.pending_frames.store(0, Ordering::SeqCst);
}
fn push_inbox_frame(shared: &Shared, frame: Vec<u8>) {
let mut inbox = shared.inbox.lock();
if inbox.len() >= MAX_INBOX_FRAMES {
inbox.pop_front();
shared.dropped_inbox_frames.fetch_add(1, Ordering::Relaxed);
}
inbox.push_back(frame);
}
fn handle_event(event: Event, cid: &mut Option<ChannelId>, shared: &Shared) {
match event {
Event::ChannelOpen(id, _label) => {
*cid = Some(id);
shared.open.store(true, Ordering::SeqCst);
}
Event::ChannelData(data) => {
push_inbox_frame(shared, data.data);
}
Event::ChannelClose(_) => {
shared.open.store(false, Ordering::SeqCst);
shared.closed.store(true, Ordering::SeqCst);
}
_ => {}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_shared() -> Shared {
Shared {
inbox: Mutex::new(VecDeque::new()),
open: AtomicBool::new(false),
closed: AtomicBool::new(false),
pending_frames: AtomicUsize::new(0),
dropped_inbox_frames: AtomicUsize::new(0),
last_error: Mutex::new(None),
}
}
#[test]
fn last_error_round_trips_apply_failure() {
let shared = make_shared();
assert!(
shared.last_error.lock().is_none(),
"last_error starts clear on a fresh peer"
);
*shared.last_error.lock() = Some("fingerprint mismatch".to_string());
assert_eq!(
shared.last_error.lock().clone(),
Some("fingerprint mismatch".to_string()),
"apply-time failure must be stored and readable for wait_open diagnosis"
);
}
#[test]
fn inbox_caps_at_max_inbox_frames_and_counts_drops() {
let shared = make_shared();
for _ in 0..(MAX_INBOX_FRAMES + 50) {
push_inbox_frame(&shared, vec![0u8]);
}
{
let inbox = shared.inbox.lock();
assert_eq!(
inbox.len(),
MAX_INBOX_FRAMES,
"inbox must be capped at MAX_INBOX_FRAMES"
);
assert_eq!(
shared.dropped_inbox_frames.load(Ordering::Relaxed),
50,
"overflowing frames must be counted as dropped"
);
}
push_inbox_frame(&shared, vec![42u8]);
{
let inbox = shared.inbox.lock();
assert_eq!(
inbox.len(),
MAX_INBOX_FRAMES,
"capped inbox must not grow further"
);
assert_eq!(
inbox.back().map(|v| v[0]),
Some(42u8),
"drop-oldest ring semantics must retain the newest frame"
);
assert_eq!(
shared.dropped_inbox_frames.load(Ordering::Relaxed),
51,
"one more drop to make room for the new frame"
);
}
}
}