pixel-change-check-client 0.1.2

Replicates your screen exactly and sends only the pixels that changed, over QUIC or a relay, to a native or browser viewer.
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
//! A relay for when a sharer and a viewer can't reach each other directly
//! (both behind NAT, no port forwarding).
//!
//! Both sides make *outbound* TLS connections to a relay host and register
//! with the same session code and the same viewer token. The relay pairs a
//! host with any number of viewers and forwards framed `Message` bytes. It
//! never inspects frame contents, so a future end-to-end encryption layer
//! needs no changes here.
//!
//! What it does enforce, because it faces untrusted networks:
//!
//! * every length prefix is checked against the frame budget *before* the
//!   buffer is allocated;
//! * the session map is never held across an `await`, so one slow viewer
//!   cannot stall another session;
//! * each viewer has its own byte budget, and a viewer that stops draining
//!   is disconnected rather than allowed to grow a queue forever;
//! * registrations are generational, so a reconnecting host's cleanup
//!   cannot erase its own replacement;
//! * reader and writer are supervised together, so EOF on either side
//!   tears the whole connection down.

use crate::network::{
    read_len_prefix, verify_token, write_encoded, Message, MessageSink, MessageSource,
    MessageTransport, NetworkConfig, ServerIdentity, SessionToken,
};
use anyhow::{Context, Result};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::mpsc;
use tokio::sync::Mutex;
use tokio_rustls::client::TlsStream;
use tracing::{info, warn};

/// Sent once, immediately after connecting, to identify the session, the
/// role, and to authorize the connection.
#[derive(Debug, Serialize, Deserialize)]
pub struct RelayRegister {
    pub session: String,
    pub role: RelayRole,
    pub token: String,
}

#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum RelayRole {
    Host,
    Viewer,
}

/// Ceiling on a registration envelope.
const MAX_REGISTRATION_BYTES: usize = 4096;

/// Queued bytes one peer may fall behind by before it is disconnected.
///
/// Receipt: a 4K snapshot is at most a few MB. 8 MiB is roughly two
/// worst-case snapshots, so a viewer on a real link never trips this, while
/// a viewer that has stopped reading cannot pin an unbounded amount of
/// memory in the relay.
const MAX_PEER_QUEUE_BYTES: usize = 8 * 1024 * 1024;

/// Queued messages per peer, so many tiny messages cannot be cheaper to
/// buffer than a few large ones.
const MAX_PEER_QUEUE_MESSAGES: usize = 64;

/// Maximum concurrent sessions the relay will host.
const MAX_SESSIONS: usize = 1024;

/// Maximum viewers per session.
const MAX_VIEWERS_PER_SESSION: usize = 64;

/// Failed registrations allowed from one peer address inside the window.
const AUTH_FAILURES_ALLOWED: u32 = 8;
const AUTH_FAILURE_WINDOW: Duration = Duration::from_secs(60);

/// A session with no traffic at all is forgotten.
const SESSION_IDLE_TTL: Duration = Duration::from_secs(3600);

/// Short rendezvous code used to find a session. Not a secret on its own --
/// the token is what authorizes.
pub fn generate_session_code() -> String {
    use rand::Rng;
    const CHARS: &[u8] = b"ABCDEFGHJKLMNPQRSTUVWXYZ23456789"; // no ambiguous chars
    let mut rng = rand::thread_rng();
    (0..6)
        .map(|_| CHARS[rng.gen_range(0..CHARS.len())] as char)
        .collect()
}

#[derive(Clone)]
struct Peer {
    /// Monotonic per connection. Cleanup only removes the registration it
    /// created, so a reconnecting host's predecessor cannot delete it.
    gen: u64,
    tx: mpsc::Sender<Vec<u8>>,
}

struct Session {
    host: Option<Peer>,
    viewers: Vec<Peer>,
    last_seen: Instant,
}

type Sessions = Arc<Mutex<HashMap<String, Session>>>;

/// A TLS-wrapped TCP connection to a relay, already registered and
/// authorized.
pub struct RelayTransport {
    stream: TlsStream<TcpStream>,
}

impl RelayTransport {
    /// Dial a relay, present the token, and wait for it to accept.
    pub async fn connect(
        relay_addr: std::net::SocketAddr,
        server_name: &str,
        pin: &[u8],
        session: String,
        token: SessionToken,
        role: RelayRole,
    ) -> Result<Self> {
        let config = NetworkConfig::client_tls_config(pin)?;
        let server_name = rustls::pki_types::ServerName::try_from(server_name.to_owned())
            .map_err(|e| anyhow::anyhow!("Invalid relay server name '{server_name}': {e}"))?;
        let connector = tokio_rustls::TlsConnector::from(Arc::new(config));
        let tcp = TcpStream::connect(relay_addr)
            .await
            .with_context(|| format!("Failed to connect to the relay at {relay_addr}"))?;
        let mut stream = connector
            .connect(server_name, tcp)
            .await
            .with_context(|| format!("TLS handshake with the relay at {relay_addr} failed"))?;

        let reg = RelayRegister {
            session,
            role,
            token: token.as_str().to_string(),
        };
        let bytes = bincode::serialize(&reg)?;
        stream
            .write_all(&(bytes.len() as u32).to_le_bytes())
            .await?;
        stream.write_all(&bytes).await?;
        stream.flush().await?;

        // The relay acknowledges a good registration with the protocol
        // version byte, so "connected" means "authorized" rather than
        // merely "TCP accepted".
        let mut ack = [0u8; 1];
        tokio::time::timeout(Duration::from_secs(10), stream.read_exact(&mut ack))
            .await
            .context("Timed out waiting for the relay to accept the registration")?
            .context("The relay closed the connection during registration")?;
        if ack[0] != crate::network::PROTOCOL_VERSION {
            anyhow::bail!(
                "The relay rejected this viewer: check --session and --token \
                 (relay speaks protocol {}, offered {})",
                ack[0],
                crate::network::PROTOCOL_VERSION
            );
        }

        Ok(Self { stream })
    }
}

struct RelaySink {
    stream: tokio::io::WriteHalf<TlsStream<TcpStream>>,
}

#[async_trait]
impl MessageSink for RelaySink {
    async fn send(&mut self, msg: &Message) -> Result<()> {
        self.send_encoded(&msg.encode()?).await
    }

    async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
        write_encoded(&mut self.stream, bytes).await
    }
}

struct RelaySource {
    stream: tokio::io::ReadHalf<TlsStream<TcpStream>>,
}

#[async_trait]
impl MessageSource for RelaySource {
    async fn recv(&mut self) -> Result<Message> {
        Message::read_framed(&mut self.stream).await
    }

    /// Read one envelope without decoding it, which is what a sealed
    /// session needs: the ciphertext is not a `Message` yet.
    async fn recv_raw(&mut self) -> Result<Vec<u8>> {
        Message::read_envelope(&mut self.stream).await
    }
}

#[async_trait]
impl MessageTransport for RelayTransport {
    async fn send(&mut self, msg: &Message) -> Result<()> {
        write_encoded(&mut self.stream, &msg.encode()?).await
    }

    async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
        write_encoded(&mut self.stream, bytes).await
    }

    async fn recv(&mut self) -> Result<Message> {
        Message::read_framed(&mut self.stream).await
    }

    fn split(self: Box<Self>) -> (Box<dyn MessageSink>, Box<dyn MessageSource>) {
        let (read, write) = tokio::io::split(self.stream);
        (
            Box::new(RelaySink { stream: write }),
            Box::new(RelaySource { stream: read }),
        )
    }
}

/// Run a relay server until the listener errors.
///
/// The listener is supplied by the caller so the bound address is known
/// before the server starts: binding here would let the caller print an
/// address that turned out to be unavailable.
pub async fn run_relay_server(
    listener: TcpListener,
    identity: Arc<ServerIdentity>,
    token: SessionToken,
) -> Result<()> {
    let acceptor =
        tokio_rustls::TlsAcceptor::from(Arc::new(NetworkConfig::server_crypto_config(&identity)?));
    info!("Relay server listening on {} (TLS)", listener.local_addr()?);
    info!(
        "Relay certificate fingerprint (sha256): {}",
        identity.fingerprint
    );

    let sessions: Sessions = Arc::new(Mutex::new(HashMap::new()));
    let next_gen = Arc::new(AtomicU64::new(1));
    let failures: Arc<Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>> =
        Arc::new(Mutex::new(HashMap::new()));

    {
        let sessions = sessions.clone();
        let failures = failures.clone();
        tokio::spawn(async move { sweep(sessions, failures).await });
    }

    loop {
        let (tcp, peer) = listener.accept().await?;
        let acceptor = acceptor.clone();
        let sessions = sessions.clone();
        let next_gen = next_gen.clone();
        let token = token.clone();
        let failures = failures.clone();
        tokio::spawn(async move {
            if let Err(e) =
                handle_client(tcp, peer, acceptor, sessions, next_gen, token, failures).await
            {
                warn!("Relay peer {peer}: {e}");
            }
        });
    }
}

/// Retire sessions that have gone quiet, and forget auth-failure counters
/// that have aged out.
async fn sweep(
    sessions: Sessions,
    failures: Arc<tokio::sync::Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>>,
) {
    let mut ticker = tokio::time::interval(Duration::from_secs(60));
    loop {
        ticker.tick().await;
        {
            let mut map = sessions.lock().await;
            map.retain(|_, s| s.last_seen.elapsed() < SESSION_IDLE_TTL);
        }
        {
            let mut map = failures.lock().await;
            map.retain(|_, (_, at)| at.elapsed() < AUTH_FAILURE_WINDOW);
        }
    }
}

async fn handle_client(
    tcp: TcpStream,
    peer: std::net::SocketAddr,
    acceptor: tokio_rustls::TlsAcceptor,
    sessions: Sessions,
    next_gen: Arc<AtomicU64>,
    expected_token: SessionToken,
    failures: Arc<tokio::sync::Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>>,
) -> Result<()> {
    if record_and_check_rate(&failures, peer).await {
        anyhow::bail!(
            "too many failed registrations (max {AUTH_FAILURES_ALLOWED} per {}s)",
            AUTH_FAILURE_WINDOW.as_secs()
        );
    }

    let stream = acceptor.accept(tcp).await.context("TLS handshake failed")?;
    // Split after the handshake so the reader and the writer can be
    // supervised together by `select!` below.
    let (mut read_half, mut write_half) = tokio::io::split(stream);
    let reg = read_registration(&mut read_half).await?;
    if !verify_token(&expected_token, &reg.token) {
        warn!("Relay: rejected {:?} from {peer} (bad token)", reg.role);
        anyhow::bail!("bad viewer token");
    }
    // The registration was good; clear the peer's failure count.
    failures.lock().await.remove(&peer);
    // Accept the connection, so the client can distinguish "authorized"
    // from "TCP accepted".
    write_half
        .write_all(&[crate::network::PROTOCOL_VERSION])
        .await?;
    write_half.flush().await?;

    info!(
        "Relay: {:?} joined session '{}' from {peer}",
        reg.role, reg.session
    );

    let (tx, mut rx) = mpsc::channel::<Vec<u8>>(MAX_PEER_QUEUE_MESSAGES);
    let gen = next_gen.fetch_add(1, Ordering::Relaxed);
    let peer_entry = Peer {
        gen,
        tx: tx.clone(),
    };

    {
        // Look up and register under the lock, then let it go. Nothing
        // below this point ever holds it across an await.
        let mut map = sessions.lock().await;
        let session = map.entry(reg.session.clone()).or_insert_with(|| Session {
            host: None,
            viewers: Vec::new(),
            last_seen: Instant::now(),
        });
        session.last_seen = Instant::now();
        match reg.role {
            RelayRole::Host => {
                // A new host generation replaces the old one. The previous
                // host's writer sees its channel close and exits; its
                // cleanup can no longer clear this registration.
                if let Some(previous) = session.host.replace(peer_entry.clone()) {
                    drop(previous);
                }
            }
            RelayRole::Viewer => {
                if session.viewers.len() >= MAX_VIEWERS_PER_SESSION {
                    anyhow::bail!(
                        "session '{}' already has {} viewers (max {MAX_VIEWERS_PER_SESSION})",
                        reg.session,
                        session.viewers.len()
                    );
                }
                session.viewers.push(peer_entry.clone());
            }
        }
        if map.len() > MAX_SESSIONS {
            anyhow::bail!(
                "relay is hosting {} sessions (max {MAX_SESSIONS})",
                map.len()
            );
        }
    }
    // Everything above this point may still be holding `tx`; from here on,
    // closing the session entry is what ends this peer's writer.
    drop(tx);

    let peer_role = reg.role;
    let session_id = reg.session.clone();
    let mut peer_entry = Some(peer_entry);
    let mut queued_bytes = 0usize;

    loop {
        tokio::select! {
            // Inbound: forward to the counterpart(s).
            inbound = read_frame(&mut read_half) => {
                let Some(payload) = inbound? else { break }; // EOF
                let framed = payload;
                let len = framed.len();

                let mut targets: Vec<mpsc::Sender<Vec<u8>>> = Vec::new();
                {
                    let mut map = sessions.lock().await;
                    if let Some(session) = map.get_mut(&session_id) {
                        session.last_seen = Instant::now();
                        match peer_role {
                            RelayRole::Host => {
                                targets.extend(session.viewers.iter().map(|v| v.tx.clone()));
                            }
                            RelayRole::Viewer => {
                                if let Some(h) = &session.host {
                                    targets.push(h.tx.clone());
                                }
                            }
                        }
                    }
                }
                // No lock held. A full queue means that peer is not
                // draining; dropping the message there would silently
                // corrupt its view, so it is disconnected instead and
                // reconnects for a fresh snapshot.
                let mut congested = false;
                for tx in targets {
                    match tx.try_send(framed.clone()) {
                        Ok(()) => {}
                        Err(mpsc::error::TrySendError::Full(_)) => congested = true,
                        Err(mpsc::error::TrySendError::Closed(_)) => {}
                    }
                }
                if congested {
                    info!("Relay: disconnecting a congested viewer in session '{session_id}'");
                    break;
                }
                queued_bytes += len;
                if queued_bytes > MAX_PEER_QUEUE_BYTES {
                    warn!("Relay: peer {peer} queued {queued_bytes} bytes without draining; dropping it");
                    break;
                }
            }
            // Outbound: write to the socket.
            outbound = rx.recv() => {
                let Some(bytes) = outbound else { break };
                queued_bytes = queued_bytes.saturating_sub(bytes.len());
                if write_half.write_all(&bytes).await.is_err() {
                    break;
                }
            }
        }
    }

    // Unregister exactly the generation we created.
    {
        let mut map = sessions.lock().await;
        if let Some(session) = map.get_mut(&session_id) {
            if let Some(entry) = peer_entry.as_mut() {
                match peer_role {
                    RelayRole::Host => {
                        if session.host.as_ref().is_some_and(|h| h.gen == entry.gen) {
                            session.host = None;
                        }
                    }
                    RelayRole::Viewer => {
                        session.viewers.retain(|v| v.gen != entry.gen);
                    }
                }
            }
            if session.host.is_none() && session.viewers.is_empty() {
                map.remove(&session_id);
            }
        }
    }
    info!("Relay: {:?} left session '{session_id}'", peer_role);
    Ok(())
}

/// One inbound frame, already length-checked, or `None` at EOF.
async fn read_frame<R: tokio::io::AsyncRead + Unpin>(reader: &mut R) -> Result<Option<Vec<u8>>> {
    let mut len_buf = [0u8; 4];
    match reader.read_exact(&mut len_buf).await {
        Ok(_) => {}
        Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
        Err(e) => return Err(e.into()),
    }
    // The cap is applied here, before the allocation it authorises.
    let len = read_len_prefix(&len_buf)?;
    let mut payload = vec![0u8; len];
    reader.read_exact(&mut payload).await?;
    let mut framed = Vec::with_capacity(4 + len);
    framed.extend_from_slice(&len_buf);
    framed.extend_from_slice(&payload);
    Ok(Some(framed))
}

async fn read_registration<R: tokio::io::AsyncRead + Unpin>(
    reader: &mut R,
) -> Result<RelayRegister> {
    let mut len_buf = [0u8; 4];
    reader.read_exact(&mut len_buf).await?;
    let len = u32::from_le_bytes(len_buf) as usize;
    if len == 0 || len > MAX_REGISTRATION_BYTES {
        anyhow::bail!("relay registration size {len} is outside 1..={MAX_REGISTRATION_BYTES}");
    }
    let mut buf = vec![0u8; len];
    reader.read_exact(&mut buf).await?;
    Ok(bincode::deserialize(&buf)?)
}

/// Per-address failure counter. Returns true when the peer is over budget.
async fn record_and_check_rate(
    failures: &Arc<Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>>,
    peer: std::net::SocketAddr,
) -> bool {
    let mut map = failures.lock().await;
    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::*;

    #[test]
    fn session_codes_are_six_unambiguous_symbols() {
        let code = generate_session_code();
        assert_eq!(code.len(), 6);
        assert!(code
            .chars()
            .all(|c| "ABCDEFGHJKLMNPQRSTUVWXYZ23456789".contains(c)));
    }

    #[tokio::test]
    async fn a_hostile_registration_length_is_refused_without_allocating() {
        let mut bytes = Vec::new();
        bytes.extend_from_slice(&u32::MAX.to_le_bytes());
        let mut cursor: &[u8] = &bytes;
        let err = read_registration(&mut cursor)
            .await
            .unwrap_err()
            .to_string();
        assert!(
            err.contains("max_message_size") || err.contains("outside"),
            "unhelpful: {err}"
        );
    }

    #[tokio::test]
    async fn a_hostile_frame_length_is_refused_without_allocating() {
        let mut bytes = Vec::new();
        bytes.extend_from_slice(&u32::MAX.to_le_bytes());
        let mut cursor: &[u8] = &bytes;
        let err = read_frame(&mut cursor).await.unwrap_err().to_string();
        assert!(err.contains("max_message_size"), "unhelpful: {err}");
    }

    #[tokio::test]
    async fn rate_limiting_trips_after_the_configured_failures() {
        let failures: Arc<Mutex<HashMap<std::net::SocketAddr, (u32, Instant)>>> =
            Arc::new(Mutex::new(HashMap::new()));
        let peer: std::net::SocketAddr = "10.0.0.1:5000".parse().unwrap();
        for i in 0..AUTH_FAILURES_ALLOWED {
            assert!(
                !record_and_check_rate(&failures, peer).await,
                "allowed {i} of {AUTH_FAILURES_ALLOWED} attempts"
            );
        }
        assert!(
            record_and_check_rate(&failures, peer).await,
            "the limit must trip"
        );
    }
}