use std::collections::HashMap;
use std::path::PathBuf;
use std::process::Stdio;
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
pub const OP_RAW: u8 = b'R';
pub const OP_PNG: u8 = b'P';
pub const STATUS_OK: u8 = 0;
pub const STATUS_UNAVAILABLE: u8 = 1;
#[derive(Clone, PartialEq, Eq)]
pub enum CapturedFrame {
Bgra {
width: u32,
height: u32,
data: Vec<u8>,
},
Png(Vec<u8>),
}
impl std::fmt::Debug for CapturedFrame {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CapturedFrame::Bgra {
width,
height,
data,
} => f
.debug_struct("CapturedFrame::Bgra")
.field("width", width)
.field("height", height)
.field("bytes", &data.len())
.finish(),
CapturedFrame::Png(b) => f
.debug_struct("CapturedFrame::Png")
.field("bytes", &b.len())
.finish(),
}
}
}
#[derive(Debug)]
pub enum HostError {
BinaryMissing(PathBuf),
ResolveFailed(String),
HostGone,
Io(std::io::Error),
Protocol(String),
}
impl std::fmt::Display for HostError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
HostError::BinaryMissing(p) => write!(f, "smix-capture-host not found at {p:?}"),
HostError::ResolveFailed(s) => write!(f, "capture-host surface resolve failed: {s}"),
HostError::HostGone => write!(f, "capture-host process gone"),
HostError::Io(e) => write!(f, "capture-host io: {e}"),
HostError::Protocol(s) => write!(f, "capture-host protocol: {s}"),
}
}
}
impl std::error::Error for HostError {}
impl From<std::io::Error> for HostError {
fn from(e: std::io::Error) -> Self {
HostError::Io(e)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FrameHeader {
pub width: u32,
pub height: u32,
pub len: u32,
}
impl FrameHeader {
pub fn parse(buf: &[u8; 12]) -> FrameHeader {
FrameHeader {
width: u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]]),
height: u32::from_le_bytes([buf[4], buf[5], buf[6], buf[7]]),
len: u32::from_le_bytes([buf[8], buf[9], buf[10], buf[11]]),
}
}
}
pub fn capture_host_bin() -> PathBuf {
std::env::var_os("SMIX_CAPTURE_HOST_BIN").map_or_else(
|| PathBuf::from("swift-bridge/.build/release/smix-capture-host"),
PathBuf::from,
)
}
pub struct SurfaceCaptureHost {
child: Child,
stdin: ChildStdin,
stdout: BufReader<ChildStdout>,
pub width: u32,
pub height: u32,
}
impl SurfaceCaptureHost {
pub async fn spawn(udid: &str) -> Result<SurfaceCaptureHost, HostError> {
let bin = capture_host_bin();
if !bin.exists() {
return Err(HostError::BinaryMissing(bin));
}
let mut child = Command::new(&bin)
.arg(udid)
.arg("serve")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.kill_on_drop(true)
.spawn()
.map_err(HostError::Io)?;
let stdin = child
.stdin
.take()
.ok_or_else(|| HostError::ResolveFailed("stdin not piped".into()))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| HostError::ResolveFailed("stdout not piped".into()))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| HostError::ResolveFailed("stderr not piped".into()))?;
let mut stderr_reader = BufReader::new(stderr);
let mut header = String::new();
let read =
tokio::time::timeout(Duration::from_secs(5), stderr_reader.read_line(&mut header))
.await;
let (width, height) = match read {
Ok(Ok(n)) if n > 0 => parse_geometry_line(&header)
.ok_or_else(|| HostError::ResolveFailed(format!("bad WxH header: {header:?}")))?,
Ok(Ok(_)) => {
return Err(HostError::ResolveFailed(
"host exited before WxH header".into(),
));
}
Ok(Err(e)) => return Err(HostError::ResolveFailed(format!("read header: {e}"))),
Err(_) => {
return Err(HostError::ResolveFailed(
"WxH header not received within 5s".into(),
));
}
};
tokio::spawn(async move {
let mut lines = stderr_reader.lines();
while let Ok(Some(_line)) = lines.next_line().await {}
});
Ok(SurfaceCaptureHost {
child,
stdin,
stdout: BufReader::new(stdout),
width,
height,
})
}
pub async fn grab(&mut self, want_png: bool) -> Result<Option<CapturedFrame>, HostError> {
let op = if want_png { OP_PNG } else { OP_RAW };
self.stdin.write_all(&[op]).await?;
self.stdin.flush().await?;
let mut status = [0u8; 1];
if let Err(e) = self.stdout.read_exact(&mut status).await {
return if e.kind() == std::io::ErrorKind::UnexpectedEof {
Err(HostError::HostGone)
} else {
Err(HostError::Io(e))
};
}
match status[0] {
STATUS_UNAVAILABLE => Ok(None),
STATUS_OK => {
let mut hdr = [0u8; 12];
self.read_exact_or_gone(&mut hdr).await?;
let h = FrameHeader::parse(&hdr);
let len = h.len as usize;
if len > 128 * 1024 * 1024 {
return Err(HostError::Protocol(format!("payload len too large: {len}")));
}
let mut payload = vec![0u8; len];
self.read_exact_or_gone(&mut payload).await?;
let frame = if want_png {
CapturedFrame::Png(payload)
} else {
CapturedFrame::Bgra {
width: h.width,
height: h.height,
data: payload,
}
};
Ok(Some(frame))
}
other => Err(HostError::Protocol(format!("unknown status byte {other}"))),
}
}
async fn read_exact_or_gone(&mut self, buf: &mut [u8]) -> Result<(), HostError> {
match self.stdout.read_exact(buf).await {
Ok(_) => Ok(()),
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => Err(HostError::HostGone),
Err(e) => Err(HostError::Io(e)),
}
}
pub async fn shutdown(mut self) {
drop(self.stdin);
let _ = tokio::time::timeout(Duration::from_secs(2), self.child.wait()).await;
}
}
pub fn parse_geometry_line(s: &str) -> Option<(u32, u32)> {
let (w, h) = s.trim().split_once('x')?;
Some((w.parse().ok()?, h.parse().ok()?))
}
#[derive(Default)]
pub struct CaptureHostRegistry {
hosts: HashMap<String, SurfaceCaptureHost>,
}
impl std::fmt::Debug for CaptureHostRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CaptureHostRegistry")
.field("resident_hosts", &self.hosts.len())
.finish()
}
}
impl CaptureHostRegistry {
pub fn take(&mut self, udid: &str) -> Option<SurfaceCaptureHost> {
self.hosts.remove(udid)
}
pub fn put(&mut self, udid: &str, host: SurfaceCaptureHost) {
self.hosts.insert(udid.to_string(), host);
}
pub fn evict(&mut self, udid: &str) -> Option<SurfaceCaptureHost> {
self.hosts.remove(udid)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn frame_header_parses_little_endian() {
let buf = [
0xB6, 0x04, 0x00, 0x00, 0x3E, 0x0A, 0x00, 0x00, 0x50, 0x00, 0xC1, 0x00,
];
let h = FrameHeader::parse(&buf);
assert_eq!(h.width, 1206);
assert_eq!(h.height, 2622);
assert_eq!(h.len, 12_648_528);
}
#[test]
fn status_constants_are_distinct_and_ops_are_ascii() {
assert_ne!(STATUS_OK, STATUS_UNAVAILABLE);
assert_eq!(OP_RAW, b'R');
assert_eq!(OP_PNG, b'P');
}
#[test]
fn geometry_line_parses_and_rejects_junk() {
assert_eq!(parse_geometry_line("1206x2622\n"), Some((1206, 2622)));
assert_eq!(parse_geometry_line(" 800x600 "), Some((800, 600)));
assert_eq!(parse_geometry_line("not-a-size"), None);
assert_eq!(parse_geometry_line("1206x"), None);
}
#[test]
fn registry_take_put_evict_roundtrip_key() {
let mut reg = CaptureHostRegistry::default();
assert!(reg.take("UDID-A").is_none());
assert!(reg.evict("UDID-A").is_none());
assert!(reg.hosts.is_empty());
}
#[test]
fn bin_path_honors_env_override() {
let prev = std::env::var_os("SMIX_CAPTURE_HOST_BIN");
unsafe { std::env::set_var("SMIX_CAPTURE_HOST_BIN", "/tmp/custom-host") };
assert_eq!(capture_host_bin(), PathBuf::from("/tmp/custom-host"));
unsafe {
match prev {
Some(v) => std::env::set_var("SMIX_CAPTURE_HOST_BIN", v),
None => std::env::remove_var("SMIX_CAPTURE_HOST_BIN"),
}
}
}
}