#[cfg(feature = "audio")]
use crate::audio::transport::EncodedFrame as EncodedAudio;
use crate::capture::CaptureSource;
use crate::encoder::SurfaceSnapshot;
use crate::network::{
verify_token, Epoch, Message, MessageSink, MessageSource, MessageTransport, NetworkConfig, Rev,
SessionToken, SNAPSHOT_CHUNK_BYTES,
};
use crate::pcc::types::{rgb_len, Frame};
use crate::pcc::{PlanLimits, Planner};
use crate::reach::Rung;
use crate::relay::{
generate_session_code, RelayFan, RelayRole, RelayTransport, K_CONTROL_BROWSER_HERE,
RELAY_CONTROL_ID,
};
use crate::server::renderer::{web, SharedSurface};
use anyhow::{Context, Result};
use parking_lot::Mutex;
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{broadcast, RwLock};
use tracing::{error, info, warn};
const REPAIR_INTERVAL: Duration = Duration::from_secs(30);
const BROADCAST_DEPTH: usize = 256;
const MOTION_PREVIEW_FRACTION: f64 = 0.40;
const SEND_BUDGET: usize = crate::network::MAX_MESSAGE_SIZE as usize;
const AUTH_FAILURES_ALLOWED: u32 = 8;
const AUTH_FAILURE_WINDOW: Duration = Duration::from_secs(60);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
const MAX_CONSECUTIVE_LAGS: u32 = 3;
#[derive(Clone)]
pub struct ShareArgs {
pub listen: Option<String>,
pub relay: Option<String>,
pub relay_pin: String,
pub session: Option<String>,
pub web: Option<String>,
pub web_cert: Option<(String, String)>,
pub synthetic: bool,
pub capture_target: crate::capture::CaptureTarget,
pub fps: u32,
pub max_fps: u32,
pub quality: f32,
pub token: SessionToken,
pub repair_interval: Duration,
pub reach: crate::reach::ReachPolicy,
pub approve: bool,
pub audio_source: crate::audio::AudioSource,
pub transport: crate::network::TransportKind,
pub broadcast_above: usize,
}
#[derive(Clone)]
pub struct Published {
pub rev: Rev,
pub epoch: Epoch,
pub snapshot: Arc<SurfaceSnapshot>,
pub encoded: Option<(Rev, Arc<Vec<u8>>)>,
pub ring: RevisionRing,
}
impl Published {
fn encoded_for(&self, rev: Rev) -> Option<&Arc<Vec<u8>>> {
self.encoded
.as_ref()
.filter(|(r, _)| *r == rev)
.map(|(_, d)| d)
}
}
pub type Shared = Arc<RwLock<Published>>;
type EncodedBroadcast = broadcast::Sender<Arc<Vec<u8>>>;
#[derive(Debug, Default, Clone)]
pub struct RevisionRing {
entries: std::collections::VecDeque<RingEntry>,
bytes: usize,
}
#[derive(Debug, Clone)]
struct RingEntry {
rev: Rev,
epoch: Epoch,
bytes: Arc<Vec<u8>>,
}
const RING_MAX_ENTRIES: usize = 512;
const RING_MAX_BYTES: usize = 16 * 1024 * 1024;
impl RevisionRing {
pub fn push(&mut self, rev: Rev, epoch: Epoch, bytes: Arc<Vec<u8>>) {
if let Some(back) = self.entries.back() {
if back.epoch != epoch {
self.entries.clear();
self.bytes = 0;
}
}
self.bytes += bytes.len();
self.entries.push_back(RingEntry { rev, epoch, bytes });
while self.entries.len() > RING_MAX_ENTRIES || self.bytes > RING_MAX_BYTES {
if let Some(old) = self.entries.pop_front() {
self.bytes = self.bytes.saturating_sub(old.bytes.len());
} else {
break;
}
}
}
pub fn push_keepalive(&mut self, rev: Rev, epoch: Epoch, bytes: Arc<Vec<u8>>) {
if let Some(back) = self.entries.back() {
if back.rev == rev && back.epoch == epoch {
return;
}
}
self.push(rev, epoch, bytes);
}
pub fn replay_from(&self, floor: Rev, epoch: Epoch) -> Option<Vec<Arc<Vec<u8>>>> {
if self.entries.is_empty() {
return None;
}
let mut start = None;
let mut prev_rev: Option<Rev> = None;
for (i, entry) in self.entries.iter().enumerate() {
if entry.epoch != epoch {
return None;
}
if entry.rev <= floor {
if entry.rev == floor {
start = Some(i + 1);
} else {
start = None;
}
prev_rev = Some(entry.rev);
continue;
}
match prev_rev {
Some(prev) if entry.rev == prev || entry.rev == prev + 1 => {}
_ => return None,
}
start?;
prev_rev = Some(entry.rev);
}
let start = start?;
if start >= self.entries.len() {
return Some(Vec::new());
}
Some(
self.entries
.range(start..)
.map(|e| e.bytes.clone())
.collect(),
)
}
}
#[cfg(test)]
mod revision_ring_tests {
use super::*;
fn msg(rev: Rev, _epoch: Epoch) -> Arc<Vec<u8>> {
Arc::new(Message::KeepAlive { rev }.encode().unwrap())
}
#[test]
fn replay_returns_everything_after_the_floor_in_order() {
let mut ring = RevisionRing::default();
for rev in 1..=5 {
ring.push(rev, 0, msg(rev, 0));
}
let out = ring.replay_from(2, 0).expect("contiguous coverage replays");
assert_eq!(out.len(), 3);
assert_eq!(crate::network::peek_rev(&out[0]), Some(3));
assert_eq!(crate::network::peek_rev(&out[2]), Some(5));
}
#[test]
fn a_floor_that_predates_the_ring_falls_back_to_snapshot() {
let mut ring = RevisionRing::default();
for rev in 10..=12 {
ring.push(rev, 0, msg(rev, 0));
}
assert!(
ring.replay_from(3, 0).is_none(),
"unknown history must not replay"
);
}
#[test]
fn a_hole_in_the_middle_falls_back_to_snapshot() {
let mut ring = RevisionRing::default();
for rev in [1u64, 2, 4, 5] {
ring.push(rev, 0, msg(rev, 0));
}
assert!(
ring.replay_from(1, 0).is_none(),
"a missing revision is not replayable"
);
}
#[test]
fn an_epoch_change_clears_the_ring() {
let mut ring = RevisionRing::default();
for rev in 1..=3 {
ring.push(rev, 0, msg(rev, 0));
}
ring.push(4, 1, msg(4, 1));
assert!(
ring.replay_from(2, 0).is_none(),
"old-epoch rectangles would corrupt"
);
let out = ring.replay_from(4, 1).expect("caught-up floor replays");
assert!(out.is_empty());
}
#[test]
fn the_bounds_evict_oldest_first() {
let mut ring = RevisionRing::default();
for rev in 1..=(RING_MAX_ENTRIES as u64 + 10) {
ring.push(rev, 0, msg(rev, 0));
}
assert_eq!(ring.entries.len(), RING_MAX_ENTRIES);
assert!(
ring.replay_from(1, 0).is_none(),
"evicted history must not replay"
);
let floor = RING_MAX_ENTRIES as u64 + 9;
assert!(
ring.replay_from(floor, 0).is_some(),
"retained tail must replay"
);
}
#[test]
fn idle_keepalives_do_not_churn_real_history() {
let mut ring = RevisionRing::default();
for rev in 1..=3 {
ring.push(rev, 0, msg(rev, 0));
}
for _ in 0..(RING_MAX_ENTRIES + 10) {
ring.push_keepalive(3, 0, msg(3, 0));
}
assert_eq!(ring.entries.len(), 3);
let out = ring.replay_from(2, 0).expect("history must survive idling");
assert_eq!(out.len(), 1);
assert_eq!(crate::network::peek_rev(&out[0]), Some(3));
}
#[test]
fn an_idle_floor_with_nothing_after_replays_empty() {
let mut ring = RevisionRing::default();
for rev in 1..=3 {
ring.push(rev, 0, msg(rev, 0));
}
let out = ring.replay_from(3, 0).expect("caught up replays empty");
assert!(out.is_empty());
}
}
#[cfg(feature = "audio")]
type AudioBroadcast = broadcast::Sender<Arc<EncodedAudio>>;
#[cfg(feature = "audio")]
const AUDIO_QUEUE: usize = 256;
#[derive(Debug, Default, Clone, Copy)]
struct ViewerStats {
lag_events: u32,
acked_rev: Rev,
rtt_ms: Option<u64>,
lost_packets: u64,
cwnd_bytes: u64,
direct: bool,
redirect_queued: bool,
}
#[derive(Debug)]
struct Pressure {
value: f32,
last_change: Instant,
}
impl Pressure {
fn penalise(&mut self, weight: f32) {
self.value = (self.value + weight).min(4.0);
}
fn decay(&mut self, dt: Duration) {
self.value = (self.value - dt.as_secs_f32() * 0.25).max(0.0);
}
fn is_pressured(&self) -> bool {
self.value >= 1.0
}
}
pub async fn run_share(args: ShareArgs, metrics: crate::telemetry::SharedMetrics) -> Result<()> {
let capture = CaptureSource::open_target(args.synthetic, &args.capture_target)?;
let (width, height) = (capture.width(), capture.height());
let token = args.token.clone();
info!("Sharing {width}x{height}");
let identity = Arc::new(
crate::network::generate_identity().context("Failed to create the sharer identity")?,
);
println!("Certificate fingerprint (sha256): {}", identity.fingerprint);
let target_quality = args.quality.clamp(0.1, 1.0);
let requested_fps = args.fps.max(1).min(args.max_fps.max(1));
let surface = args
.web
.as_ref()
.map(|_| Arc::new(SharedSurface::new(width, height, target_quality)));
let (tx, keepalive) = broadcast::channel::<Arc<Vec<u8>>>(BROADCAST_DEPTH);
#[cfg(feature = "audio")]
let (audio_tx, audio_rx) = tokio::sync::broadcast::channel::<Arc<EncodedAudio>>(AUDIO_QUEUE);
#[cfg(feature = "audio")]
if args.audio_source != crate::audio::AudioSource::None_ {
start_audio_capture(audio_tx.clone(), args.audio_source);
}
let first = capture.capture_frame()?;
let first_frame = Frame::new(0, first.width, first.height, first.data)?;
let published: Shared = Arc::new(RwLock::new(Published {
rev: 0,
epoch: 0,
snapshot: Arc::new(SurfaceSnapshot::new(
first_frame.width,
first_frame.height,
first_frame.data.clone(),
)?),
encoded: None,
ring: RevisionRing::default(),
}));
if let Some(surface) = &surface {
surface.publish(&first_frame, target_quality);
}
if args.reach != crate::reach::ReachPolicy::RelayOnly {
let reach = args.reach;
let relay_configured = args.relay.is_some();
tokio::spawn(async move {
let Some(rung) = resolve_ladder(reach).await else {
return;
};
match (rung.is_direct(), relay_configured) {
(true, _) => info!(
"Direct path available via {}; the relay stays a fallback",
rung.label()
),
(false, true) => info!("No direct path found; using the relay"),
(false, false) => warn!(
"No direct path found and no --relay was given, so viewers must be on \
this network or reach the port directly"
),
}
});
}
if args.reach == crate::reach::ReachPolicy::DirectOnly && args.relay.is_some() {
anyhow::bail!("--reach direct was given but --relay was also set; pick one");
}
if let (Some(web_addr), Some(surface)) = (&args.web, &surface) {
let web_updates = tx.subscribe();
let source = published.clone();
let snapshot: web::SnapshotFn = Arc::new(move || {
let source = source.clone();
match try_current_snapshot(&source) {
Ok(msgs) => msgs,
Err(e) => {
warn!("Could not build a web snapshot: {e}");
Vec::new()
}
}
});
start_web(
web_addr,
args.web_cert.clone(),
surface,
web_updates,
&token,
&identity,
snapshot,
)
.await?;
}
let viewers: Arc<Mutex<HashMap<u64, ViewerStats>>> = Arc::new(Mutex::new(HashMap::new()));
let next_viewer_id = Arc::new(AtomicU64::new(1));
let redirect_to = redirect_target(&args);
if args.transport == crate::network::TransportKind::WebRtc {
let offer = crate::network::webrtc_host_offer().await?;
info!("WebRTC viewers with this offer (paste it to the viewer):");
info!(" OFFER: {}", offer.blob);
let answer = tokio::task::spawn_blocking(|| {
use std::io::{BufRead, Write};
let _ = writeln!(
std::io::stdout(),
"Paste the viewer answer blob, then Enter:"
);
let _ = std::io::stdout().flush();
let mut line = String::new();
std::io::BufReader::new(std::io::stdin().lock())
.read_line(&mut line)
.ok()
.filter(|_| !line.trim().is_empty())
.map(|_| line.trim().to_string())
})
.await
.context("reading the answer blob")?
.context("no answer blob pasted (stdin closed or empty)")?;
let _ = offer.answer_tx.send(answer);
let mut cand_rx = offer.candidate_rx;
let cand_in2 = offer.candidate_tx;
tokio::spawn(async move {
while let Some(c) = cand_rx.recv().await {
if let Ok(json) = serde_json::to_string(&c) {
println!("CANDIDATE: {json}");
}
}
});
tokio::task::spawn_blocking(move || {
use std::io::BufRead;
for line in std::io::BufReader::new(std::io::stdin().lock()).lines() {
let line = line.unwrap_or_default();
if line.trim().is_empty() {
break;
}
let body = line
.trim()
.strip_prefix("CANDIDATE:")
.map(str::trim)
.unwrap_or(line.trim());
if let Ok(c) = serde_json::from_str(body) {
let _ = cand_in2.try_send(c);
}
}
});
match tokio::time::timeout(std::time::Duration::from_secs(20), offer.transport_rx).await {
Ok(Ok(transport)) => {
info!("WebRTC viewer connected; serving");
spawn_webrtc_session(
transport,
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_viewer_id.clone(),
#[cfg(feature = "audio")]
Some((audio_tx.clone(), audio_rx.resubscribe())),
#[cfg(not(feature = "audio"))]
None,
metrics.clone(),
args.approve,
);
}
other => {
warn!("WebRTC viewer never completed the handshake");
let _ = other;
}
}
}
if args.transport == crate::network::TransportKind::Iroh {
let (endpoint, ticket) = crate::network::iroh_host_endpoint().await?;
info!("Iroh viewers with ticket:");
info!(
" pcc view --transport iroh --ticket {ticket} --token {}",
token.as_str()
);
spawn_iroh_accept_loop(
endpoint,
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_viewer_id.clone(),
#[cfg(feature = "audio")]
Some((audio_tx.clone(), audio_rx.resubscribe())),
#[cfg(not(feature = "audio"))]
None,
metrics.clone(),
args.approve,
);
}
if args.transport == crate::network::TransportKind::Quic {
if let Some(listen_addr) = &args.listen {
let addr = crate::network::resolve(listen_addr)
.await
.with_context(|| format!("Invalid --listen address '{listen_addr}'"))?;
let endpoint =
crate::network::server_endpoint(&NetworkConfig::default(), &identity, addr)?;
let dial_addr = crate::reach::lan_address_for(addr.port()).unwrap_or(addr);
info!(
"Direct viewers on {dial_addr} (pin the sharer: {})",
identity.fingerprint
);
info!(
" pcc view --connect {dial_addr} --token {} --pin {}",
token.as_str(),
identity.fingerprint
);
spawn_accept_loop(
endpoint,
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_viewer_id.clone(),
#[cfg(feature = "audio")]
Some((audio_tx.clone(), audio_rx.resubscribe())),
#[cfg(not(feature = "audio"))]
None,
metrics.clone(),
args.approve,
redirect_to.clone(),
);
}
}
if args.transport == crate::network::TransportKind::Quic {
if let Some(relays) = &args.relay {
let pin = crate::network::hex_to_der(&args.relay_pin)?;
let session = args.session.clone().unwrap_or_else(generate_session_code);
for relay_addr in split_relays(relays) {
let addr = match crate::network::resolve(&relay_addr).await {
Ok(a) => a,
Err(e) => {
warn!("Skipping relay '{relay_addr}': {e:#}");
continue;
}
};
info!("Relay {relay_addr}, session '{session}' (pin the relay, not the sharer)");
println!(
"Viewer link: {}",
crate::reach::pair_url_relay(
&relay_addr,
&args.relay_pin,
&session,
token.as_str()
)
);
println!(
"Browser link: https://{relay_addr}/v/{session}/#token={} \
(needs `pcc relay --web`)",
token.as_str()
);
info!(
" pcc view --relay {relay_addr} --pin {} --session '{session}' --token {}",
args.relay_pin,
token.as_str()
);
spawn_relay_loop(
addr,
pin.clone(),
session.clone(),
token.clone(),
tx.clone(),
published.clone(),
viewers.clone(),
next_viewer_id.clone(),
#[cfg(feature = "audio")]
Some((audio_tx.clone(), audio_rx.resubscribe())),
#[cfg(not(feature = "audio"))]
None,
metrics.clone(),
args.approve,
None,
);
}
}
}
drop(keepalive);
let origin = capture.origin();
let cursor_sampler: Box<dyn crate::capture::CursorSampler> = if args.synthetic {
Box::new(crate::capture::SweepSampler::new())
} else {
Box::new(crate::capture::PlatformCursorSampler::new(origin))
};
capture_loop(
metrics,
first_frame,
capture,
cursor_sampler,
Planner::default(),
published,
tx,
viewers,
surface,
target_quality,
requested_fps,
args.max_fps.max(1),
if args.repair_interval.is_zero() {
REPAIR_INTERVAL
} else {
args.repair_interval
},
args.broadcast_above,
redirect_to,
)
.await
}
#[cfg(feature = "audio")]
fn start_audio_capture(audio_tx: AudioBroadcast, which: crate::audio::AudioSource) {
let source = match crate::audio::capture::open_source(which) {
Ok(s) => s,
Err(e) => {
warn!("No audio capture device: {e}");
return;
}
};
let name = source.device_name().to_string();
let mut source = source;
match source.start() {
Ok(rx) => {
let start = Instant::now();
std::thread::spawn(move || {
let Ok(mut encoder) = crate::audio::AudioSender::new() else {
return;
};
while let Ok(frame) = rx.recv() {
let Ok(opus) = encoder.encode_frame(&frame.pcm) else {
continue;
};
let _ = audio_tx.send(Arc::new(EncodedAudio {
pts_us: frame.capture_time.duration_since(start).as_micros() as u64,
pcm_len: frame.pcm.len() as u32,
opus: Arc::new(opus),
}));
}
});
info!("Audio capture started from {name}");
}
Err(e) => warn!("Audio capture could not start: {e}"),
}
}
async fn start_web(
web_addr: &str,
cert: Option<(String, String)>,
surface: &Arc<SharedSurface>,
updates: broadcast::Receiver<Arc<Vec<u8>>>,
token: &SessionToken,
identity: &crate::network::ServerIdentity,
snapshot: web::SnapshotFn,
) -> Result<()> {
let addr = crate::network::resolve(web_addr)
.await
.with_context(|| format!("Invalid --web address '{web_addr}'"))?;
if cert.is_none() && !addr.ip().is_loopback() {
anyhow::bail!(
"refusing to serve the browser viewer on non-loopback {addr} without TLS: pass --web-cert/--web-key, or bind a loopback --web address"
);
}
let tls = match cert {
Some((cert_path, key_path)) => {
let certs = std::fs::read(&cert_path)
.with_context(|| format!("Failed to read the certificate at {cert_path}"))?;
let key = std::fs::read(&key_path)
.with_context(|| format!("Failed to read the private key at {key_path}"))?;
let parsed = rustls::ServerConfig::builder_with_provider(Arc::new(
rustls::crypto::ring::default_provider(),
))
.with_protocol_versions(rustls::ALL_VERSIONS)
.expect("the ring provider supports these versions")
.with_no_client_auth()
.with_single_cert(
vec![rustls::pki_types::CertificateDer::from(certs)],
rustls::pki_types::PrivatePkcs8KeyDer::from(key).into(),
)
.map_err(|e| anyhow::anyhow!("Invalid web certificate/key pair: {e}"))?;
Some(Arc::new(parsed))
}
None => None,
};
let surface = surface.clone();
let token = token.clone();
let fingerprint = identity.fingerprint.clone();
let snapshot_for_thread = snapshot.clone();
let scheme = if tls.is_some() { "https" } else { "http" };
let listener = std::net::TcpListener::bind(addr)
.with_context(|| format!("Failed to bind the web viewer on {addr}"))?;
listener
.set_nonblocking(true)
.context("Failed to put the web listener into non-blocking mode")?;
println!(
"Browser viewer: {scheme}://{addr}/#token={}",
token.as_str()
);
std::thread::spawn(move || {
if let Err(e) = web::run_web_server(
listener,
surface,
updates,
token,
fingerprint,
tls,
snapshot_for_thread,
) {
error!("Web viewer stopped: {e}");
}
});
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn spawn_webrtc_session(
transport: crate::network::WebrtcTransport,
tx: EncodedBroadcast,
published: Shared,
token: SessionToken,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_id: Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))] audio: Option<()>,
metrics: crate::telemetry::SharedMetrics,
approve: bool,
) {
tokio::spawn(async move {
serve_viewer(
Box::new(transport),
SocketAddr::from(([127, 0, 0, 1], 0)),
"webrtc",
tx,
published,
token,
viewers,
next_id,
audio,
None,
metrics,
approve,
None,
)
.await;
});
}
#[allow(clippy::too_many_arguments)]
fn spawn_iroh_accept_loop(
endpoint: iroh::Endpoint,
tx: EncodedBroadcast,
published: Shared,
token: SessionToken,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_id: Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))] audio: Option<()>,
metrics: crate::telemetry::SharedMetrics,
approve: bool,
) {
tokio::spawn(async move {
loop {
let Some(incoming) = endpoint.accept().await else {
break;
};
let Ok(accepting) = incoming.accept() else {
continue;
};
let conn = match accepting.await {
Ok(c) => c,
Err(e) => {
tracing::debug!("iroh handshake failed: {e:#}");
continue;
}
};
if conn.alpn() != crate::network::IROH_ALPN {
continue;
}
let peer = conn.remote_id().to_string();
tracing::debug!("iroh connection from {peer} (alpn ok)");
let Ok((send, recv)) = conn.accept_bi().await else {
warn!("Iroh viewer {peer} never opened a stream");
continue;
};
tracing::debug!("iroh stream from {peer} open; serving");
let (tx, published, token, viewers, next_id, audio, metrics, endpoint) = (
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_id.clone(),
#[cfg(feature = "audio")]
audio.as_ref().map(|(t, r)| (t.clone(), r.resubscribe())),
#[cfg(not(feature = "audio"))]
audio,
metrics.clone(),
endpoint.clone(),
);
let endpoint2 = endpoint.clone();
tokio::spawn(async move {
let _endpoint = endpoint2;
serve_viewer(
Box::new(crate::network::IrohTransport::new(
send,
recv,
conn,
endpoint.clone(),
)),
peer.parse()
.unwrap_or(SocketAddr::from(([127, 0, 0, 1], 0))),
"iroh",
tx,
published,
token,
viewers,
next_id,
audio,
None,
metrics,
approve,
None,
)
.await;
});
}
});
}
#[allow(clippy::too_many_arguments)]
fn spawn_accept_loop(
endpoint: quinn::Endpoint,
tx: EncodedBroadcast,
published: Shared,
token: SessionToken,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_id: Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))] audio: Option<()>,
metrics: crate::telemetry::SharedMetrics,
approve: bool,
redirect_to: Option<(String, String)>,
) {
tokio::spawn(async move {
let redirect_to = redirect_to;
let failures: Arc<Mutex<HashMap<SocketAddr, (u32, Instant)>>> =
Arc::new(Mutex::new(HashMap::new()));
while let Some(connecting) = endpoint.accept().await {
let peer = connecting.remote_address();
if rate_limited(&failures, peer) {
warn!("Refusing {peer}: too many failed handshakes");
continue;
}
match connecting.await {
Ok(connection) => {
let (tx, published, token, viewers, next_id, audio, metrics, redirect_to) = (
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_id.clone(),
#[cfg(feature = "audio")]
audio.as_ref().map(|(t, r)| (t.clone(), r.resubscribe())),
#[cfg(not(feature = "audio"))]
audio,
metrics.clone(),
redirect_to.clone(),
);
tokio::spawn(async move {
match connection.accept_bi().await {
Ok((send, recv)) => {
let transport = crate::network::QuicTransport::new(
send,
recv,
connection.clone(),
);
serve_viewer(
Box::new(transport),
peer,
"QUIC",
tx,
published,
token,
viewers,
next_id,
audio,
Some(connection),
metrics,
approve,
redirect_to,
)
.await;
}
Err(e) => warn!("Viewer {peer} never opened a stream: {e}"),
}
});
}
Err(e) => warn!("Failed to accept connection from {peer}: {e}"),
}
}
});
}
#[allow(clippy::too_many_arguments)]
fn spawn_relay_loop(
addr: SocketAddr,
pin: Vec<u8>,
session: String,
token: SessionToken,
tx: EncodedBroadcast,
published: Shared,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_id: Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))] audio: Option<()>,
metrics: crate::telemetry::SharedMetrics,
approve: bool,
_redirect_to: Option<(String, String)>,
) {
tokio::spawn(async move {
let resilience = crate::network::NetworkResilience::new(crate::network::ResilienceConfig {
max_retries: 0,
retry_delay: Duration::from_secs(3),
});
loop {
match RelayTransport::connect(
addr,
"pcc-relay",
&pin,
session.clone(),
token.clone(),
RelayRole::Host,
)
.await
{
Ok(transport) => {
info!("Relay connected (session '{session}')");
serve_relay_fan(
transport,
addr,
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_id.clone(),
#[cfg(feature = "audio")]
audio.as_ref().map(|(t, r)| (t.clone(), r.resubscribe())),
#[cfg(not(feature = "audio"))]
audio,
metrics.clone(),
approve,
)
.await;
warn!("Relay connection lost; reconnecting");
}
Err(e) => warn!("Failed to reach the relay: {e}"),
}
resilience.note_failure().await;
tokio::time::sleep(Duration::from_secs(3)).await;
}
});
}
struct FanSession {
id: u32,
tx: tokio::sync::mpsc::Sender<(u32, Vec<u8>)>,
rx: tokio::sync::mpsc::Receiver<Vec<u8>>,
}
impl FanSession {
fn new(
id: u32,
tx: tokio::sync::mpsc::Sender<(u32, Vec<u8>)>,
rx: tokio::sync::mpsc::Receiver<Vec<u8>>,
) -> Self {
Self { id, tx, rx }
}
}
struct FanSink {
id: u32,
tx: tokio::sync::mpsc::Sender<(u32, Vec<u8>)>,
}
struct FanSource {
rx: tokio::sync::mpsc::Receiver<Vec<u8>>,
}
#[async_trait::async_trait]
impl MessageSink for FanSink {
async fn send(&mut self, msg: &Message) -> anyhow::Result<()> {
self.send_encoded(&msg.encode()?).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> anyhow::Result<()> {
self.tx
.send((self.id, bytes.to_vec()))
.await
.map_err(|_| anyhow::anyhow!("relay fan is gone"))?;
Ok(())
}
}
#[async_trait::async_trait]
impl MessageSource for FanSource {
async fn recv(&mut self) -> anyhow::Result<Message> {
let bytes = self
.rx
.recv()
.await
.ok_or_else(|| anyhow::anyhow!("relay fan session closed"))?;
Ok(Message::decode(&bytes)?)
}
async fn recv_raw(&mut self) -> anyhow::Result<Vec<u8>> {
self.rx
.recv()
.await
.ok_or_else(|| anyhow::anyhow!("relay fan session closed"))
}
}
#[async_trait::async_trait]
impl MessageTransport for FanSession {
async fn send(&mut self, msg: &Message) -> anyhow::Result<()> {
self.send_encoded(&msg.encode()?).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> anyhow::Result<()> {
self.tx
.send((self.id, bytes.to_vec()))
.await
.map_err(|_| anyhow::anyhow!("relay fan is gone"))?;
Ok(())
}
async fn recv(&mut self) -> anyhow::Result<Message> {
let bytes = self
.rx
.recv()
.await
.ok_or_else(|| anyhow::anyhow!("relay fan session closed"))?;
Ok(Message::decode(&bytes)?)
}
fn split(self: Box<Self>) -> (Box<dyn MessageSink>, Box<dyn MessageSource>) {
let FanSession { id, tx, rx } = *self;
(Box::new(FanSink { id, tx }), Box::new(FanSource { rx }))
}
}
const FAN_SESSION_QUEUE: usize = 64;
#[allow(clippy::too_many_arguments)]
async fn serve_relay_fan(
transport: RelayTransport,
peer: SocketAddr,
tx: EncodedBroadcast,
published: Shared,
token: SessionToken,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_viewer_id: Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))] audio: Option<()>,
metrics: crate::telemetry::SharedMetrics,
approve: bool,
) {
let fan: RelayFan = Box::new(transport).into_fan();
let mut fan = fan;
let mut sessions: HashMap<u32, tokio::sync::mpsc::Sender<Vec<u8>>> = HashMap::new();
let (exit_tx, mut exit_rx) = tokio::sync::mpsc::channel::<u32>(64);
loop {
tokio::select! {
frame = fan.rx.recv() => {
let Some((id, bytes)) = frame else { break }; if id == RELAY_CONTROL_ID {
if bytes.len() == 5 {
let viewer = u32::from_le_bytes(bytes[1..5].try_into().unwrap());
if bytes[0] == 0xEE {
sessions.remove(&viewer);
} else if bytes[0] == K_CONTROL_BROWSER_HERE
&& !sessions.contains_key(&viewer)
{
sessions.insert(
viewer,
spawn_browser_session(
viewer, &fan.tx, &tx, &published, &token, &exit_tx,
),
);
} else if bytes[0] == 0xEF && !sessions.contains_key(&viewer) {
sessions.insert(
viewer,
spawn_fan_session(
viewer, peer, &fan.tx, &tx, &published, &token,
&viewers, &next_viewer_id, &audio, &metrics,
approve, &exit_tx,
),
);
}
}
continue;
}
let send = match sessions.get(&id) {
Some(send) => send.clone(),
None => {
let send = spawn_fan_session(
id, peer, &fan.tx, &tx, &published, &token,
&viewers, &next_viewer_id, &audio, &metrics,
approve, &exit_tx,
);
sessions.insert(id, send.clone());
send
}
};
if send.send(bytes).await.is_err() {
sessions.remove(&id);
}
}
exited = exit_rx.recv() => {
let Some(id) = exited else { break };
sessions.remove(&id);
}
}
}
sessions.clear();
}
#[allow(clippy::too_many_arguments)]
fn spawn_fan_session(
id: u32,
peer: SocketAddr,
fan_tx: &tokio::sync::mpsc::Sender<(u32, Vec<u8>)>,
tx: &EncodedBroadcast,
published: &Shared,
token: &SessionToken,
viewers: &Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_viewer_id: &Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: &Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))] audio: &Option<()>,
metrics: &crate::telemetry::SharedMetrics,
approve: bool,
exit_tx: &tokio::sync::mpsc::Sender<u32>,
) -> tokio::sync::mpsc::Sender<Vec<u8>> {
let (session_tx, session_rx) = tokio::sync::mpsc::channel::<Vec<u8>>(FAN_SESSION_QUEUE);
let (fan_tx, tx, published, token, viewers, next_id, audio, metrics, approve, exit_tx) = (
fan_tx.clone(),
tx.clone(),
published.clone(),
token.clone(),
viewers.clone(),
next_viewer_id.clone(),
#[cfg(feature = "audio")]
audio.as_ref().map(|(t, r)| (t.clone(), r.resubscribe())),
#[cfg(not(feature = "audio"))]
(*audio),
metrics.clone(),
approve,
exit_tx.clone(),
);
tokio::spawn(async move {
serve_viewer(
Box::new(FanSession::new(id, fan_tx, session_rx)),
peer,
&format!("relay id={id}"),
tx,
published,
token,
viewers,
next_id,
audio,
None,
metrics,
approve,
None,
)
.await;
let _ = exit_tx.send(id).await;
});
session_tx
}
struct FanBrowserChannel {
id: u32,
tx: tokio::sync::mpsc::Sender<(u32, Vec<u8>)>,
rx: tokio::sync::mpsc::Receiver<Vec<u8>>,
}
#[async_trait::async_trait]
impl web::FrameChannel for FanBrowserChannel {
async fn send_frame(&mut self, payload: &[u8]) -> Result<()> {
self.tx
.send((self.id, payload.to_vec()))
.await
.map_err(|_| anyhow::anyhow!("relay fan is gone"))
}
async fn recv_frame(&mut self) -> Result<Option<Vec<u8>>> {
Ok(self.rx.recv().await)
}
}
fn spawn_browser_session(
id: u32,
fan_tx: &tokio::sync::mpsc::Sender<(u32, Vec<u8>)>,
tx: &EncodedBroadcast,
published: &Shared,
token: &SessionToken,
exit_tx: &tokio::sync::mpsc::Sender<u32>,
) -> tokio::sync::mpsc::Sender<Vec<u8>> {
let (session_tx, session_rx) = tokio::sync::mpsc::channel::<Vec<u8>>(FAN_SESSION_QUEUE);
let updates = tx.subscribe();
let source = published.clone();
let token = token.as_str().to_string();
let fan_tx = fan_tx.clone();
let exit_tx = exit_tx.clone();
tokio::spawn(async move {
let snapshot: web::SnapshotFn = Arc::new(move || match try_current_snapshot(&source) {
Ok(msgs) => msgs,
Err(e) => {
warn!("Could not build a browser snapshot: {e}");
Vec::new()
}
});
let mut channel = FanBrowserChannel {
id,
tx: fan_tx,
rx: session_rx,
};
if let Err(e) = web::run_browser_session(&mut channel, &updates, snapshot, token).await {
warn!("relay browser viewer {id} ended: {e}");
}
let _ = exit_tx.send(id).await;
});
session_tx
}
#[allow(clippy::too_many_arguments)]
fn split_relays(spec: &str) -> Vec<String> {
spec.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string)
.collect()
}
fn direct_viewer_count(viewers: &Arc<Mutex<HashMap<u64, ViewerStats>>>) -> usize {
viewers.lock().values().filter(|v| v.direct).count()
}
fn redirect_target(args: &ShareArgs) -> Option<(String, String)> {
let relay = split_relays(args.relay.as_deref()?).first()?.clone();
let session = args.session.clone().unwrap_or_else(generate_session_code);
Some((relay, session))
}
fn is_stdin_closed(stdin: &std::io::Stdin) -> bool {
use std::io::BufRead;
let mut buf = [0u8; 1];
let mut locked = stdin.lock();
matches!(locked.fill_buf(), Ok(b) if b.is_empty())
&& matches!(std::io::Read::read(&mut locked, &mut buf), Ok(0))
}
#[allow(clippy::too_many_arguments)]
async fn serve_viewer(
transport: Box<dyn MessageTransport>,
peer: SocketAddr,
path: &str,
tx: EncodedBroadcast,
published: Shared,
token: SessionToken,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
next_viewer_id: Arc<AtomicU64>,
#[cfg(feature = "audio")] audio: Option<(
AudioBroadcast,
broadcast::Receiver<Arc<EncodedAudio>>,
)>,
#[cfg(not(feature = "audio"))]
#[cfg_attr(not(feature = "audio"), allow(unused_variables))]
audio: Option<()>,
#[cfg_attr(not(feature = "audio"), allow(unused_variables))] quic: Option<quinn::Connection>,
#[cfg_attr(not(feature = "audio"), allow(unused_variables))]
metrics: crate::telemetry::SharedMetrics,
approve: bool,
redirect_to: Option<(String, String)>,
) {
let mut rx = tx.subscribe();
let (mut sink, mut source) = transport.split();
let label = format!("{path} viewer {peer}");
let handshake = tokio::time::timeout(HANDSHAKE_TIMEOUT, async {
match source.recv().await {
Ok(Message::Hello { token, resume }) => Some((token, resume)),
Ok(other) => {
let _ = sink
.send(&Message::Error(format!(
"Send Hello with your viewer token first (got a message with rev {:?})",
other.rev()
)))
.await;
None
}
Err(e) => {
warn!("{label} handshake read failed: {e}");
None
}
}
})
.await;
let (presented, resume) = match handshake {
Ok(Some((presented, resume))) if verify_token(&token, &presented) => (presented, resume),
_ => {
let _ = sink
.send(&Message::Error(
"Unauthorized. Check the --token this sharer printed.".into(),
))
.await;
let _ = tokio::time::timeout(Duration::from_secs(2), source.recv()).await;
warn!("{label} rejected: bad or missing viewer token");
return;
}
};
let _ = presented;
if approve {
let approved = tokio::task::spawn_blocking(move || {
use std::io::{BufRead, Write};
let stdin = std::io::stdin();
if is_stdin_closed(&stdin) {
return true;
}
let mut out = std::io::stdout();
let _ = writeln!(out, "Viewer {peer} requests the screen. Admit? [y/N]");
let _ = out.flush();
let mut line = String::new();
std::io::BufReader::new(stdin.lock())
.read_line(&mut line)
.is_ok()
&& matches!(line.trim().to_ascii_lowercase().as_str(), "y" | "yes")
})
.await
.unwrap_or(false);
if !approved {
let _ = sink
.send(&Message::Error("The sharer declined this viewer.".into()))
.await;
let _ = tokio::time::timeout(Duration::from_secs(2), source.recv()).await;
warn!("{label} declined by approval prompt");
return;
}
}
let id = next_viewer_id.fetch_add(1, Ordering::Relaxed);
{
let mut guard = viewers.lock();
guard.insert(
id,
ViewerStats {
direct: quic.is_some(),
..Default::default()
},
);
}
info!("{label} authorized");
let host_keys = crate::network::e2e::KeyPair::generate();
let mut sink: Box<dyn crate::network::MessageSink> = sink;
let mut source: Box<dyn crate::network::MessageSource> = source;
let session = match crate::network::e2e::host_handshake(
&mut sink,
&mut source,
host_keys,
token.as_str(),
)
.await
{
Ok(s) => s,
Err(e) => {
warn!("{label} encryption handshake failed: {e}");
return;
}
};
let mut sink: Box<dyn crate::network::MessageSink> = Box::new(crate::network::SealedSink::new(
sink,
session.host_to_viewer,
));
let mut source: Box<dyn crate::network::MessageSource> = Box::new(
crate::network::SealedSource::new(source, session.viewer_to_host),
);
#[cfg(feature = "audio")]
if let (Some((_, _)), Some(mut arx), Some(connection)) = (
audio.as_ref(),
audio.as_ref().map(|(_, r)| r.resubscribe()),
quic.as_ref(),
) {
let audio_label = label.clone();
let send_conn = connection.clone();
let audio_metrics = metrics.clone();
tokio::spawn(async move {
loop {
let frame = match arx.recv().await {
Ok(f) => f,
Err(broadcast::error::RecvError::Lagged(n)) => {
warn!("{audio_label} audio fell behind by {n} frames");
continue;
}
Err(broadcast::error::RecvError::Closed) => break,
};
let payload = match crate::audio::transport::AudioSender::datagram_payload(&frame) {
Ok(payload) => payload,
Err(e) => {
warn!("{audio_label} oversize audio frame dropped: {e:#}");
audio_metrics.audio_frames_dropped.incr();
continue;
}
};
if let Err(e) = send_conn.send_datagram(bytes::Bytes::from(payload)) {
tracing::debug!("{audio_label} audio datagram not queued: {e}");
audio_metrics.audio_frames_dropped.incr();
continue;
}
audio_metrics.audio_frames_sent.incr();
}
});
info!("{label} audio over QUIC datagrams enabled");
}
let mut floor = match resume {
Some((epoch, rev)) => {
let replay = {
let p = published.read().await;
if epoch == p.epoch {
p.ring.replay_from(rev, epoch)
} else {
None
}
};
match replay {
Some(msgs) => {
let mut caught = rev;
for bytes in &msgs {
if let Err(e) = sink.send_encoded(bytes).await {
warn!("{label} resume replay failed: {e}");
return;
}
if let Some(r) = crate::network::peek_rev(bytes) {
caught = caught.max(r);
}
}
info!(
"{label} resumed from rev {rev} with {} replayed messages",
msgs.len()
);
caught
}
None => match send_snapshot(&mut sink, &published).await {
Ok(rev) => rev,
Err(e) => {
warn!("{label} could not be caught up: {e}");
return;
}
},
}
}
None => match send_snapshot(&mut sink, &published).await {
Ok(rev) => rev,
Err(e) => {
warn!("{label} could not be caught up: {e}");
return;
}
},
};
if let Some(v) = viewers.lock().get_mut(&id) {
v.acked_rev = floor;
}
let (control_tx, mut control_rx) = tokio::sync::mpsc::channel::<Message>(8);
tokio::spawn(async move {
while let Ok(msg) = source.recv().await {
if control_tx.send(msg).await.is_err() {
break;
}
}
});
let quic_stats = quic.clone();
let mut last_lost: Option<u64> = None;
let mut consecutive_lags = 0u32;
let mut last_sample = Instant::now() - Duration::from_secs(1);
loop {
if viewers.lock().get(&id).is_some_and(|v| v.redirect_queued) {
if let Some((relay, session)) = redirect_to.clone() {
let _ = sink.send(&Message::Redirect { relay, session }).await;
info!("{label} redirected to relay (broadcast mode)");
}
break;
}
tokio::select! {
update = rx.recv() => match update {
Ok(bytes) => {
let Some(rev) = crate::network::peek_rev(&bytes) else {
continue;
};
if rev <= floor {
continue; }
if let Err(e) = sink.send_encoded(&bytes).await {
if let Some(v) = viewers.lock().get_mut(&id) {
v.lag_events += 1;
}
warn!("{label} write failed: {e}");
break;
}
if last_sample.elapsed() >= Duration::from_millis(250) {
if let Some(conn) = quic_stats.as_ref() {
let stats = conn.stats();
let lost = stats.path.lost_packets;
let is_new_loss = last_lost.is_some_and(|prev| lost > prev);
last_lost = Some(lost);
if let Some(v) = viewers.lock().get_mut(&id) {
v.rtt_ms = Some(conn.rtt().as_millis().min(u128::from(u64::MAX)) as u64);
v.lost_packets = lost;
v.cwnd_bytes = stats.path.cwnd;
if is_new_loss {
v.lag_events += 1;
}
}
last_sample = Instant::now();
}
}
}
Err(broadcast::error::RecvError::Lagged(missed)) => {
consecutive_lags += 1;
if let Some(v) = viewers.lock().get_mut(&id) {
v.lag_events += 1;
}
if consecutive_lags > MAX_CONSECUTIVE_LAGS {
warn!("{label} fell behind {consecutive_lags} times; disconnecting");
break;
}
warn!("{label} fell behind by {missed} updates; repairing");
match repair_viewer(&mut sink, &published, floor).await {
Ok(rev) => {
floor = rev;
consecutive_lags = 0;
}
Err(e) => {
warn!("{label} repair failed: {e}");
break;
}
}
}
Err(broadcast::error::RecvError::Closed) => break,
},
control = control_rx.recv() => {
let Some(msg) = control else { break };
match msg {
Message::RequestKeyframe => {
match send_snapshot(&mut sink, &published).await {
Ok(rev) => floor = rev,
Err(e) => { warn!("{label} refresh failed: {e}"); break; }
}
}
Message::Ack { rev } => {
if let Some(v) = viewers.lock().get_mut(&id) {
v.acked_rev = v.acked_rev.max(rev);
}
}
Message::Bye => break,
Message::Hello { .. } => {
let _ = sink.send(&Message::Error("Already authorized.".into())).await;
}
Message::Error(e) => warn!("{label} reported: {e}"),
_ => {}
}
}
}
}
viewers.lock().remove(&id);
info!("{label} disconnected");
}
pub async fn repair_viewer(
sink: &mut Box<dyn MessageSink>,
published: &Shared,
floor: Rev,
) -> Result<Rev> {
let replay = {
let p = published.read().await;
p.ring.replay_from(floor, p.epoch)
};
let Some(msgs) = replay else {
return send_snapshot(sink, published).await;
};
let mut new_floor = floor;
for bytes in &msgs {
sink.send_encoded(bytes).await?;
if let Some(rev) = crate::network::peek_rev(bytes) {
new_floor = new_floor.max(rev);
}
}
Ok(new_floor)
}
async fn send_snapshot(sink: &mut Box<dyn MessageSink>, published: &Shared) -> Result<Rev> {
let (rev, epoch, width, height, data) = {
let mut p = published.write().await;
let rev = p.rev;
if p.encoded_for(rev).is_none() {
p.encoded = Some((
rev,
Arc::new(
crate::encoder::encode_snapshot(
p.snapshot.width,
p.snapshot.height,
&p.snapshot.rgb,
)?
.1,
),
));
}
let encoded = p.encoded_for(rev).expect("just produced");
(
rev,
p.epoch,
p.snapshot.width,
p.snapshot.height,
encoded.clone(),
)
};
for msg in snapshot_messages(width, height, &data, rev, epoch)? {
sink.send_encoded(&msg).await?;
}
Ok(rev)
}
fn snapshot_messages(
width: u32,
height: u32,
data: &[u8],
rev: Rev,
epoch: Epoch,
) -> Result<Vec<Vec<u8>>> {
let chunks: Vec<&[u8]> = data.chunks(SNAPSHOT_CHUNK_BYTES).collect();
if chunks.is_empty() || chunks.len() as u32 > crate::network::MAX_SNAPSHOT_CHUNKS {
anyhow::bail!(
"snapshot of {} bytes needs {} chunks (max {})",
data.len(),
chunks.len().max(1),
crate::network::MAX_SNAPSHOT_CHUNKS
);
}
let mut out = Vec::with_capacity(chunks.len() + 2);
out.push(
Message::SnapshotBegin {
rev,
pts_us: 0,
epoch,
width,
height,
format: crate::encoder::SnapshotFormat::Png,
total_len: data.len() as u32,
chunks: chunks.len() as u32,
}
.encode()?,
);
for (index, chunk) in chunks.iter().enumerate() {
out.push(
Message::SnapshotChunk {
rev,
index: index as u32,
data: chunk.to_vec(),
}
.encode()?,
);
}
out.push(
Message::SnapshotCommit {
rev,
pts_us: 0,
epoch,
}
.encode()?,
);
Ok(out)
}
#[allow(clippy::too_many_arguments)]
async fn capture_loop(
metrics: crate::telemetry::SharedMetrics,
mut first: Frame,
capture: CaptureSource,
mut cursor_sampler: Box<dyn crate::capture::CursorSampler>,
planner: Planner,
published: Shared,
tx: EncodedBroadcast,
viewers: Arc<Mutex<HashMap<u64, ViewerStats>>>,
surface: Option<Arc<SharedSurface>>,
target_quality: f32,
requested_fps: u32,
max_fps: u32,
repair_interval: Duration,
broadcast_above: usize,
redirect_to: Option<(String, String)>,
) -> Result<()> {
let mut reference: Arc<Vec<u8>> = Arc::new(std::mem::take(&mut first.data));
let mut surface_size = (first.width, first.height);
let mut epoch: Epoch = 0;
let mut rev: Rev = 0;
let mut pending_snapshot = true;
let mut effective_fps = requested_fps;
let mut last_total_lost: u64 = 0;
let mut last_repair = Instant::now();
let mut quality = target_quality;
let mut pressure = Pressure {
value: 0.0,
last_change: Instant::now(),
};
let mut last_cursor: Option<(u32, u32)> = None;
let mut cursor_hidden = false;
loop {
let loop_start = Instant::now();
let frame = capture.capture_frame()?;
if surface_size != (frame.width, frame.height) {
epoch += 1;
metrics.epoch_bumps.incr();
warn!(
"Capture geometry changed to {}x{}; starting epoch {epoch}",
frame.width, frame.height
);
reference = Arc::new(vec![0u8; rgb_len(frame.width, frame.height)?]);
surface_size = (frame.width, frame.height);
pending_snapshot = true;
}
if pending_snapshot && surface_size == (0, 0) {
reference = Arc::new(frame.data.clone());
}
let snapshot_bytes = published
.read()
.await
.encoded
.as_ref()
.map(|(_, d)| d.len())
.unwrap_or(usize::MAX);
let (width, height) = surface_size;
let rgb: &mut Vec<u8> = Arc::make_mut(&mut reference);
let detect_start = Instant::now();
let plan = planner.plan(
rgb,
width,
height,
&frame,
PlanLimits {
snapshot_bytes,
max_update_bytes: SEND_BUDGET,
},
)?;
let total = detect_start.elapsed();
metrics.detect.record_duration(plan.detect);
metrics
.plan
.record_duration(total.saturating_sub(plan.detect));
metrics.frames_total.incr();
let total_pixels = (width as f64) * (height as f64);
metrics.changed_area_fraction.set(if total_pixels > 0.0 {
(plan.changed_pixels as f64 / total_pixels).min(1.0)
} else {
0.0
});
let idle = plan.ops.is_empty() && plan.wire_len == 0;
let wants_snapshot = !idle && plan.ops.is_empty();
let repair_due = last_repair.elapsed() >= repair_interval;
let changed_fraction = plan.changed_fraction(width, height);
let changed_bounds = plan.changed_bounds();
if pending_snapshot || wants_snapshot || (!idle && repair_due && snapshot_bytes > 0) {
if wants_snapshot {
reference = Arc::new(frame.data.clone());
}
rev += 1;
last_repair = Instant::now();
pending_snapshot = false;
metrics.frames_snapshot_fallback.incr();
let encode_start = Instant::now();
for msg in publish_snapshot(&published, &reference, width, height, rev, epoch).await? {
metrics.observe_message_bytes(msg.len());
metrics.bytes_snapshot.add(msg.len() as u64);
let shared = Arc::new(msg);
published
.write()
.await
.ring
.push(rev, epoch, shared.clone());
let _ = tx.send(shared);
}
metrics.encode.record_duration(encode_start.elapsed());
} else if !idle {
rev += 1;
for op in &plan.ops {
let n = op.wire_len() as u64;
match op {
crate::network::WireOp::Fill { .. } => metrics.bytes_fill.add(n),
crate::network::WireOp::Copy { .. } => metrics.bytes_copy.add(n),
crate::network::WireOp::Rect { .. } => metrics.bytes_patch.add(n),
}
}
let serialize_start = Instant::now();
let bytes = Message::PartialUpdate {
rev,
pts_us: 0,
epoch,
ops: plan.ops,
}
.encode();
match bytes {
Ok(bytes) => {
metrics.observe_message_bytes(bytes.len());
metrics.serialize.record_duration(serialize_start.elapsed());
let shared = Arc::new(bytes);
published
.write()
.await
.ring
.push(rev, epoch, shared.clone());
let _ = tx.send(shared);
}
Err(e) => {
warn!("Update rejected ({e}); sending a snapshot instead");
reference = Arc::new(frame.data.clone());
rev += 1;
last_repair = Instant::now();
for msg in
publish_snapshot(&published, &reference, width, height, rev, epoch).await?
{
let shared = Arc::new(msg);
published
.write()
.await
.ring
.push(rev, epoch, shared.clone());
let _ = tx.send(shared);
}
}
}
} else {
metrics.frames_idle.incr();
if let Ok(bytes) = (Message::KeepAlive { rev }).encode() {
metrics.bytes_keepalive.add(bytes.len() as u64);
let shared = Arc::new(bytes);
published
.write()
.await
.ring
.push_keepalive(rev, epoch, shared.clone());
let _ = tx.send(shared);
}
}
match cursor_sampler.sample(width, height) {
Some(c)
if last_cursor.is_none_or(|(x, y)| x.abs_diff(c.x).max(y.abs_diff(c.y)) >= 2) =>
{
last_cursor = Some((c.x, c.y));
cursor_hidden = false;
if let Ok(bytes) = (Message::CursorMove { x: c.x, y: c.y }).encode() {
let shared = Arc::new(bytes);
published
.write()
.await
.ring
.push_keepalive(rev, epoch, shared.clone());
let _ = tx.send(shared);
}
}
None if !cursor_hidden => {
cursor_hidden = true;
last_cursor = None;
if let Ok(bytes) = Message::CursorHide.encode() {
let shared = Arc::new(bytes);
published
.write()
.await
.ring
.push_keepalive(rev, epoch, shared.clone());
let _ = tx.send(shared);
}
}
_ => {}
}
if changed_fraction >= MOTION_PREVIEW_FRACTION {
if let Some((x, y, w, h)) = changed_bounds {
let preview = Message::MotionPreview {
rev,
x,
y,
width: w,
height: h,
changed_fraction: changed_fraction as f32,
};
if let Ok(bytes) = preview.encode() {
metrics.bytes_preview.add(bytes.len() as u64);
let _ = tx.send(Arc::new(bytes));
}
}
}
{
let mut p = published.write().await;
p.rev = rev;
p.epoch = epoch;
p.snapshot = Arc::new(SurfaceSnapshot {
width,
height,
rgb: reference.clone(),
});
}
if let Some(surface) = &surface {
surface.publish(&frame, quality);
}
pressure.decay(loop_start.elapsed());
let current = rev;
let stats: Vec<ViewerStats> = viewers.lock().values().cloned().collect();
let lagging = stats.iter().filter(|v| v.lag_events > 0).count();
let stalled = stats
.iter()
.filter(|v| v.lag_events > 0 && current.saturating_sub(v.acked_rev) > 30)
.count();
let worst_rtt_ms = stats.iter().filter_map(|v| v.rtt_ms).max().unwrap_or(0);
let total_lost: u64 = stats.iter().map(|v| v.lost_packets).sum();
let min_cwnd = stats.iter().map(|v| v.cwnd_bytes).filter(|c| *c > 0).min();
let cwnd_collapsed = min_cwnd.is_some_and(|c| c < 64 * 1024);
let new_loss = total_lost > last_total_lost;
last_total_lost = total_lost;
if lagging > 0 {
pressure.penalise(lagging as f32 + stalled as f32);
}
if worst_rtt_ms >= 250 || new_loss {
pressure.penalise(1.0);
}
let step = decide_congestion_step(
effective_fps,
requested_fps,
worst_rtt_ms,
new_loss,
cwnd_collapsed,
loop_start.elapsed() > frame_interval_at(effective_fps) * 2,
);
if pressure.is_pressured() && step.fps != effective_fps {
effective_fps = step.fps;
info!("{lagging} viewer(s) falling behind; dropping to {effective_fps} fps");
pressure.last_change = Instant::now();
} else if !pressure.is_pressured()
&& step.fps != effective_fps
&& pressure.last_change.elapsed() > Duration::from_secs(5)
{
effective_fps = step.fps;
pressure.last_change = Instant::now();
}
if broadcast_above > 0
&& redirect_to.is_some()
&& direct_viewer_count(&viewers) > broadcast_above
{
let mut guard = viewers.lock();
if let Some((_, v)) = guard
.iter_mut()
.filter(|(_, v)| v.direct && !v.redirect_queued)
.max_by_key(|(id, _)| **id)
{
v.redirect_queued = true;
}
}
metrics.viewer_rtt_ms.set_raw(worst_rtt_ms);
let max_gap = stats
.iter()
.map(|v| current.saturating_sub(v.acked_rev))
.max()
.unwrap_or(0);
metrics.viewer_queue_depth.set_raw(max_gap);
metrics
.oldest_pending_ms
.set_raw(max_gap.saturating_mul(1_000 / u64::from(effective_fps.max(1))));
let frame_interval = frame_interval_at(effective_fps);
if step.degrade_quality && quality > 0.3 {
quality = if cwnd_collapsed {
(quality - 0.2).max(0.3)
} else {
(quality - 0.1).max(0.3)
};
let cfg = crate::pcc::QualityConfig {
target_fps: effective_fps,
max_fps,
quality,
};
if let Ok(bytes) = Message::QualityConfig(cfg).encode() {
let _ = tx.send(Arc::new(bytes));
}
}
let elapsed = loop_start.elapsed();
if elapsed > frame_interval {
metrics.loop_overruns.incr();
}
metrics.captured_frames.incr();
metrics
.effective_fps
.set(1.0 / frame_interval.as_secs_f64());
if elapsed < frame_interval {
tokio::time::sleep(frame_interval - elapsed).await;
}
}
}
fn frame_interval_at(fps: u32) -> Duration {
Duration::from_secs(1) / fps.max(1)
}
#[derive(Debug, Clone, Copy, PartialEq)]
struct CongestionDecision {
fps: u32,
degrade_quality: bool,
}
fn decide_congestion_step(
fps: u32,
requested_fps: u32,
worst_rtt_ms: u64,
new_loss: bool,
cwnd_collapsed: bool,
loop_overran: bool,
) -> CongestionDecision {
let congested_path = worst_rtt_ms >= 250 || new_loss;
let mut fps = fps;
if congested_path && fps > 10 {
fps = (fps * 3 / 4).max(10);
} else if !congested_path && fps < requested_fps && worst_rtt_ms < 100 {
fps = (fps * 4 / 3).min(requested_fps);
}
let degrade_quality = loop_overran || congested_path || cwnd_collapsed;
CongestionDecision {
fps,
degrade_quality,
}
}
async fn resolve_ladder(policy: crate::reach::ReachPolicy) -> Option<Rung> {
if policy == crate::reach::ReachPolicy::RelayOnly {
return None;
}
let report = crate::reach::diagnose(None).await;
let rung = if report.has_global_ipv6 {
Rung::Ipv6Direct
} else if report.reflexive.is_some() {
Rung::StunIce
} else {
Rung::Relay
};
for note in &report.notes {
info!("reach: {note}");
}
Some(rung)
}
fn try_current_snapshot(published: &Shared) -> Result<Vec<Vec<u8>>> {
let mut state = published
.try_write()
.map_err(|_| anyhow::anyhow!("the capture loop holds the surface; try again"))?;
let width = state.snapshot.width;
let height = state.snapshot.height;
let rev = state.rev;
let epoch = state.epoch;
let rgb = state.snapshot.rgb.clone();
if state.encoded_for(rev).is_none() {
state.encoded = Some((
rev,
Arc::new(crate::encoder::encode_snapshot(width, height, &rgb)?.1),
));
}
let data = state.encoded_for(rev).expect("just produced").clone();
snapshot_messages(width, height, &data, rev, epoch)
}
async fn publish_snapshot(
published: &Shared,
reference: &Arc<Vec<u8>>,
width: u32,
height: u32,
rev: Rev,
epoch: Epoch,
) -> Result<Vec<Vec<u8>>> {
let data = {
let mut p = published.write().await;
p.snapshot = Arc::new(SurfaceSnapshot {
width,
height,
rgb: reference.clone(),
});
if p.encoded_for(rev).is_none() {
p.encoded = Some((
rev,
Arc::new(crate::encoder::encode_snapshot(width, height, reference)?.1),
));
}
p.encoded_for(rev).expect("just produced").clone()
};
snapshot_messages(width, height, &data, rev, epoch)
}
fn rate_limited(
failures: &Arc<Mutex<HashMap<SocketAddr, (u32, Instant)>>>,
peer: SocketAddr,
) -> bool {
let mut map = failures.lock();
let entry = map.entry(peer).or_insert((0, Instant::now()));
if entry.1.elapsed() > AUTH_FAILURE_WINDOW {
*entry = (0, Instant::now());
}
entry.0 += 1;
entry.0 > AUTH_FAILURES_ALLOWED
}
#[cfg(test)]
mod tests {
use super::*;
use crate::network::PROTOCOL_VERSION;
#[test]
fn broadcast_redirect_needs_a_relay() {
let mut args = ShareArgs {
listen: None,
relay: None,
relay_pin: String::new(),
session: Some("ABC123".into()),
web: None,
web_cert: None,
synthetic: true,
capture_target: crate::capture::CaptureTarget::Display(0),
fps: 30,
max_fps: 60,
quality: 0.8,
token: crate::network::SessionToken::parse("TOKENTOKENTOKENTOKENTOKENTOKENTOK1")
.unwrap(),
repair_interval: std::time::Duration::from_secs(30),
reach: crate::reach::ReachPolicy::TryDirect,
approve: false,
broadcast_above: 2,
audio_source: crate::audio::AudioSource::None_,
transport: crate::network::TransportKind::Quic,
};
assert!(redirect_target(&args).is_none());
args.relay = Some("r1:1, r2:2".into());
assert_eq!(
redirect_target(&args),
Some(("r1:1".to_string(), "ABC123".to_string()))
);
}
#[test]
fn redirect_flags_only_unflagged_direct_viewers() {
let viewers: Arc<Mutex<HashMap<u64, ViewerStats>>> = Arc::new(Mutex::new(HashMap::new()));
assert_eq!(direct_viewer_count(&viewers), 0);
for (id, direct) in [(1u64, true), (2, true), (3, false)] {
viewers.lock().insert(
id,
ViewerStats {
direct,
..Default::default()
},
);
}
assert_eq!(direct_viewer_count(&viewers), 2);
}
#[test]
fn relay_list_splits_and_skips_empties() {
assert_eq!(split_relays("a:1"), vec!["a:1"]);
assert_eq!(split_relays("a:1, b:2 ,c:3"), vec!["a:1", "b:2", "c:3"]);
assert!(split_relays(" , ,").is_empty());
assert_eq!(split_relays("b:2,a:1,b:2"), vec!["b:2", "a:1", "b:2"]);
}
#[test]
fn closed_stdin_counts_as_closed() {
use std::io::BufRead;
let empty: &[u8] = &[];
let mut locked = std::io::BufReader::new(empty);
let mut buf = Vec::new();
assert_eq!(locked.read_until(b'x', &mut buf).unwrap(), 0);
}
#[test]
fn high_rtt_backs_off_before_queues_fill() {
let d = decide_congestion_step(30, 30, 300, false, false, false);
assert_eq!(d.fps, 22, "a 300 ms path must shed fps, got {}", d.fps);
assert!(d.degrade_quality);
}
#[test]
fn new_loss_backs_off_before_queues_fill() {
let d = decide_congestion_step(30, 30, 40, true, false, false);
assert_eq!(d.fps, 22);
assert!(d.degrade_quality);
}
#[test]
fn healthy_path_recovers_and_holds_quality() {
let d = decide_congestion_step(22, 30, 20, false, false, false);
assert_eq!(d.fps, 29, "recovery multiplies by 4/3, got {}", d.fps);
assert!(!d.degrade_quality);
}
#[test]
fn recovery_does_not_ramp_into_a_bad_path() {
let d = decide_congestion_step(22, 30, 300, false, false, false);
assert_eq!(
d.fps, 16,
"congested step wins over recovery, got {}",
d.fps
);
}
#[test]
fn collapsed_window_degrades_quality_without_touching_fps() {
let d = decide_congestion_step(30, 30, 40, false, true, false);
assert_eq!(d.fps, 30);
assert!(d.degrade_quality);
}
fn web_tls_gate(addr: &str, has_cert: bool) -> bool {
let addr: SocketAddr = addr.parse().unwrap();
has_cert || addr.ip().is_loopback()
}
#[test]
fn loopback_without_tls_is_allowed() {
assert!(web_tls_gate("127.0.0.1:8080", false));
}
#[test]
fn remote_without_tls_is_refused() {
assert!(!web_tls_gate("0.0.0.0:8080", false));
assert!(!web_tls_gate("192.168.1.20:8080", false));
}
#[test]
fn remote_with_tls_is_allowed() {
assert!(web_tls_gate("0.0.0.0:8080", true));
}
#[test]
fn peek_rev_reads_a_partial_update_revision() {
let bytes = Message::PartialUpdate {
rev: 42,
pts_us: 0,
epoch: 1,
ops: vec![],
}
.encode()
.unwrap();
assert_eq!(crate::network::peek_rev(&bytes), Some(42));
}
#[test]
fn peek_rev_reads_a_keepalive_revision() {
let bytes = Message::KeepAlive { rev: 7 }.encode().unwrap();
assert_eq!(crate::network::peek_rev(&bytes), Some(7));
}
#[test]
fn peek_rev_ignores_messages_without_a_revision() {
let bytes = Message::Bye.encode().unwrap();
assert_eq!(crate::network::peek_rev(&bytes), None);
assert_eq!(crate::network::peek_rev(b""), None);
assert_eq!(
crate::network::peek_rev(&[PROTOCOL_VERSION, 0, 0, 0, 0]),
None
);
}
}