1use std::collections::HashMap;
29use std::path::PathBuf;
30use std::process::Stdio;
31use std::time::Duration;
32
33use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
34use tokio::process::{Child, ChildStdin, ChildStdout, Command};
35
36pub const OP_RAW: u8 = b'R';
38pub const OP_PNG: u8 = b'P';
40pub const STATUS_OK: u8 = 0;
42pub const STATUS_UNAVAILABLE: u8 = 1;
45
46#[derive(Clone, PartialEq, Eq)]
49pub enum CapturedFrame {
50 Bgra {
52 width: u32,
54 height: u32,
56 data: Vec<u8>,
58 },
59 Png(Vec<u8>),
61}
62
63impl std::fmt::Debug for CapturedFrame {
64 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
65 match self {
66 CapturedFrame::Bgra {
67 width,
68 height,
69 data,
70 } => f
71 .debug_struct("CapturedFrame::Bgra")
72 .field("width", width)
73 .field("height", height)
74 .field("bytes", &data.len())
75 .finish(),
76 CapturedFrame::Png(b) => f
77 .debug_struct("CapturedFrame::Png")
78 .field("bytes", &b.len())
79 .finish(),
80 }
81 }
82}
83
84#[derive(Debug)]
87pub enum HostError {
88 BinaryMissing(PathBuf),
90 ResolveFailed(String),
93 HostGone,
95 Io(std::io::Error),
97 Protocol(String),
99}
100
101impl std::fmt::Display for HostError {
102 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
103 match self {
104 HostError::BinaryMissing(p) => write!(f, "smix-capture-host not found at {p:?}"),
105 HostError::ResolveFailed(s) => write!(f, "capture-host surface resolve failed: {s}"),
106 HostError::HostGone => write!(f, "capture-host process gone"),
107 HostError::Io(e) => write!(f, "capture-host io: {e}"),
108 HostError::Protocol(s) => write!(f, "capture-host protocol: {s}"),
109 }
110 }
111}
112
113impl std::error::Error for HostError {}
114
115impl From<std::io::Error> for HostError {
116 fn from(e: std::io::Error) -> Self {
117 HostError::Io(e)
118 }
119}
120
121#[derive(Debug, Clone, Copy, PartialEq, Eq)]
123pub struct FrameHeader {
124 pub width: u32,
126 pub height: u32,
128 pub len: u32,
130}
131
132impl FrameHeader {
133 pub fn parse(buf: &[u8; 12]) -> FrameHeader {
135 FrameHeader {
136 width: u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]]),
137 height: u32::from_le_bytes([buf[4], buf[5], buf[6], buf[7]]),
138 len: u32::from_le_bytes([buf[8], buf[9], buf[10], buf[11]]),
139 }
140 }
141}
142
143pub fn capture_host_bin() -> PathBuf {
146 std::env::var_os("SMIX_CAPTURE_HOST_BIN").map_or_else(
147 || PathBuf::from("swift-bridge/.build/release/smix-capture-host"),
148 PathBuf::from,
149 )
150}
151
152pub struct SurfaceCaptureHost {
158 child: Child,
159 stdin: ChildStdin,
160 stdout: BufReader<ChildStdout>,
161 pub width: u32,
164 pub height: u32,
166}
167
168impl SurfaceCaptureHost {
169 pub async fn spawn(udid: &str) -> Result<SurfaceCaptureHost, HostError> {
172 let bin = capture_host_bin();
173 if !bin.exists() {
174 return Err(HostError::BinaryMissing(bin));
175 }
176 let mut child = Command::new(&bin)
177 .arg(udid)
178 .arg("serve")
179 .stdin(Stdio::piped())
180 .stdout(Stdio::piped())
181 .stderr(Stdio::piped())
182 .kill_on_drop(true)
183 .spawn()
184 .map_err(HostError::Io)?;
185
186 let stdin = child
187 .stdin
188 .take()
189 .ok_or_else(|| HostError::ResolveFailed("stdin not piped".into()))?;
190 let stdout = child
191 .stdout
192 .take()
193 .ok_or_else(|| HostError::ResolveFailed("stdout not piped".into()))?;
194 let stderr = child
195 .stderr
196 .take()
197 .ok_or_else(|| HostError::ResolveFailed("stderr not piped".into()))?;
198
199 let mut stderr_reader = BufReader::new(stderr);
200 let mut header = String::new();
201 let read =
202 tokio::time::timeout(Duration::from_secs(5), stderr_reader.read_line(&mut header))
203 .await;
204 let (width, height) = match read {
205 Ok(Ok(n)) if n > 0 => parse_geometry_line(&header)
206 .ok_or_else(|| HostError::ResolveFailed(format!("bad WxH header: {header:?}")))?,
207 Ok(Ok(_)) => {
208 return Err(HostError::ResolveFailed(
209 "host exited before WxH header".into(),
210 ));
211 }
212 Ok(Err(e)) => return Err(HostError::ResolveFailed(format!("read header: {e}"))),
213 Err(_) => {
214 return Err(HostError::ResolveFailed(
215 "WxH header not received within 5s".into(),
216 ));
217 }
218 };
219
220 tokio::spawn(async move {
223 let mut lines = stderr_reader.lines();
224 while let Ok(Some(_line)) = lines.next_line().await {}
225 });
226
227 Ok(SurfaceCaptureHost {
228 child,
229 stdin,
230 stdout: BufReader::new(stdout),
231 width,
232 height,
233 })
234 }
235
236 pub async fn grab(&mut self, want_png: bool) -> Result<Option<CapturedFrame>, HostError> {
243 let op = if want_png { OP_PNG } else { OP_RAW };
244 self.stdin.write_all(&[op]).await?;
245 self.stdin.flush().await?;
246
247 let mut status = [0u8; 1];
248 if let Err(e) = self.stdout.read_exact(&mut status).await {
249 return if e.kind() == std::io::ErrorKind::UnexpectedEof {
250 Err(HostError::HostGone)
251 } else {
252 Err(HostError::Io(e))
253 };
254 }
255 match status[0] {
256 STATUS_UNAVAILABLE => Ok(None),
257 STATUS_OK => {
258 let mut hdr = [0u8; 12];
259 self.read_exact_or_gone(&mut hdr).await?;
260 let h = FrameHeader::parse(&hdr);
261 let len = h.len as usize;
262 if len > 128 * 1024 * 1024 {
266 return Err(HostError::Protocol(format!("payload len too large: {len}")));
267 }
268 let mut payload = vec![0u8; len];
269 self.read_exact_or_gone(&mut payload).await?;
270 let frame = if want_png {
271 CapturedFrame::Png(payload)
272 } else {
273 CapturedFrame::Bgra {
274 width: h.width,
275 height: h.height,
276 data: payload,
277 }
278 };
279 Ok(Some(frame))
280 }
281 other => Err(HostError::Protocol(format!("unknown status byte {other}"))),
282 }
283 }
284
285 async fn read_exact_or_gone(&mut self, buf: &mut [u8]) -> Result<(), HostError> {
286 match self.stdout.read_exact(buf).await {
287 Ok(_) => Ok(()),
288 Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Err(HostError::HostGone),
289 Err(e) => Err(HostError::Io(e)),
290 }
291 }
292
293 pub async fn shutdown(mut self) {
296 drop(self.stdin);
297 let _ = tokio::time::timeout(Duration::from_secs(2), self.child.wait()).await;
298 }
299}
300
301pub fn parse_geometry_line(s: &str) -> Option<(u32, u32)> {
303 let (w, h) = s.trim().split_once('x')?;
304 Some((w.parse().ok()?, h.parse().ok()?))
305}
306
307#[derive(Default)]
310pub struct CaptureHostRegistry {
311 hosts: HashMap<String, SurfaceCaptureHost>,
312}
313
314impl std::fmt::Debug for CaptureHostRegistry {
315 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
316 f.debug_struct("CaptureHostRegistry")
317 .field("resident_hosts", &self.hosts.len())
318 .finish()
319 }
320}
321
322impl CaptureHostRegistry {
323 pub fn take(&mut self, udid: &str) -> Option<SurfaceCaptureHost> {
325 self.hosts.remove(udid)
326 }
327
328 pub fn put(&mut self, udid: &str, host: SurfaceCaptureHost) {
330 self.hosts.insert(udid.to_string(), host);
331 }
332
333 pub fn evict(&mut self, udid: &str) -> Option<SurfaceCaptureHost> {
335 self.hosts.remove(udid)
336 }
337}
338
339#[cfg(test)]
340mod tests {
341 use super::*;
342
343 #[test]
344 fn frame_header_parses_little_endian() {
345 let buf = [
347 0xB6, 0x04, 0x00, 0x00, 0x3E, 0x0A, 0x00, 0x00, 0x50, 0x00, 0xC1, 0x00,
348 ];
349 let h = FrameHeader::parse(&buf);
350 assert_eq!(h.width, 1206);
351 assert_eq!(h.height, 2622);
352 assert_eq!(h.len, 12_648_528);
353 }
354
355 #[test]
356 fn status_constants_are_distinct_and_ops_are_ascii() {
357 assert_ne!(STATUS_OK, STATUS_UNAVAILABLE);
358 assert_eq!(OP_RAW, b'R');
359 assert_eq!(OP_PNG, b'P');
360 }
361
362 #[test]
363 fn geometry_line_parses_and_rejects_junk() {
364 assert_eq!(parse_geometry_line("1206x2622\n"), Some((1206, 2622)));
365 assert_eq!(parse_geometry_line(" 800x600 "), Some((800, 600)));
366 assert_eq!(parse_geometry_line("not-a-size"), None);
367 assert_eq!(parse_geometry_line("1206x"), None);
368 }
369
370 #[test]
371 fn registry_take_put_evict_roundtrip_key() {
372 let mut reg = CaptureHostRegistry::default();
375 assert!(reg.take("UDID-A").is_none());
376 assert!(reg.evict("UDID-A").is_none());
377 assert!(reg.hosts.is_empty());
380 }
381
382 #[test]
383 fn bin_path_honors_env_override() {
384 let prev = std::env::var_os("SMIX_CAPTURE_HOST_BIN");
386 unsafe { std::env::set_var("SMIX_CAPTURE_HOST_BIN", "/tmp/custom-host") };
388 assert_eq!(capture_host_bin(), PathBuf::from("/tmp/custom-host"));
389 unsafe {
390 match prev {
391 Some(v) => std::env::set_var("SMIX_CAPTURE_HOST_BIN", v),
392 None => std::env::remove_var("SMIX_CAPTURE_HOST_BIN"),
393 }
394 }
395 }
396}