Skip to main content

vnc_worker/
lib.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! A worker for running a VNC server.
5
6#![forbid(unsafe_code)]
7
8use anyhow::Context;
9use anyhow::anyhow;
10use futures::FutureExt;
11use futures::StreamExt;
12use input_core::InputData;
13use input_core::KeyboardData;
14use input_core::MouseData;
15use mesh::message::MeshField;
16use mesh_worker::Worker;
17use mesh_worker::WorkerId;
18use mesh_worker::WorkerRpc;
19use pal_async::local::LocalDriver;
20use pal_async::local::block_with_io;
21use pal_async::socket::Listener;
22use pal_async::socket::PolledSocket;
23use pal_async::timer::PolledTimer;
24use parking_lot::Mutex;
25use std::future::Future;
26use std::net::TcpListener;
27use std::pin::Pin;
28use std::sync::Arc;
29use std::sync::atomic::AtomicBool;
30use std::sync::atomic::Ordering;
31use std::time::Duration;
32use tracing_helpers::AnyhowValueExt;
33use vnc_worker_defs::VncParameters;
34
35/// A worker for running a VNC server.
36pub struct VncWorker<T: Listener> {
37    listener: T,
38    view: ViewWrapper,
39    input_send: mesh::Sender<InputData>,
40    dirty_recv: Option<mesh::Receiver<Vec<video_core::DirtyRect>>>,
41    max_clients: usize,
42    evict_oldest: bool,
43}
44
45impl Worker for VncWorker<TcpListener> {
46    type Parameters = VncParameters<TcpListener>;
47    type State = VncParameters<TcpListener>;
48    const ID: WorkerId<Self::Parameters> = vnc_worker_defs::VNC_WORKER_TCP;
49
50    fn new(params: Self::Parameters) -> anyhow::Result<Self> {
51        Self::new_inner(params)
52    }
53
54    fn restart(state: Self::State) -> anyhow::Result<Self> {
55        Self::new(state)
56    }
57
58    fn run(self, rpc_recv: mesh::Receiver<WorkerRpc<Self::State>>) -> anyhow::Result<()> {
59        self.run_inner(rpc_recv)
60    }
61}
62
63#[cfg(any(windows, target_os = "linux"))]
64impl Worker for VncWorker<vmsocket::VmListener> {
65    type Parameters = VncParameters<vmsocket::VmListener>;
66    type State = VncParameters<vmsocket::VmListener>;
67    const ID: WorkerId<Self::Parameters> = vnc_worker_defs::VNC_WORKER_VMSOCKET;
68
69    fn new(params: Self::Parameters) -> anyhow::Result<Self> {
70        Self::new_inner(params)
71    }
72
73    fn restart(state: Self::State) -> anyhow::Result<Self> {
74        Self::new(state)
75    }
76
77    fn run(self, rpc_recv: mesh::Receiver<WorkerRpc<Self::State>>) -> anyhow::Result<()> {
78        self.run_inner(rpc_recv)
79    }
80}
81
82impl<T: 'static + Listener + MeshField + Send> VncWorker<T> {
83    fn new_inner(params: VncParameters<T>) -> anyhow::Result<Self> {
84        Ok(Self {
85            listener: params.listener,
86            view: ViewWrapper(
87                params
88                    .framebuffer
89                    .view()
90                    .context("failed to map framebuffer")?,
91            ),
92            input_send: params.input_send,
93            dirty_recv: params.dirty_recv,
94            max_clients: params.max_clients,
95            evict_oldest: params.evict_oldest,
96        })
97    }
98
99    fn run_inner(
100        self,
101        mut rpc_recv: mesh::Receiver<WorkerRpc<VncParameters<T>>>,
102    ) -> anyhow::Result<()> {
103        block_with_io(async |driver| {
104            tracing::info!(
105                address = ?self.listener.local_addr().unwrap(),
106                "VNC server listening",
107            );
108
109            let listener = PolledSocket::new(&driver, self.listener)?;
110            let mut server = MultiClientServer {
111                listener,
112                view: Arc::new(Mutex::new(self.view)),
113                input_send: self.input_send,
114                dirty_recv: self.dirty_recv,
115                dirty_senders: Vec::new(),
116                clients: unicycle::FuturesUnordered::new(),
117                abort_senders: Vec::new(),
118                next_client_id: 0,
119                max_clients: self.max_clients,
120                evict_oldest: self.evict_oldest,
121            };
122
123            let rpc = loop {
124                let r = futures::select! { // merge semantics
125                    r = rpc_recv.recv().fuse() => r,
126                    r = server.process(&driver).fuse() => break r.map(|_| None)?,
127                };
128                match r {
129                    Ok(message) => match message {
130                        WorkerRpc::Stop => break None,
131                        WorkerRpc::Inspect(deferred) => deferred.inspect(&server),
132                        WorkerRpc::Restart(response) => break Some(response),
133                    },
134                    Err(_) => break None,
135                }
136            };
137            if let Some(rpc) = rpc {
138                // Abort all active clients before recovering shared state.
139                server.abort_all_clients().await;
140                let view = Arc::try_unwrap(server.view)
141                    .expect("all clients terminated")
142                    .into_inner();
143                let state = VncParameters {
144                    listener: server.listener.into_inner(),
145                    framebuffer: view.0.access(),
146                    input_send: server.input_send,
147                    dirty_recv: server.dirty_recv,
148                    max_clients: server.max_clients,
149                    evict_oldest: server.evict_oldest,
150                };
151                rpc.complete(Ok(state));
152            }
153            Ok(())
154        })
155    }
156}
157
158/// Coordinator-side handle for one connected VNC client's dirty-rect broadcast
159/// channel. `try_send` on `sender` is non-blocking; if the channel is full the
160/// coordinator sets `missed_dirty` so the client knows to do a full refresh
161/// once it catches up.
162struct ClientDirtySender {
163    id: u64,
164    sender: async_channel::Sender<Arc<Vec<video_core::DirtyRect>>>,
165    missed_dirty: Arc<AtomicBool>,
166}
167
168/// A multi-client VNC server that accepts and manages concurrent connections.
169struct MultiClientServer<T: Listener> {
170    listener: PolledSocket<T>,
171    /// Shared framebuffer view, protected by a mutex since reads mutate
172    /// internal state (channel polling in `resolution()`).
173    view: Arc<Mutex<ViewWrapper>>,
174    /// Cloneable input sender -- each client gets its own clone.
175    input_send: mesh::Sender<InputData>,
176    /// Dirty rectangles from the synthetic video device. None if no video
177    /// device is configured.
178    dirty_recv: Option<mesh::Receiver<Vec<video_core::DirtyRect>>>,
179    /// Per-client dirty rect senders. The coordinator broadcasts device rects
180    /// to all clients via these channels. Each client accumulates rects in
181    /// its own pending bitmap -- no shared state, no clearing issues.
182    dirty_senders: Vec<ClientDirtySender>,
183    /// Futures for all active client connections. Each resolves to the
184    /// client's id when the connection ends.
185    clients: unicycle::FuturesUnordered<Pin<Box<dyn Future<Output = u64>>>>,
186    /// Abort senders for each client, keyed by client id. Dropping the
187    /// sender closes the oneshot channel, which the client detects as
188    /// an abort signal.
189    abort_senders: Vec<(u64, mesh::OneshotSender<()>)>,
190    next_client_id: u64,
191    /// Maximum concurrent clients. Bounds memory (~8MB per client for
192    /// framebuffer buffers) and prevents VRAM mutex contention.
193    max_clients: usize,
194    /// When true, evict the oldest client instead of rejecting new ones.
195    evict_oldest: bool,
196}
197
198impl<T: Listener> MultiClientServer<T> {
199    /// Main loop: accept new clients, reap finished ones, and broadcast
200    /// device dirty rects to per-client channels.
201    async fn process(&mut self, driver: &LocalDriver) -> anyhow::Result<()> {
202        enum Event<A> {
203            Accepted(A),
204            ClientDone(u64),
205            DirtyRects(Vec<video_core::DirtyRect>),
206        }
207
208        let mut device_dirty_seen = false;
209
210        loop {
211            let listener = &mut self.listener;
212            let clients = &mut self.clients;
213            let dirty_recv = &mut self.dirty_recv;
214
215            // Optional future for dirty rect reception (pending if no video device).
216            let dirty_fut = async {
217                match dirty_recv {
218                    Some(recv) => recv.recv().await,
219                    None => std::future::pending().await,
220                }
221            };
222
223            // Optional future for client completion (pending if no clients).
224            // A separate future avoids duplicating the entire select! block.
225            let client_done = async {
226                if clients.is_empty() {
227                    std::future::pending().await
228                } else {
229                    clients.select_next_some().await
230                }
231            };
232
233            let event = futures::select! {
234                accept = listener.accept().fuse() => {
235                    let (socket, addr) = accept?;
236                    Event::Accepted((socket, addr))
237                }
238                id = client_done.fuse() => Event::ClientDone(id),
239                msg = dirty_fut.fuse() => match msg {
240                    Ok(rects) => Event::DirtyRects(rects),
241                    Err(_) => {
242                        // Upstream dirty channel closed (video device reset
243                        // or teardown). Drop it so clients fall back to tile
244                        // diff instead of freezing with device_dirty_seen.
245                        tracing::warn!("device dirty channel closed, falling back to tile diff");
246                        self.dirty_recv = None;
247                        // Close all per-client dirty senders so clients
248                        // detect the closure and reset device_dirty_seen.
249                        self.dirty_senders.clear();
250                        continue;
251                    }
252                }
253            };
254
255            match event {
256                Event::Accepted((socket, remote_addr)) => {
257                    // Use abort_senders.len() as the active client count,
258                    // not clients.len(). Evicted clients are removed from
259                    // abort_senders immediately but their futures may linger
260                    // in self.clients until the next poll reaps them via
261                    // ClientDone. This means self.clients.len() can transiently
262                    // exceed max_clients during rapid connection churn (e.g.,
263                    // max_clients=1 with A→B→C arriving before A is reaped).
264                    // This is acceptable: the dying futures are in their abort
265                    // cleanup path and resolve within one poll cycle. Only
266                    // abort_senders.len() clients are actively running the VNC
267                    // protocol. Awaiting the evicted client before spawning
268                    // the replacement would add latency to every new connection.
269                    if self.abort_senders.len() >= self.max_clients {
270                        if self.evict_oldest && !self.abort_senders.is_empty() {
271                            // Disconnect the oldest client to make room.
272                            let (oldest_id, abort) = self.abort_senders.remove(0);
273                            tracing::info!(
274                                id = oldest_id,
275                                addr = ?remote_addr,
276                                "evicting oldest VNC client for new connection"
277                            );
278                            abort.send(());
279                            self.dirty_senders.retain(|s| s.id != oldest_id);
280                        } else {
281                            // Drop the socket to close the connection immediately.
282                            tracing::warn!(
283                                addr = ?remote_addr,
284                                max = self.max_clients,
285                                "VNC client rejected, limit reached"
286                            );
287                            continue;
288                        }
289                    }
290                    let sock: socket2::Socket = socket.into();
291                    let _ = sock.set_tcp_nodelay(true);
292                    match PolledSocket::new(driver, sock) {
293                        Ok(socket) => self.spawn_client(driver, socket, remote_addr),
294                        Err(e) => {
295                            tracing::error!(
296                                error = %e,
297                                "failed to register VNC client socket, dropping connection"
298                            );
299                        }
300                    }
301                }
302                Event::ClientDone(id) => {
303                    self.abort_senders.retain(|(cid, _)| *cid != id);
304                    self.dirty_senders.retain(|s| s.id != id);
305                    tracing::info!(id, count = self.clients.len(), "VNC client disconnected");
306                }
307                Event::DirtyRects(rects) => {
308                    if !device_dirty_seen {
309                        device_dirty_seen = true;
310                        tracing::info!("device dirty rects active, preferring over tile diff");
311                    }
312                    // Broadcast to all connected clients. Arc avoids cloning
313                    // the rect Vec for each client (only ref-count bump).
314                    let rects = Arc::new(rects);
315                    for s in &mut self.dirty_senders {
316                        if s.sender.try_send(Arc::clone(&rects)).is_err() {
317                            s.missed_dirty.store(true, Ordering::Relaxed);
318                            tracing::debug!(
319                                id = s.id,
320                                "client dirty channel full, flagged for full refresh"
321                            );
322                        }
323                    }
324                    tracing::trace!(
325                        rect_count = rects.len(),
326                        clients = self.dirty_senders.len(),
327                        "broadcast device dirty rects"
328                    );
329                }
330            }
331        }
332    }
333
334    /// Creates a new client connection future and adds it to the active set.
335    fn spawn_client(
336        &mut self,
337        driver: &LocalDriver,
338        socket: PolledSocket<socket2::Socket>,
339        remote_addr: impl std::fmt::Debug,
340    ) {
341        let id = self.next_client_id;
342        self.next_client_id += 1;
343        let addr_str = format!("{:?}", remote_addr);
344
345        tracing::info!(
346            id,
347            addr = %addr_str,
348            count = self.clients.len() + 1,
349            "VNC client connected",
350        );
351
352        let view = self.view.clone();
353        let input_send = self.input_send.clone();
354        let (abort_send, abort_recv) = mesh::oneshot();
355        // Per-client channel for receiving device dirty rects from the coordinator.
356        // Capacity 4: enough to buffer a few batches without blocking the coordinator.
357        // Bounded so a slow client can't unboundedly buffer broadcast batches —
358        // when full, the coordinator sets `missed_dirty` and the client falls
359        // back to a full refresh on the next pass.
360        let (dirty_send, dirty_recv) = async_channel::bounded::<Arc<Vec<video_core::DirtyRect>>>(4);
361        let missed_dirty = Arc::new(AtomicBool::new(false));
362        self.dirty_senders.push(ClientDirtySender {
363            id,
364            sender: dirty_send,
365            missed_dirty: missed_dirty.clone(),
366        });
367
368        // Each client gets its own VNC server instance with independent
369        // zlib state and pixel format, sharing only the framebuffer and
370        // input channel. The first frame is always a full screen refresh
371        // (force_full_update=true in run_internal), so new clients don't
372        // need to receive prior device dirty rects.
373        let driver = driver.clone();
374        let client_future = Box::pin(async move {
375            let fb = SharedView(view);
376            let input = SharedInput(input_send);
377            let mut vncserver = vnc::Server::new(
378                "OpenVMM VM".into(),
379                socket,
380                fb,
381                input,
382                Some(dirty_recv),
383                Some(missed_dirty),
384            );
385            let mut updater = vncserver.updater();
386
387            let mut timer = PolledTimer::new(&driver);
388            let update_task = async {
389                loop {
390                    timer.sleep(Duration::from_millis(30)).await;
391                    updater.update();
392                }
393            };
394
395            let r = futures::select! { // race semantics
396                r = vncserver.run().fuse() => r.context("VNC error"),
397                _ = abort_recv.fuse() => Err(anyhow!("VNC connection aborted")),
398                _ = update_task.fuse() => unreachable!(),
399            };
400            match r {
401                Ok(_) => {}
402                Err(err) => tracing::error!(error = err.as_error(), id, "VNC client error"),
403            }
404            id
405        });
406
407        // Store the abort sender separately so we can drop it to cancel
408        // the client on shutdown without needing to drive the future.
409        self.abort_senders.push((id, abort_send));
410        self.clients.push(client_future);
411    }
412
413    /// Aborts all active clients and waits for them to finish.
414    async fn abort_all_clients(&mut self) {
415        // Drop all abort senders, which closes their oneshot channels and
416        // causes each client's abort_recv to resolve.
417        self.abort_senders.clear();
418        self.dirty_senders.clear();
419        // Drive all client futures to completion so they can clean up.
420        while self.clients.next().await.is_some() {}
421    }
422}
423
424impl<T: Listener> inspect::Inspect for MultiClientServer<T> {
425    fn inspect(&self, req: inspect::Request<'_>) {
426        let mut resp = req.respond();
427        resp.display_debug("local_addr", &self.listener.get().local_addr().unwrap());
428        resp.field("client_count", self.clients.len());
429        resp.field("has_dirty_recv", self.dirty_recv.is_some());
430    }
431}
432
433/// Wrapper around `mesh::Sender<InputData>` that implements `vnc::Input`.
434///
435/// Each client gets its own clone; input from any client goes to the same VM.
436struct SharedInput(mesh::Sender<InputData>);
437
438impl vnc::Input for SharedInput {
439    fn key(&mut self, scancode: u16, is_down: bool) {
440        self.0.send(InputData::Keyboard(KeyboardData {
441            code: scancode,
442            make: is_down,
443        }));
444    }
445
446    fn mouse(&mut self, button_mask: u8, x: u16, y: u16) {
447        self.0
448            .send(InputData::Mouse(MouseData { button_mask, x, y }));
449    }
450}
451
452/// Wrapper around `Arc<Mutex<ViewWrapper>>` that implements `vnc::Framebuffer`.
453///
454/// The mutex is needed because `View::resolution()` mutates internal state
455/// (drains a channel). Lock durations are trivially short (memory reads).
456struct SharedView(Arc<Mutex<ViewWrapper>>);
457
458impl vnc::Framebuffer for SharedView {
459    fn read_line(&mut self, line: u16, data: &mut [u8]) {
460        self.0.lock().0.read_line(line, data)
461    }
462
463    fn resolution(&mut self) -> (u16, u16) {
464        self.0.lock().0.resolution()
465    }
466}
467
468#[derive(Debug)]
469struct ViewWrapper(framebuffer::View);
470
471#[cfg(test)]
472mod tests {
473    use super::*;
474    use framebuffer::FRAMEBUFFER_SIZE;
475    use futures::FutureExt;
476    use input_core::InputData;
477    use sparse_mmap::SparseMapping;
478    use sparse_mmap::alloc_shared_memory;
479    use std::io::Read;
480    use std::io::Write;
481    use std::net::SocketAddr;
482    use std::net::TcpStream;
483    use std::thread;
484    use std::thread::JoinHandle;
485    use std::time::Duration;
486    use video_core::DirtyRect;
487    use video_core::FramebufferFormat;
488
489    const ENCODING_TYPE_RAW: u32 = 0;
490    const ENCODING_TYPE_DESKTOP_SIZE: u32 = -223i32 as u32;
491
492    #[derive(Debug)]
493    struct UpdateRect {
494        x: u16,
495        y: u16,
496        width: u16,
497        height: u16,
498        encoding: u32,
499        payload: Vec<u8>,
500    }
501
502    struct Client {
503        stream: TcpStream,
504        width: u16,
505        height: u16,
506    }
507
508    impl Client {
509        fn connect(addr: SocketAddr) -> Self {
510            let mut stream = TcpStream::connect(addr).unwrap();
511            stream.set_nodelay(true).unwrap();
512            stream
513                .set_read_timeout(Some(Duration::from_secs(5)))
514                .unwrap();
515            stream
516                .set_write_timeout(Some(Duration::from_secs(5)))
517                .unwrap();
518
519            let mut version = [0; 12];
520            stream.read_exact(&mut version).unwrap();
521            assert_eq!(&version, b"RFB 003.008\n");
522            stream.write_all(b"RFB 003.008\n").unwrap();
523
524            let mut sec_count = [0; 1];
525            stream.read_exact(&mut sec_count).unwrap();
526            assert_eq!(sec_count, [1]);
527            let mut sec_types = vec![0; sec_count[0] as usize];
528            stream.read_exact(&mut sec_types).unwrap();
529            assert_eq!(sec_types, [1]);
530            stream.write_all(&[1]).unwrap();
531
532            let mut sec_result = [0; 4];
533            stream.read_exact(&mut sec_result).unwrap();
534            assert_eq!(sec_result, [0; 4]);
535            stream.write_all(&[1]).unwrap();
536
537            let mut init = [0; 24];
538            stream.read_exact(&mut init).unwrap();
539            let width = u16::from_be_bytes([init[0], init[1]]);
540            let height = u16::from_be_bytes([init[2], init[3]]);
541            let name_len = u32::from_be_bytes([init[20], init[21], init[22], init[23]]) as usize;
542            let mut name = vec![0; name_len];
543            stream.read_exact(&mut name).unwrap();
544            assert_eq!(name, b"OpenVMM VM");
545
546            Self {
547                stream,
548                width,
549                height,
550            }
551        }
552
553        fn send_update_request(&mut self, incremental: bool) {
554            let mut request = [0; 10];
555            request[0] = 3;
556            request[1] = u8::from(incremental);
557            request[6..8].copy_from_slice(&self.width.to_be_bytes());
558            request[8..10].copy_from_slice(&self.height.to_be_bytes());
559            self.stream.write_all(&request).unwrap();
560        }
561
562        fn send_pointer_event(&mut self, button_mask: u8, x: u16, y: u16) {
563            let mut request = [0; 6];
564            request[0] = 5;
565            request[1] = button_mask;
566            request[2..4].copy_from_slice(&x.to_be_bytes());
567            request[4..6].copy_from_slice(&y.to_be_bytes());
568            self.stream.write_all(&request).unwrap();
569        }
570
571        fn send_set_encodings(&mut self, encodings: &[u32]) {
572            let mut request = [0; 4];
573            request[0] = 2;
574            request[2..4].copy_from_slice(&(encodings.len() as u16).to_be_bytes());
575            self.stream.write_all(&request).unwrap();
576            for &encoding in encodings {
577                self.stream.write_all(&encoding.to_be_bytes()).unwrap();
578            }
579        }
580
581        fn read_update(&mut self) -> Vec<UpdateRect> {
582            self.try_read_update(Duration::from_secs(5))
583                .expect("timed out waiting for framebuffer update")
584        }
585
586        fn try_read_update(&mut self, timeout: Duration) -> Option<Vec<UpdateRect>> {
587            self.stream.set_read_timeout(Some(timeout)).unwrap();
588            let mut header = [0; 4];
589            match self.stream.read_exact(&mut header) {
590                Ok(()) => {}
591                Err(err)
592                    if matches!(
593                        err.kind(),
594                        std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
595                    ) =>
596                {
597                    self.stream
598                        .set_read_timeout(Some(Duration::from_secs(5)))
599                        .unwrap();
600                    return None;
601                }
602                Err(err) => panic!("failed to read framebuffer update header: {err}"),
603            }
604            assert_eq!(header[0], 0);
605            let rect_count = u16::from_be_bytes([header[2], header[3]]) as usize;
606            let mut rects = Vec::with_capacity(rect_count);
607            for _ in 0..rect_count {
608                let mut rect = [0; 12];
609                self.stream.read_exact(&mut rect).unwrap();
610                let x = u16::from_be_bytes([rect[0], rect[1]]);
611                let y = u16::from_be_bytes([rect[2], rect[3]]);
612                let width = u16::from_be_bytes([rect[4], rect[5]]);
613                let height = u16::from_be_bytes([rect[6], rect[7]]);
614                let encoding = u32::from_be_bytes([rect[8], rect[9], rect[10], rect[11]]);
615                let payload_len = match encoding {
616                    ENCODING_TYPE_RAW => width as usize * height as usize * 4,
617                    ENCODING_TYPE_DESKTOP_SIZE => {
618                        self.width = width;
619                        self.height = height;
620                        0
621                    }
622                    other => panic!("unsupported test encoding {other:#x}"),
623                };
624                let mut payload = vec![0; payload_len];
625                self.stream.read_exact(&mut payload).unwrap();
626                rects.push(UpdateRect {
627                    x,
628                    y,
629                    width,
630                    height,
631                    encoding,
632                    payload,
633                });
634            }
635            self.stream
636                .set_read_timeout(Some(Duration::from_secs(5)))
637                .unwrap();
638            Some(rects)
639        }
640    }
641
642    struct WorkerServer {
643        addr: SocketAddr,
644        vram: SparseMapping,
645        format_send: mesh::Sender<FramebufferFormat>,
646        input_recv: mesh::Receiver<InputData>,
647        dirty_send: Option<mesh::Sender<Vec<DirtyRect>>>,
648        stop_send: Option<mesh::OneshotSender<()>>,
649        join: Option<JoinHandle<anyhow::Result<()>>>,
650    }
651
652    impl WorkerServer {
653        fn stop(mut self) {
654            if let Some(stop) = self.stop_send.take() {
655                stop.send(());
656            }
657            if let Some(join) = self.join.take() {
658                let result = join.join().unwrap();
659                assert!(result.is_ok(), "{result:?}");
660            }
661        }
662    }
663
664    fn pixel(r: u8, g: u8, b: u8) -> [u8; 4] {
665        ((r as u32) << 16 | (g as u32) << 8 | b as u32).to_le_bytes()
666    }
667
668    fn start_server(width: u16, height: u16, with_dirty: bool) -> WorkerServer {
669        start_server_with_options(width, height, with_dirty, 16, false)
670    }
671
672    fn start_server_with_max_clients(
673        width: u16,
674        height: u16,
675        with_dirty: bool,
676        max_clients: usize,
677    ) -> WorkerServer {
678        start_server_with_options(width, height, with_dirty, max_clients, false)
679    }
680
681    fn start_server_with_options(
682        width: u16,
683        height: u16,
684        with_dirty: bool,
685        max_clients: usize,
686        evict_oldest: bool,
687    ) -> WorkerServer {
688        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
689        let addr = listener.local_addr().unwrap();
690
691        let vram = alloc_shared_memory(FRAMEBUFFER_SIZE, "vnc-worker-test").unwrap();
692        let writer_vram = vram.try_clone().unwrap();
693        let (fb, access) = framebuffer::framebuffer(vram, FRAMEBUFFER_SIZE, 0).unwrap();
694        let format_send = fb.format_send();
695        format_send.send(FramebufferFormat {
696            width: width as usize,
697            height: height as usize,
698            bytes_per_line: width as usize * 4,
699            offset: 0,
700        });
701
702        let mapping = SparseMapping::new(FRAMEBUFFER_SIZE).unwrap();
703        mapping
704            .map_file(0, FRAMEBUFFER_SIZE, &writer_vram, 0, true)
705            .unwrap();
706
707        let (input_send, input_recv) = mesh::channel();
708        let (dirty_send, dirty_recv) = if with_dirty {
709            let (send, recv) = mesh::channel();
710            (Some(send), Some(recv))
711        } else {
712            (None, None)
713        };
714        let (stop_send, stop_recv) = mesh::oneshot();
715
716        let join = thread::spawn(move || {
717            block_with_io(async |driver| -> anyhow::Result<()> {
718                let mut server = MultiClientServer {
719                    listener: PolledSocket::new(&driver, listener)?,
720                    view: Arc::new(Mutex::new(ViewWrapper(access.view().unwrap()))),
721                    input_send,
722                    dirty_recv,
723                    dirty_senders: Vec::new(),
724                    clients: unicycle::FuturesUnordered::new(),
725                    abort_senders: Vec::new(),
726                    next_client_id: 0,
727                    max_clients,
728                    evict_oldest,
729                };
730
731                futures::select! {
732                    result = server.process(&driver).fuse() => result,
733                    _ = stop_recv.fuse() => Ok(()),
734                }
735            })
736        });
737
738        WorkerServer {
739            addr,
740            vram: mapping,
741            format_send,
742            input_recv,
743            dirty_send,
744            stop_send: Some(stop_send),
745            join: Some(join),
746        }
747    }
748
749    fn wait_for_input(recv: &mut mesh::Receiver<InputData>) -> InputData {
750        for _ in 0..100 {
751            if let Ok(data) = recv.try_recv() {
752                return data;
753            }
754            thread::sleep(Duration::from_millis(10));
755        }
756        panic!("timed out waiting for input");
757    }
758
759    #[test]
760    fn e2e_multiclient_broadcasts_updates_and_forwards_input() {
761        let mut server = start_server(32, 32, true);
762        let mut client1 = Client::connect(server.addr);
763        let mut client2 = Client::connect(server.addr);
764
765        client1.send_update_request(false);
766        let first = client1.read_update();
767        assert_eq!(first.len(), 1);
768        assert_eq!(first[0].width, 32);
769        assert_eq!(first[0].height, 32);
770        assert_eq!(first[0].encoding, 0);
771
772        client2.send_update_request(false);
773        let second = client2.read_update();
774        assert_eq!(second.len(), 1);
775        assert_eq!(second[0].width, 32);
776        assert_eq!(second[0].height, 32);
777        assert_eq!(second[0].encoding, 0);
778
779        let pixel_offset = (20usize * 32 + 20) * 4;
780        server
781            .vram
782            .write_at(pixel_offset, &pixel(0x12, 0x34, 0x56))
783            .unwrap();
784        server.dirty_send.as_ref().unwrap().send(vec![DirtyRect {
785            left: 16,
786            top: 16,
787            right: 32,
788            bottom: 32,
789        }]);
790
791        client1.send_update_request(true);
792        let update1 = client1.read_update();
793        assert_eq!(update1.len(), 1);
794        assert_eq!(update1[0].x, 16);
795        assert_eq!(update1[0].y, 16);
796        assert_eq!(update1[0].width, 16);
797        assert_eq!(update1[0].height, 16);
798        assert_eq!(update1[0].encoding, 0);
799        assert_eq!(update1[0].payload.len(), 16 * 16 * 4);
800
801        client2.send_update_request(true);
802        let update2 = client2.read_update();
803        assert_eq!(update2.len(), 1);
804        assert_eq!(update2[0].x, 16);
805        assert_eq!(update2[0].y, 16);
806        assert_eq!(update2[0].width, 16);
807        assert_eq!(update2[0].height, 16);
808        assert_eq!(update2[0].encoding, 0);
809        assert_eq!(update2[0].payload.len(), 16 * 16 * 4);
810
811        client1.send_pointer_event(1, 31, 31);
812        match wait_for_input(&mut server.input_recv) {
813            InputData::Mouse(MouseData {
814                button_mask: 1,
815                x: 0x7fff,
816                y: 0x7fff,
817            }) => {}
818            other => panic!("unexpected input event: {other:?}"),
819        }
820
821        drop(client1);
822        drop(client2);
823        server.stop();
824    }
825
826    #[test]
827    fn e2e_dirty_channel_close_falls_back_to_tile_diff() {
828        let mut server = start_server(32, 32, true);
829        let mut client = Client::connect(server.addr);
830
831        client.send_update_request(false);
832        let initial = client.read_update();
833        assert_eq!(initial.len(), 1);
834        assert_eq!(initial[0].width, 32);
835        assert_eq!(initial[0].height, 32);
836
837        let first_offset = (2usize * 32 + 2) * 4;
838        server
839            .vram
840            .write_at(first_offset, &pixel(0xaa, 0xbb, 0xcc))
841            .unwrap();
842        server.dirty_send.as_ref().unwrap().send(vec![DirtyRect {
843            left: 0,
844            top: 0,
845            right: 16,
846            bottom: 16,
847        }]);
848        client.send_update_request(true);
849        let dirty_update = client.read_update();
850        assert_eq!(dirty_update.len(), 1);
851        assert_eq!(dirty_update[0].x, 0);
852        assert_eq!(dirty_update[0].y, 0);
853        assert_eq!(dirty_update[0].width, 16);
854        assert_eq!(dirty_update[0].height, 16);
855
856        let second_offset = (20usize * 32 + 20) * 4;
857        drop(server.dirty_send.take());
858        server
859            .vram
860            .write_at(second_offset, &pixel(0x11, 0x22, 0x33))
861            .unwrap();
862        let mut fallback_update = None;
863        for _ in 0..20 {
864            client.send_update_request(true);
865            if let Some(update) = client.try_read_update(Duration::from_millis(50)) {
866                if update.len() == 1
867                    && update[0].encoding == ENCODING_TYPE_RAW
868                    && update[0].x == 16
869                    && update[0].y == 16
870                    && update[0].width == 16
871                    && update[0].height == 16
872                {
873                    fallback_update = Some(update);
874                    break;
875                }
876            }
877        }
878        let fallback_update = fallback_update.expect("dirty-close fallback update not observed");
879        assert_eq!(fallback_update.len(), 1);
880
881        drop(client);
882        server.stop();
883    }
884
885    #[test]
886    fn e2e_rejects_connections_over_limit() {
887        let server = start_server(1, 1, false);
888        let mut accepted = Vec::new();
889        for _ in 0..16 {
890            let mut stream = TcpStream::connect(server.addr).unwrap();
891            stream
892                .set_read_timeout(Some(Duration::from_secs(5)))
893                .unwrap();
894            let mut version = [0; 12];
895            stream.read_exact(&mut version).unwrap();
896            assert_eq!(&version, b"RFB 003.008\n");
897            accepted.push(stream);
898        }
899
900        let mut rejected = TcpStream::connect(server.addr).unwrap();
901        rejected
902            .set_read_timeout(Some(Duration::from_secs(2)))
903            .unwrap();
904        let mut version = [0; 12];
905        let read = rejected.read(&mut version);
906        assert!(matches!(read, Ok(0) | Err(_)));
907
908        drop(rejected);
909        drop(accepted);
910        server.stop();
911    }
912
913    #[test]
914    fn e2e_max_clients_1_allows_single_connection() {
915        let server = start_server_with_max_clients(4, 4, false, 1);
916        let mut client = Client::connect(server.addr);
917
918        // First client works.
919        client.send_update_request(false);
920        let update = client.read_update();
921        assert!(!update.is_empty());
922
923        // Second client is rejected.
924        let mut rejected = TcpStream::connect(server.addr).unwrap();
925        rejected
926            .set_read_timeout(Some(Duration::from_secs(2)))
927            .unwrap();
928        let mut buf = [0; 12];
929        let read = rejected.read(&mut buf);
930        assert!(matches!(read, Ok(0) | Err(_)));
931
932        drop(rejected);
933        drop(client);
934        server.stop();
935    }
936
937    #[test]
938    fn e2e_max_clients_custom_limit_accepts_up_to_limit() {
939        let limit = 3;
940        let server = start_server_with_max_clients(4, 4, false, limit);
941        let mut accepted = Vec::new();
942
943        // Connect exactly `limit` clients.
944        for _ in 0..limit {
945            let mut stream = TcpStream::connect(server.addr).unwrap();
946            stream
947                .set_read_timeout(Some(Duration::from_secs(5)))
948                .unwrap();
949            let mut version = [0; 12];
950            stream.read_exact(&mut version).unwrap();
951            assert_eq!(&version, b"RFB 003.008\n");
952            accepted.push(stream);
953        }
954
955        // The (limit+1)th client is rejected.
956        let mut rejected = TcpStream::connect(server.addr).unwrap();
957        rejected
958            .set_read_timeout(Some(Duration::from_secs(2)))
959            .unwrap();
960        let mut buf = [0; 12];
961        let read = rejected.read(&mut buf);
962        assert!(matches!(read, Ok(0) | Err(_)));
963
964        drop(rejected);
965        drop(accepted);
966        server.stop();
967    }
968
969    #[test]
970    fn e2e_max_clients_slot_freed_after_disconnect() {
971        let server = start_server_with_max_clients(4, 4, false, 1);
972        let mut client1 = Client::connect(server.addr);
973        client1.send_update_request(false);
974        let _ = client1.read_update();
975
976        // Disconnect first client.
977        drop(client1);
978        // Give the server time to reap the disconnected client.
979        thread::sleep(Duration::from_millis(100));
980
981        // A new client can now connect.
982        let mut client2 = Client::connect(server.addr);
983        client2.send_update_request(false);
984        let update = client2.read_update();
985        assert!(!update.is_empty());
986
987        drop(client2);
988        server.stop();
989    }
990
991    #[test]
992    fn e2e_evict_oldest_disconnects_first_client() {
993        let server = start_server_with_options(4, 4, false, 1, true);
994
995        // Connect client A — should work.
996        let mut client_a = Client::connect(server.addr);
997        client_a.send_update_request(false);
998        let _ = client_a.read_update();
999
1000        // Connect client B — should evict client A.
1001        let mut client_b = Client::connect(server.addr);
1002
1003        // Client A should be disconnected.
1004        client_a
1005            .stream
1006            .set_read_timeout(Some(Duration::from_secs(2)))
1007            .unwrap();
1008        let mut buf = [0; 1];
1009        let read = client_a.stream.read(&mut buf);
1010        assert!(
1011            matches!(read, Ok(0) | Err(_)),
1012            "client A should be disconnected"
1013        );
1014
1015        // Client B should work.
1016        client_b.send_update_request(false);
1017        let update = client_b.read_update();
1018        assert!(!update.is_empty());
1019
1020        drop(client_a);
1021        drop(client_b);
1022        server.stop();
1023    }
1024
1025    #[test]
1026    fn e2e_evict_oldest_false_rejects_new_client() {
1027        let server = start_server_with_options(4, 4, false, 1, false);
1028
1029        // Connect client A.
1030        let mut client_a = Client::connect(server.addr);
1031        client_a.send_update_request(false);
1032        let _ = client_a.read_update();
1033
1034        // Client B should be rejected.
1035        let mut rejected = TcpStream::connect(server.addr).unwrap();
1036        rejected
1037            .set_read_timeout(Some(Duration::from_secs(2)))
1038            .unwrap();
1039        let mut buf = [0; 12];
1040        let read = rejected.read(&mut buf);
1041        assert!(matches!(read, Ok(0) | Err(_)));
1042
1043        // Client A should still work — send a non-incremental request
1044        // to force a full update (incremental with no changes produces
1045        // no response).
1046        client_a.send_update_request(false);
1047        let update = client_a.read_update();
1048        assert!(!update.is_empty());
1049
1050        drop(rejected);
1051        drop(client_a);
1052        server.stop();
1053    }
1054
1055    #[test]
1056    fn e2e_evict_oldest_with_multiple_clients() {
1057        let server = start_server_with_options(4, 4, false, 2, true);
1058
1059        // Connect A and B.
1060        let mut client_a = Client::connect(server.addr);
1061        client_a.send_update_request(false);
1062        let _ = client_a.read_update();
1063
1064        let mut client_b = Client::connect(server.addr);
1065        client_b.send_update_request(false);
1066        let _ = client_b.read_update();
1067
1068        // Connect C — should evict A (the oldest).
1069        let mut client_c = Client::connect(server.addr);
1070
1071        // Client A should be disconnected.
1072        client_a
1073            .stream
1074            .set_read_timeout(Some(Duration::from_secs(2)))
1075            .unwrap();
1076        let mut buf = [0; 1];
1077        let read = client_a.stream.read(&mut buf);
1078        assert!(matches!(read, Ok(0) | Err(_)), "client A should be evicted");
1079
1080        // Client B should still work (non-incremental to force response).
1081        client_b.send_update_request(false);
1082        let update_b = client_b.read_update();
1083        assert!(!update_b.is_empty());
1084
1085        // Client C should work.
1086        client_c.send_update_request(false);
1087        let update_c = client_c.read_update();
1088        assert!(!update_c.is_empty());
1089
1090        drop(client_a);
1091        drop(client_b);
1092        drop(client_c);
1093        server.stop();
1094    }
1095
1096    #[test]
1097    fn e2e_resize_parses_desktop_size_before_full_refresh() {
1098        let server = start_server(4, 2, false);
1099        let mut client = Client::connect(server.addr);
1100
1101        client.send_set_encodings(&[ENCODING_TYPE_DESKTOP_SIZE]);
1102        client.send_update_request(false);
1103        let initial = client.read_update();
1104        assert_eq!(initial.len(), 1);
1105        assert_eq!(initial[0].encoding, ENCODING_TYPE_RAW);
1106
1107        server.format_send.send(FramebufferFormat {
1108            width: 6,
1109            height: 3,
1110            bytes_per_line: 6 * 4,
1111            offset: 0,
1112        });
1113        client.send_update_request(true);
1114        let resize = client.read_update();
1115        assert_eq!(resize.len(), 1);
1116        assert_eq!(resize[0].encoding, ENCODING_TYPE_DESKTOP_SIZE);
1117        assert_eq!(resize[0].width, 6);
1118        assert_eq!(resize[0].height, 3);
1119
1120        let refresh = client.read_update();
1121        assert_eq!(refresh.len(), 1);
1122        assert_eq!(refresh[0].encoding, ENCODING_TYPE_RAW);
1123        assert_eq!(refresh[0].width, 6);
1124        assert_eq!(refresh[0].height, 3);
1125
1126        drop(client);
1127        server.stop();
1128    }
1129}