From 0655bdc71b69988967197d44bb9abaf1aa97d604 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sat, 1 Aug 2026 13:46:06 +0200 Subject: [PATCH 01/52] Centralise vsock frame proto --- src/daemon.rs | 34 ++++------------------- src/main.rs | 20 +++----------- src/vscomm/mod.rs | 69 +++++++++++++++++++++++++++++++++++++++-------- 3 files changed, 66 insertions(+), 57 deletions(-) diff --git a/src/daemon.rs b/src/daemon.rs index b763969..8a8d073 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -446,23 +446,12 @@ fn is_allowed(passthrough: &[String], command: &str, args: &[String]) -> bool { } async fn read_exec_request(reader: &mut R) -> Result { - let mut header = [0u8; 6]; - reader.read_exact(&mut header).await.map_err(|e| format!("read header: {e}"))?; - - let frame_type_raw = u16::from_le_bytes([header[0], header[1]]); - let ft = FrameType::from_u16(frame_type_raw).ok_or_else(|| format!("unknown frame type: {frame_type_raw}"))?; - - if !matches!(ft, FrameType::ExecReq) { - return Err(format!("expected ExecReq, got {:?}", ft as u16)); + let frame = Frame::read_async(reader).await.map_err(|e| format!("read frame: {e}"))?; + if !matches!(frame.frame_type, FrameType::ExecReq) { + return Err(format!("expected ExecReq, got {:?}", frame.frame_type as u16)); } - let payload_len = u32::from_le_bytes([header[2], header[3], header[4], header[5]]) as usize; - let mut payload = vec![0u8; payload_len]; - if payload_len > 0 { - reader.read_exact(&mut payload).await.map_err(|e| format!("read payload: {e}"))?; - } - - ExecRequest::deserialize(&payload) + ExecRequest::deserialize(&frame.payload) } fn bwrap_status_pipe() -> Result<(File, File), String> { @@ -546,20 +535,7 @@ async fn pump_to_channel(mut reader: R, frame_type: Fra } async fn write_frame(writer: &mut W, frame: &Frame) -> Result<(), String> { - let frame_type_raw = frame.frame_type as u16; - let payload_len = frame.payload.len() as u32; - - let mut header = [0u8; 6]; - header[0..2].copy_from_slice(&frame_type_raw.to_le_bytes()); - header[2..6].copy_from_slice(&payload_len.to_le_bytes()); - - writer.write_all(&header).await.map_err(|e| format!("write header: {e}"))?; - if !frame.payload.is_empty() { - writer.write_all(&frame.payload).await.map_err(|e| format!("write payload: {e}"))?; - } - writer.flush().await.map_err(|e| format!("flush: {e}"))?; - - Ok(()) + frame.write_async(writer).await.map_err(|e| format!("write frame: {e}")) } fn find_netrelay_binary() -> Result { diff --git a/src/main.rs b/src/main.rs index 83ecce4..784d643 100644 --- a/src/main.rs +++ b/src/main.rs @@ -435,28 +435,14 @@ async fn status_listener( let overlay = overlay.clone(); tokio::spawn(async move { - let mut header = [0u8; 6]; - if stream.read_exact(&mut header).await.is_err() { - return; - } - - let frame_type_raw = u16::from_le_bytes([header[0], header[1]]); - let payload_len = u32::from_le_bytes([header[2], header[3], header[4], header[5]]) as usize; - - let Some(ft) = vscomm::FrameType::from_u16(frame_type_raw) else { + let Ok(frame) = vscomm::Frame::read_async(&mut stream).await else { return; }; - - if !matches!(ft, vscomm::FrameType::UiCommand) { - return; - } - - let mut payload = vec![0u8; payload_len]; - if payload_len > 0 && stream.read_exact(&mut payload).await.is_err() { + if !matches!(frame.frame_type, vscomm::FrameType::UiCommand) { return; } - if let Some((widget, cmd, opts, val)) = vscomm::decode_ui_payload(&payload) { + if let Some((widget, cmd, opts, val)) = vscomm::decode_ui_payload(&frame.payload) { if widget == "error" && cmd == "show" { let title = if opts.is_empty() { "Bunkerbox error" } else { opts }; logging::diagnostic(&format!("TUI error [{title}]: {val}")); diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index 00365f6..3ee4b17 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -7,12 +7,15 @@ use std::path::Path; #[cfg(unix)] use std::os::unix::ffi::OsStrExt; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; pub mod buildsys; pub const TOOLCHAIN_PORT: u32 = 9999; // Keep UI traffic on a separate vsock endpoint from command execution. pub const TUI_STATUS_PORT: u32 = 10000; pub const VSCOMM_BIN_DIR: &str = "/usr/local/bunkerbox/bin"; +/// Maximum payload accepted in one vsock frame. +pub const MAX_FRAME_PAYLOAD: usize = 1024 * 1024; #[repr(u16)] #[derive(Clone, Copy)] @@ -119,11 +122,7 @@ impl Frame { let mut header = [0u8; 6]; reader.read_exact(&mut header)?; - let frame_type_raw = u16::from_le_bytes([header[0], header[1]]); - let frame_type = FrameType::from_u16(frame_type_raw) - .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, format!("unknown frame type: {frame_type_raw}")))?; - - let payload_len = u32::from_le_bytes([header[2], header[3], header[4], header[5]]) as usize; + let (frame_type, payload_len) = decode_header(&header)?; let mut payload = vec![0u8; payload_len]; if payload_len > 0 { @@ -133,13 +132,21 @@ impl Frame { Ok(Self { frame_type, payload }) } - pub fn write(&self, writer: &mut W) -> io::Result<()> { - let frame_type_raw = self.frame_type as u16; - let payload_len = self.payload.len() as u32; - + pub async fn read_async(reader: &mut R) -> io::Result { let mut header = [0u8; 6]; - header[0..2].copy_from_slice(&frame_type_raw.to_le_bytes()); - header[2..6].copy_from_slice(&payload_len.to_le_bytes()); + reader.read_exact(&mut header).await?; + + let (frame_type, payload_len) = decode_header(&header)?; + let mut payload = vec![0u8; payload_len]; + if payload_len > 0 { + reader.read_exact(&mut payload).await?; + } + + Ok(Self { frame_type, payload }) + } + + pub fn write(&self, writer: &mut W) -> io::Result<()> { + let header = self.header()?; writer.write_all(&header)?; if !self.payload.is_empty() { @@ -149,6 +156,46 @@ impl Frame { writer.flush()?; Ok(()) } + + pub async fn write_async(&self, writer: &mut W) -> io::Result<()> { + let header = self.header()?; + writer.write_all(&header).await?; + + if !self.payload.is_empty() { + writer.write_all(&self.payload).await?; + } + + writer.flush().await + } + + fn header(&self) -> io::Result<[u8; 6]> { + validate_payload_size(self.payload.len())?; + + let mut header = [0u8; 6]; + header[0..2].copy_from_slice(&(self.frame_type as u16).to_le_bytes()); + header[2..6].copy_from_slice(&(self.payload.len() as u32).to_le_bytes()); + Ok(header) + } +} + +fn decode_header(header: &[u8; 6]) -> io::Result<(FrameType, usize)> { + let frame_type_raw = u16::from_le_bytes([header[0], header[1]]); + let frame_type = FrameType::from_u16(frame_type_raw) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, format!("unknown frame type: {frame_type_raw}")))?; + let payload_len = u32::from_le_bytes([header[2], header[3], header[4], header[5]]) as usize; + validate_payload_size(payload_len)?; + Ok((frame_type, payload_len)) +} + +fn validate_payload_size(payload_len: usize) -> io::Result<()> { + if payload_len > MAX_FRAME_PAYLOAD { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + format!("frame payload too large: {payload_len} bytes (maximum {MAX_FRAME_PAYLOAD})"), + )); + } + + Ok(()) } impl ExecRequest { From e5dfb116bdb87e3868de943326917c544aff0e29 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sat, 1 Aug 2026 13:46:21 +0200 Subject: [PATCH 02/52] Add centralised vsock proto UT --- src/vscomm/ut.rs | 191 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 191 insertions(+) create mode 100644 src/vscomm/ut.rs diff --git a/src/vscomm/ut.rs b/src/vscomm/ut.rs new file mode 100644 index 0000000..eede4af --- /dev/null +++ b/src/vscomm/ut.rs @@ -0,0 +1,191 @@ +use super::*; +use std::io::{self, Cursor, Read}; +use std::pin::Pin; +use std::task::{Context, Poll}; +use tokio::io::{AsyncRead, ReadBuf}; + +fn encoded_frame(frame_type: u16, payload_len: u32, payload: &[u8]) -> Vec { + let mut encoded = Vec::with_capacity(6 + payload.len()); + encoded.extend_from_slice(&frame_type.to_le_bytes()); + encoded.extend_from_slice(&payload_len.to_le_bytes()); + encoded.extend_from_slice(payload); + encoded +} + +fn frame_error(result: io::Result) -> io::Error { + match result { + Ok(_) => panic!("expected frame read to fail"), + Err(error) => error, + } +} + +struct FragmentedReader { + data: Vec, + offset: usize, + chunk_size: usize, +} + +impl FragmentedReader { + fn new(data: Vec, chunk_size: usize) -> Self { + Self { data, offset: 0, chunk_size } + } +} + +impl Read for FragmentedReader { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + if self.offset == self.data.len() { + return Ok(0); + } + + let amount = self.chunk_size.min(buf.len()).min(self.data.len() - self.offset); + buf[..amount].copy_from_slice(&self.data[self.offset..self.offset + amount]); + self.offset += amount; + Ok(amount) + } +} + +impl AsyncRead for FragmentedReader { + fn poll_read(mut self: Pin<&mut Self>, _cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + if self.offset == self.data.len() { + return Poll::Ready(Ok(())); + } + + let amount = self.chunk_size.min(buf.remaining()).min(self.data.len() - self.offset); + buf.put_slice(&self.data[self.offset..self.offset + amount]); + self.offset += amount; + Poll::Ready(Ok(())) + } +} + +#[test] +fn read_zero_length_frame() { + let frame = Frame::read(&mut Cursor::new(encoded_frame(FrameType::Stdout as u16, 0, &[]))).unwrap(); + assert!(matches!(frame.frame_type, FrameType::Stdout)); + assert!(frame.payload.is_empty()); +} + +#[test] +fn read_maximum_size_frame() { + let payload = vec![0xA5; MAX_FRAME_PAYLOAD]; + let frame = Frame::read(&mut Cursor::new(encoded_frame(FrameType::Stdout as u16, MAX_FRAME_PAYLOAD as u32, &payload))).unwrap(); + assert_eq!(frame.payload, payload); +} + +#[test] +fn read_oversized_frame_before_payload_read() { + let encoded = encoded_frame(FrameType::Stdout as u16, (MAX_FRAME_PAYLOAD + 1) as u32, &[]); + let error = frame_error(Frame::read(&mut Cursor::new(encoded))); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert!(error.to_string().contains("frame payload too large")); +} + +#[test] +fn read_unknown_frame_type() { + let error = frame_error(Frame::read(&mut Cursor::new(encoded_frame(99, 0, &[])))); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert!(error.to_string().contains("unknown frame type")); +} + +#[test] +fn read_unknown_frame_type_with_oversized_payload() { + let error = frame_error(Frame::read(&mut Cursor::new(encoded_frame(99, (MAX_FRAME_PAYLOAD + 1) as u32, &[])))); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert!(error.to_string().contains("unknown frame type")); +} + +#[test] +fn read_truncated_header() { + let error = frame_error(Frame::read(&mut Cursor::new(vec![1, 0, 0]))); + assert_eq!(error.kind(), io::ErrorKind::UnexpectedEof); +} + +#[test] +fn read_truncated_payload() { + let error = frame_error(Frame::read(&mut Cursor::new(encoded_frame(FrameType::Stdout as u16, 4, &[1, 2])))); + assert_eq!(error.kind(), io::ErrorKind::UnexpectedEof); +} + +#[test] +fn read_fragmented_frame() { + let payload = b"fragmented"; + let frame = Frame::read(&mut FragmentedReader::new(encoded_frame(FrameType::Stderr as u16, payload.len() as u32, payload), 1)).unwrap(); + assert!(matches!(frame.frame_type, FrameType::Stderr)); + assert_eq!(frame.payload, payload); +} + +#[test] +fn write_preserves_wire_format() { + let mut encoded = Vec::new(); + Frame::new(FrameType::Exit, vec![1, 2, 3]).write(&mut encoded).unwrap(); + assert_eq!(encoded, encoded_frame(FrameType::Exit as u16, 3, &[1, 2, 3])); +} + +#[test] +fn write_rejects_oversized_frame() { + let mut encoded = Vec::new(); + let error = Frame::new(FrameType::Stdout, vec![0; MAX_FRAME_PAYLOAD + 1]).write(&mut encoded).unwrap_err(); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert!(encoded.is_empty()); +} + +#[tokio::test] +async fn async_read_fragmented_frame() { + let payload = b"fragmented async"; + let mut reader = FragmentedReader::new(encoded_frame(FrameType::Stdout as u16, payload.len() as u32, payload), 1); + let frame = Frame::read_async(&mut reader).await.unwrap(); + assert!(matches!(frame.frame_type, FrameType::Stdout)); + assert_eq!(frame.payload, payload); +} + +#[tokio::test] +async fn async_read_zero_length_frame() { + let mut reader = FragmentedReader::new(encoded_frame(FrameType::Stdout as u16, 0, &[]), 1); + let frame = Frame::read_async(&mut reader).await.unwrap(); + assert!(matches!(frame.frame_type, FrameType::Stdout)); + assert!(frame.payload.is_empty()); +} + +#[tokio::test] +async fn async_read_maximum_size_frame() { + let payload = vec![0x5A; MAX_FRAME_PAYLOAD]; + let mut reader = FragmentedReader::new(encoded_frame(FrameType::Stdout as u16, MAX_FRAME_PAYLOAD as u32, &payload), MAX_FRAME_PAYLOAD); + let frame = Frame::read_async(&mut reader).await.unwrap(); + assert_eq!(frame.payload, payload); +} + +#[tokio::test] +async fn async_read_oversized_frame() { + let mut reader = FragmentedReader::new(encoded_frame(FrameType::Stdout as u16, (MAX_FRAME_PAYLOAD + 1) as u32, &[]), 1); + let error = frame_error(Frame::read_async(&mut reader).await); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); +} + +#[tokio::test] +async fn async_read_unknown_frame_type() { + let mut reader = FragmentedReader::new(encoded_frame(99, 0, &[]), 1); + let error = frame_error(Frame::read_async(&mut reader).await); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert!(error.to_string().contains("unknown frame type")); +} + +#[tokio::test] +async fn async_read_truncated_header() { + let mut reader = FragmentedReader::new(vec![1, 0, 0], 1); + let error = frame_error(Frame::read_async(&mut reader).await); + assert_eq!(error.kind(), io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn async_read_truncated_payload() { + let mut reader = FragmentedReader::new(encoded_frame(FrameType::Stdout as u16, 4, &[1, 2]), 1); + let error = frame_error(Frame::read_async(&mut reader).await); + assert_eq!(error.kind(), io::ErrorKind::UnexpectedEof); +} + +#[tokio::test] +async fn async_write_rejects_oversized_frame() { + let mut encoded = Vec::new(); + let error = Frame::new(FrameType::Stdout, vec![0; MAX_FRAME_PAYLOAD + 1]).write_async(&mut encoded).await.unwrap_err(); + assert_eq!(error.kind(), io::ErrorKind::InvalidData); + assert!(encoded.is_empty()); +} From 212222129f16d5d152324744fe79edbf5f4bea20 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sat, 1 Aug 2026 14:06:03 +0200 Subject: [PATCH 03/52] Validate guest working directories and workspace paths --- src/daemon.rs | 26 +++++++++----------- src/workspace.rs | 63 +++++++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 73 insertions(+), 16 deletions(-) diff --git a/src/daemon.rs b/src/daemon.rs index 8a8d073..006d305 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -3,6 +3,7 @@ use crate::logging; use crate::proxy::{FilterProxy, UnixProxyHandle}; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; use crate::vscomm::{validate_exec_request, validate_process_path, validate_process_string, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; +use crate::workspace::WorkspaceCwd; use rand::Rng; use std::fs::File; use std::io::{BufRead, BufReader}; @@ -183,12 +184,7 @@ async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSess async fn execute_request(writer: &mut W, session: &VsockSession, req: &ExecRequest) -> Result<(), String> { validate_exec_request(req)?; - let sandbox_cwd = req.cwd.clone(); - let host_cwd = if req.cwd.starts_with("/workspace") { - session.workspace.join(req.cwd.strip_prefix("/workspace").unwrap_or(&req.cwd).trim_start_matches('/')) - } else { - PathBuf::from(&req.cwd) - }; + let cwd = WorkspaceCwd::resolve(&session.workspace, Path::new(&req.cwd))?; let (status_reader, status_writer) = if session.merged_profile.is_some() { let (reader, writer) = bwrap_status_pipe()?; @@ -196,7 +192,7 @@ async fn execute_request(writer: &mut W, session: &Vso } else { (None, None) }; - let mut cmd = build_command(session, req, &host_cwd, &sandbox_cwd)?; + let mut cmd = build_command(session, req, &cwd)?; if let Some(status_writer) = status_writer.as_ref() { attach_bwrap_status_fd(&mut cmd, status_writer.as_raw_fd()); } @@ -267,12 +263,12 @@ async fn execute_request(writer: &mut W, session: &Vso Ok(()) } -fn build_command(session: &VsockSession, req: &ExecRequest, host_cwd: &Path, sandbox_cwd: &str) -> Result { +fn build_command(session: &VsockSession, req: &ExecRequest, cwd: &WorkspaceCwd) -> Result { validate_exec_request(req)?; validate_process_path("workspace path", &session.workspace)?; - validate_process_path("host working directory", host_cwd)?; - validate_process_string("sandbox working directory", sandbox_cwd)?; - + validate_process_path("host working directory", cwd.host_path())?; + let sandbox_cwd = cwd.guest_path(); + validate_process_path("sandbox working directory", &sandbox_cwd)?; if let Some(ref merged) = session.merged_profile { let mut cmd = Command::new("bwrap"); @@ -321,10 +317,10 @@ fn build_command(session: &VsockSession, req: &ExecRequest, host_cwd: &Path, san cmd.arg("--dev").arg("/dev"); cmd.arg("--tmpfs").arg("/tmp"); - if !sandbox_cwd.is_empty() && sandbox_cwd != "/" { - cmd.arg("--dir").arg(sandbox_cwd); + if sandbox_cwd != Path::new("/") { + cmd.arg("--dir").arg(&sandbox_cwd); } - cmd.arg("--chdir").arg(sandbox_cwd); + cmd.arg("--chdir").arg(&sandbox_cwd); cmd.arg("--clearenv"); cmd.arg("--setenv").arg("PATH").arg("/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"); @@ -383,7 +379,7 @@ fn build_command(session: &VsockSession, req: &ExecRequest, host_cwd: &Path, san } else { let mut cmd = Command::new(&req.command); cmd.args(&req.args); - cmd.current_dir(host_cwd); + cmd.current_dir(cwd.host_path()); if session.env_mode == EnvMode::Relaxed { for (key, val) in &req.env { diff --git a/src/workspace.rs b/src/workspace.rs index 769121a..8a3e7cd 100644 --- a/src/workspace.rs +++ b/src/workspace.rs @@ -3,7 +3,7 @@ use crate::overlay::CowWorkspace; use std::ffi::OsStr; use std::fs; use std::os::unix::fs::PermissionsExt; -use std::path::{Path, PathBuf}; +use std::path::{Component, Path, PathBuf}; use std::process::{Command, Stdio}; pub fn prepare(reset: bool) -> Result<(), String> { @@ -18,6 +18,63 @@ pub enum WorkspaceHandle { Isolated { path: PathBuf }, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkspaceCwd { + host: PathBuf, + relative: PathBuf, +} + +impl WorkspaceCwd { + pub fn resolve(workspace: &Path, guest_cwd: &Path) -> Result { + if !guest_cwd.is_absolute() { + return Err(format!("working directory must be absolute: {}", guest_cwd.display())); + } + + let relative = guest_cwd + .strip_prefix(Path::new("/workspace")) + .map_err(|_| format!("working directory must be under /workspace: {}", guest_cwd.display()))? + .components() + .try_fold(PathBuf::new(), |mut relative, component| match component { + Component::CurDir => Ok(relative), + Component::Normal(name) => { + relative.push(name); + Ok(relative) + } + Component::ParentDir => Err(format!("working directory contains '..': {}", guest_cwd.display())), + Component::RootDir | Component::Prefix(_) => Err(format!("invalid workspace path: {}", guest_cwd.display())), + })?; + + let canonical_workspace = fs::canonicalize(workspace).map_err(|err| format!("failed to resolve workspace {}: {err}", workspace.display()))?; + let candidate = canonical_workspace.join(&relative); + let canonical_host = + fs::canonicalize(&candidate).map_err(|err| format!("failed to resolve working directory {}: {err}", guest_cwd.display()))?; + + canonical_host.strip_prefix(&canonical_workspace).map_err(|_| format!("working directory escapes workspace: {}", guest_cwd.display()))?; + + if !fs::metadata(&canonical_host).map(|metadata| metadata.is_dir()).unwrap_or(false) { + return Err(format!("working directory is not a directory: {}", guest_cwd.display())); + } + + Ok(Self { host: canonical_host, relative }) + } + + pub fn host_path(&self) -> &Path { + &self.host + } + + pub fn relative_path(&self) -> &Path { + &self.relative + } + + pub fn guest_path(&self) -> PathBuf { + if self.relative.as_os_str().is_empty() { + PathBuf::from("/workspace") + } else { + Path::new("/workspace").join(&self.relative) + } + } +} + impl WorkspaceHandle { pub fn path(&self) -> &Path { match self { @@ -152,3 +209,7 @@ fn copy_dir(source: &Path, destination: &Path) -> Result<(), String> { fn should_skip(name: &OsStr) -> bool { matches!(name.to_str(), Some(".bunker") | Some(".bunkerbox") | Some(".git") | Some("target")) } + +#[cfg(test)] +#[path = "workspace_ut.rs"] +mod tests; From c238907d227a049ff51d6976fb5f109805b737d8 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sat, 1 Aug 2026 14:06:11 +0200 Subject: [PATCH 04/52] Add workspace unit tests --- src/workspace_ut.rs | 166 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 166 insertions(+) create mode 100644 src/workspace_ut.rs diff --git a/src/workspace_ut.rs b/src/workspace_ut.rs new file mode 100644 index 0000000..1cd1071 --- /dev/null +++ b/src/workspace_ut.rs @@ -0,0 +1,166 @@ +use super::*; +use std::os::unix::fs::symlink; +use tempfile::TempDir; + +fn workspace() -> TempDir { + TempDir::new().unwrap() +} + +fn directory(root: &Path, relative: &str) -> PathBuf { + let path = root.join(relative); + std::fs::create_dir_all(&path).unwrap(); + path +} + +fn resolve(root: &Path, cwd: &str) -> WorkspaceCwd { + WorkspaceCwd::resolve(root, Path::new(cwd)).unwrap() +} + +#[test] +fn resolves_workspace_root() { + let root = workspace(); + let cwd = resolve(root.path(), "/workspace"); + + assert_eq!(cwd.host_path(), root.path().canonicalize().unwrap()); + assert_eq!(cwd.relative_path(), Path::new("")); + assert_eq!(cwd.guest_path(), Path::new("/workspace")); +} + +#[test] +fn resolves_nested_directory() { + let root = workspace(); + let nested = directory(root.path(), "src/lib"); + let cwd = resolve(root.path(), "/workspace/src/lib"); + + assert_eq!(cwd.host_path(), nested.canonicalize().unwrap()); + assert_eq!(cwd.relative_path(), Path::new("src/lib")); + assert_eq!(cwd.guest_path(), Path::new("/workspace/src/lib")); +} + +#[test] +fn rejects_parent_directory_component() { + let root = workspace(); + directory(root.path(), "bar"); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace/foo/../bar")).is_err()); +} + +#[test] +fn normalizes_repeated_separators() { + let root = workspace(); + directory(root.path(), "foo"); + let cwd = resolve(root.path(), "/workspace//foo"); + + assert_eq!(cwd.relative_path(), Path::new("foo")); + assert_eq!(cwd.guest_path(), Path::new("/workspace/foo")); +} + +#[test] +fn normalizes_current_directory_components() { + let root = workspace(); + directory(root.path(), "foo"); + let cwd = resolve(root.path(), "/workspace/./foo"); + + assert_eq!(cwd.relative_path(), Path::new("foo")); + assert_eq!(cwd.guest_path(), Path::new("/workspace/foo")); +} + +#[test] +fn rejects_workspace_string_prefix_sibling() { + let root = workspace(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace-other/foo")).is_err()); +} + +#[test] +fn rejects_unrelated_absolute_path() { + let root = workspace(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/tmp/foo")).is_err()); +} + +#[test] +fn rejects_relative_path() { + let root = workspace(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("foo")).is_err()); +} + +#[test] +fn rejects_missing_path() { + let root = workspace(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace/missing")).is_err()); +} + +#[test] +fn rejects_file_as_working_directory() { + let root = workspace(); + std::fs::write(root.path().join("file"), b"data").unwrap(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace/file")).is_err()); +} + +#[test] +fn rejects_symlink_outside_workspace() { + let root = workspace(); + let outside = workspace(); + symlink(outside.path(), root.path().join("outside-link")).unwrap(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace/outside-link")).is_err()); +} + +#[test] +fn accepts_inside_symlink_and_preserves_guest_path() { + let root = workspace(); + let target = directory(root.path(), "real/foo"); + symlink(&target, root.path().join("foo-link")).unwrap(); + + let cwd = resolve(root.path(), "/workspace/foo-link"); + + assert_eq!(cwd.host_path(), target.canonicalize().unwrap()); + assert_eq!(cwd.relative_path(), Path::new("foo-link")); + assert_eq!(cwd.guest_path(), Path::new("/workspace/foo-link")); +} + +#[test] +fn rejects_dangling_symlink() { + let root = workspace(); + symlink(root.path().join("missing"), root.path().join("dangling-link")).unwrap(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace/dangling-link")).is_err()); +} + +#[test] +fn rejects_symlink_to_common_prefix_sibling() { + let parent = workspace(); + let root = parent.path().join("project"); + let sibling = parent.path().join("project-other"); + std::fs::create_dir_all(&root).unwrap(); + std::fs::create_dir_all(&sibling).unwrap(); + symlink(&sibling, root.join("sibling-link")).unwrap(); + + assert!(WorkspaceCwd::resolve(&root, Path::new("/workspace/sibling-link")).is_err()); +} + +#[test] +fn accepts_workspace_root_symlink() { + let parent = workspace(); + let target = workspace(); + let link = parent.path().join("workspace-link"); + symlink(target.path(), &link).unwrap(); + + let cwd = resolve(&link, "/workspace"); + + assert_eq!(cwd.host_path(), target.path().canonicalize().unwrap()); + assert_eq!(cwd.relative_path(), Path::new("")); +} + +#[test] +fn rejects_symlink_loop() { + let root = workspace(); + symlink(root.path().join("loop-b"), root.path().join("loop-a")).unwrap(); + symlink(root.path().join("loop-a"), root.path().join("loop-b")).unwrap(); + + assert!(WorkspaceCwd::resolve(root.path(), Path::new("/workspace/loop-a")).is_err()); +} From 52aff66cbb508485c3bf8c937516be1d4c242847 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sat, 1 Aug 2026 21:37:22 +0200 Subject: [PATCH 05/52] Lints --- src/daemon.rs | 2 +- src/main.rs | 2 -- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/src/daemon.rs b/src/daemon.rs index 006d305..4640315 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -2,7 +2,7 @@ use crate::cfg::EnvMode; use crate::logging; use crate::proxy::{FilterProxy, UnixProxyHandle}; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; -use crate::vscomm::{validate_exec_request, validate_process_path, validate_process_string, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; +use crate::vscomm::{validate_exec_request, validate_process_path, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; use crate::workspace::WorkspaceCwd; use rand::Rng; use std::fs::File; diff --git a/src/main.rs b/src/main.rs index 784d643..a7ec69c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -422,8 +422,6 @@ fn start_status_listener(overlay: Arc>) -> Result>, mut shutdown_rx: tokio::sync::oneshot::Receiver<()>, ) { - use tokio::io::AsyncReadExt; - loop { let (mut stream, _peer) = tokio::select! { result = listener.accept() => match result { From 2caa623b76c35355de0ec223055515ee9352288e Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 01:36:39 +0200 Subject: [PATCH 06/52] Define the versioned remote operation proto --- src/vscomm/mod.rs | 461 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 461 insertions(+) diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index 3ee4b17..be54f54 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -16,6 +16,15 @@ pub const TUI_STATUS_PORT: u32 = 10000; pub const VSCOMM_BIN_DIR: &str = "/usr/local/bunkerbox/bin"; /// Maximum payload accepted in one vsock frame. pub const MAX_FRAME_PAYLOAD: usize = 1024 * 1024; +pub const REMOTE_PROTOCOL_VERSION: u16 = 1; +pub const MAX_REMOTE_STRING_BYTES: usize = 4 * 1024; +pub const MAX_REMOTE_TOOL_BYTES: usize = 256; +pub const MAX_REMOTE_ARG_COUNT: usize = 256; +pub const MAX_REMOTE_ARG_BYTES: usize = 4 * 1024; +pub const MAX_REMOTE_ENV_COUNT: usize = 64; +pub const MAX_REMOTE_ENV_KEY_BYTES: usize = 256; +pub const MAX_REMOTE_ENV_VALUE_BYTES: usize = 4 * 1024; +pub const MAX_REMOTE_ERROR_BYTES: usize = 4 * 1024; #[repr(u16)] #[derive(Clone, Copy)] @@ -26,6 +35,8 @@ pub enum FrameType { Exit = 4, Disconnect = 5, UiCommand = 10, + RemoteRequest = 20, + RemoteEvent = 21, } impl FrameType { @@ -37,6 +48,8 @@ impl FrameType { 4 => Some(Self::Exit), 5 => Some(Self::Disconnect), 10 => Some(Self::UiCommand), + 20 => Some(Self::RemoteRequest), + 21 => Some(Self::RemoteEvent), _ => None, } } @@ -49,6 +62,454 @@ pub struct ExecRequest { pub env: Vec<(String, String)>, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RequestId(pub [u8; 16]); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct WorkspaceSessionId(pub [u8; 16]); + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkspaceRelativePath(String); + +impl WorkspaceRelativePath { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + validate_remote_string("remote cwd", &value, MAX_REMOTE_STRING_BYTES)?; + if value.is_empty() { + return Ok(Self(value)); + } + + let path = Path::new(&value); + if path.is_absolute() || value.split('/').any(|component| component.is_empty() || component == "." || component == "..") { + return Err("remote cwd must be a normalized relative path".to_string()); + } + + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteTool(String); + +impl RemoteTool { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + validate_remote_string("remote tool", &value, MAX_REMOTE_TOOL_BYTES)?; + if value.is_empty() { + return Err("remote tool is empty".to_string()); + } + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteBuild { + pub cwd: WorkspaceRelativePath, + pub tool: RemoteTool, + pub argv: Vec, + pub env: Vec<(String, String)>, +} + +impl RemoteBuild { + pub fn new(cwd: WorkspaceRelativePath, tool: RemoteTool, argv: Vec, env: Vec<(String, String)>) -> Result { + validate_remote_build_fields(&cwd, &tool, &argv, &env)?; + Ok(Self { cwd, tool, argv, env }) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RemoteOperation { + Sync(RemoteSync), + Build(RemoteBuild), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteSync; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteRequest { + pub request_id: RequestId, + pub workspace_session_id: WorkspaceSessionId, + pub operation: RemoteOperation, +} + +impl RemoteRequest { + pub fn sync(request_id: RequestId, workspace_session_id: WorkspaceSessionId) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Sync(RemoteSync) } + } + + pub fn build(request_id: RequestId, workspace_session_id: WorkspaceSessionId, build: RemoteBuild) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Build(build) } + } + + pub fn to_frame(&self) -> Result { + let mut writer = WireWriter::new(*b"BBR1"); + writer.u16(REMOTE_PROTOCOL_VERSION); + writer.u8(match &self.operation { + RemoteOperation::Sync(_) => 1, + RemoteOperation::Build(_) => 2, + }); + writer.u8(0); + writer.bytes(&self.request_id.0); + writer.bytes(&self.workspace_session_id.0); + + if let RemoteOperation::Build(build) = &self.operation { + encode_remote_build(&mut writer, build)?; + } + + writer.into_frame(FrameType::RemoteRequest) + } + + pub fn from_frame(frame: Frame) -> Result { + if !matches!(frame.frame_type, FrameType::RemoteRequest) { + return Err("expected RemoteRequest frame".to_string()); + } + + let mut reader = WireReader::new(&frame.payload); + reader.magic(*b"BBR1")?; + reader.version()?; + let operation_kind = reader.u8()?; + reader.zero_reserved()?; + let request_id = RequestId(reader.array16()?); + let workspace_session_id = WorkspaceSessionId(reader.array16()?); + let operation = match operation_kind { + 1 => RemoteOperation::Sync(RemoteSync), + 2 => RemoteOperation::Build(decode_remote_build(&mut reader)?), + value => return Err(format!("unknown remote operation: {value}")), + }; + let request = Self { request_id, workspace_session_id, operation }; + reader.finish()?; + Ok(request) + } +} + +fn encode_remote_build(writer: &mut WireWriter, build: &RemoteBuild) -> Result<(), String> { + validate_remote_build_fields(&build.cwd, &build.tool, &build.argv, &build.env)?; + writer.string(build.cwd.as_str(), MAX_REMOTE_STRING_BYTES, "remote cwd")?; + writer.string(build.tool.as_str(), MAX_REMOTE_TOOL_BYTES, "remote tool")?; + writer.count(build.argv.len(), MAX_REMOTE_ARG_COUNT, "remote argv")?; + for arg in &build.argv { + writer.string(arg, MAX_REMOTE_ARG_BYTES, "remote argument")?; + } + writer.count(build.env.len(), MAX_REMOTE_ENV_COUNT, "remote environment")?; + for (key, value) in &build.env { + writer.string(key, MAX_REMOTE_ENV_KEY_BYTES, "remote environment key")?; + writer.string(value, MAX_REMOTE_ENV_VALUE_BYTES, "remote environment value")?; + } + Ok(()) +} + +fn decode_remote_build(reader: &mut WireReader<'_>) -> Result { + let cwd = WorkspaceRelativePath::new(reader.string(MAX_REMOTE_STRING_BYTES, "remote cwd")?)?; + let tool = RemoteTool::new(reader.string(MAX_REMOTE_TOOL_BYTES, "remote tool")?)?; + let argv = (0..reader.count(MAX_REMOTE_ARG_COUNT, "remote argv")?) + .map(|_| reader.string(MAX_REMOTE_ARG_BYTES, "remote argument")) + .collect::, _>>()?; + let env = (0..reader.count(MAX_REMOTE_ENV_COUNT, "remote environment")?) + .map(|_| { + Ok(( + reader.string(MAX_REMOTE_ENV_KEY_BYTES, "remote environment key")?, + reader.string(MAX_REMOTE_ENV_VALUE_BYTES, "remote environment value")?, + )) + }) + .collect::, String>>()?; + RemoteBuild::new(cwd, tool, argv, env) +} + +fn validate_remote_build_fields(cwd: &WorkspaceRelativePath, tool: &RemoteTool, argv: &[String], env: &[(String, String)]) -> Result<(), String> { + validate_remote_string("remote cwd", cwd.as_str(), MAX_REMOTE_STRING_BYTES)?; + validate_remote_string("remote tool", tool.as_str(), MAX_REMOTE_TOOL_BYTES)?; + validate_remote_count(argv.len(), MAX_REMOTE_ARG_COUNT, "remote argv")?; + argv.iter().try_for_each(|arg| validate_remote_string("remote argument", arg, MAX_REMOTE_ARG_BYTES))?; + validate_remote_count(env.len(), MAX_REMOTE_ENV_COUNT, "remote environment")?; + env.iter().try_for_each(|(key, value)| { + validate_remote_string("remote environment key", key, MAX_REMOTE_ENV_KEY_BYTES)?; + validate_env_key("remote environment key", key)?; + validate_remote_string("remote environment value", value, MAX_REMOTE_ENV_VALUE_BYTES) + }) +} + +fn validate_remote_string(field: &str, value: &str, max: usize) -> Result<(), String> { + validate_process_string(field, value)?; + if value.len() > max { + return Err(format!("{field} exceeds maximum length {max}")); + } + Ok(()) +} + +fn validate_remote_count(count: usize, max: usize, field: &str) -> Result<(), String> { + if count > max { + return Err(format!("{field} exceeds maximum count {max}")); + } + Ok(()) +} + +struct WireWriter { + bytes: Vec, +} + +impl WireWriter { + fn new(magic: [u8; 4]) -> Self { + Self { bytes: magic.to_vec() } + } + + fn u8(&mut self, value: u8) { + self.bytes.push(value); + } + + fn u16(&mut self, value: u16) { + self.bytes.extend_from_slice(&value.to_le_bytes()); + } + + fn u64(&mut self, value: u64) { + self.bytes.extend_from_slice(&value.to_le_bytes()); + } + + fn i32(&mut self, value: i32) { + self.bytes.extend_from_slice(&value.to_le_bytes()); + } + + fn bytes(&mut self, value: &[u8]) { + self.bytes.extend_from_slice(value); + } + + fn count(&mut self, count: usize, max: usize, field: &str) -> Result<(), String> { + validate_remote_count(count, max, field)?; + self.u16(count as u16); + Ok(()) + } + + fn string(&mut self, value: &str, max: usize, field: &str) -> Result<(), String> { + validate_remote_string(field, value, max)?; + let length = u16::try_from(value.len()).map_err(|_| format!("{field} is too long"))?; + self.u16(length); + self.bytes(value.as_bytes()); + Ok(()) + } + + fn blob(&mut self, value: &[u8], max: usize, field: &str) -> Result<(), String> { + if value.len() > max { + return Err(format!("{field} exceeds maximum length {max}")); + } + let length = u32::try_from(value.len()).map_err(|_| format!("{field} is too long"))?; + self.bytes.extend_from_slice(&length.to_le_bytes()); + self.bytes(value); + Ok(()) + } + + fn into_frame(self, frame_type: FrameType) -> Result { + if self.bytes.len() > MAX_FRAME_PAYLOAD { + return Err(format!("remote payload exceeds frame limit {MAX_FRAME_PAYLOAD}")); + } + Ok(Frame::new(frame_type, self.bytes)) + } +} + +struct WireReader<'a> { + bytes: &'a [u8], + offset: usize, +} + +impl<'a> WireReader<'a> { + fn new(bytes: &'a [u8]) -> Self { + Self { bytes, offset: 0 } + } + + fn take(&mut self, length: usize) -> Result<&'a [u8], String> { + let end = self.offset.checked_add(length).ok_or_else(|| "remote payload length overflow".to_string())?; + let value = self.bytes.get(self.offset..end).ok_or_else(|| "truncated remote payload".to_string())?; + self.offset = end; + Ok(value) + } + + fn magic(&mut self, expected: [u8; 4]) -> Result<(), String> { + if self.take(4)? != expected { + return Err("invalid remote payload magic".to_string()); + } + Ok(()) + } + + fn version(&mut self) -> Result<(), String> { + let version = self.u16()?; + if version != REMOTE_PROTOCOL_VERSION { + return Err(format!("unsupported remote protocol version: {version}")); + } + Ok(()) + } + + fn zero_reserved(&mut self) -> Result<(), String> { + if self.u8()? != 0 { + return Err("remote payload reserved byte is nonzero".to_string()); + } + Ok(()) + } + + fn u8(&mut self) -> Result { + Ok(self.take(1)?[0]) + } + + fn u16(&mut self) -> Result { + let bytes = self.take(2)?; + Ok(u16::from_le_bytes([bytes[0], bytes[1]])) + } + + fn u64(&mut self) -> Result { + let bytes = self.take(8)?; + Ok(u64::from_le_bytes(bytes.try_into().map_err(|_| "invalid remote integer".to_string())?)) + } + + fn i32(&mut self) -> Result { + let bytes = self.take(4)?; + Ok(i32::from_le_bytes(bytes.try_into().map_err(|_| "invalid remote integer".to_string())?)) + } + + fn array16(&mut self) -> Result<[u8; 16], String> { + self.take(16)?.try_into().map_err(|_| "invalid remote identifier".to_string()) + } + + fn count(&mut self, max: usize, field: &str) -> Result { + let count = self.u16()? as usize; + validate_remote_count(count, max, field)?; + Ok(count) + } + + fn string(&mut self, max: usize, field: &str) -> Result { + let length = self.u16()? as usize; + if length > max { + return Err(format!("{field} exceeds maximum length {max}")); + } + let value = std::str::from_utf8(self.take(length)?).map_err(|_| format!("{field} is not valid UTF-8"))?; + validate_remote_string(field, value, max)?; + Ok(value.to_string()) + } + + fn blob(&mut self, max: usize, field: &str) -> Result, String> { + let bytes = self.take(4)?; + let length = u32::from_le_bytes(bytes.try_into().map_err(|_| "invalid remote blob length".to_string())?) as usize; + if length > max { + return Err(format!("{field} exceeds maximum length {max}")); + } + Ok(self.take(length)?.to_vec()) + } + + fn finish(self) -> Result<(), String> { + if self.offset == self.bytes.len() { + Ok(()) + } else { + Err("trailing bytes in remote payload".to_string()) + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RemoteErrorCode { + Failed = 1, +} + +impl RemoteErrorCode { + fn from_u16(value: u16) -> Result { + match value { + 1 => Ok(Self::Failed), + _ => Err(format!("unknown remote error code: {value}")), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RemoteEventKind { + SyncProgress { completed_bytes: u64, total_bytes: Option }, + Stdout(Vec), + Stderr(Vec), + Error { code: RemoteErrorCode, message: String }, + Cancelled, + Completed { exit_code: i32 }, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteEvent { + pub request_id: RequestId, + pub kind: RemoteEventKind, +} + +impl RemoteEvent { + pub fn to_frame(&self) -> Result { + let mut writer = WireWriter::new(*b"BBE1"); + writer.u16(REMOTE_PROTOCOL_VERSION); + writer.u8(match &self.kind { + RemoteEventKind::SyncProgress { .. } => 1, + RemoteEventKind::Stdout(_) => 2, + RemoteEventKind::Stderr(_) => 3, + RemoteEventKind::Error { .. } => 4, + RemoteEventKind::Cancelled => 5, + RemoteEventKind::Completed { .. } => 6, + }); + writer.u8(0); + writer.bytes(&self.request_id.0); + + match &self.kind { + RemoteEventKind::SyncProgress { completed_bytes, total_bytes } => { + writer.u64(*completed_bytes); + writer.u8(u8::from(total_bytes.is_some())); + if let Some(total_bytes) = total_bytes { + writer.u64(*total_bytes); + } + } + RemoteEventKind::Stdout(data) | RemoteEventKind::Stderr(data) => writer.blob(data, MAX_FRAME_PAYLOAD, "remote output")?, + RemoteEventKind::Error { code, message } => { + writer.u16(*code as u16); + writer.string(message, MAX_REMOTE_ERROR_BYTES, "remote error")?; + } + RemoteEventKind::Cancelled => {} + RemoteEventKind::Completed { exit_code } => writer.i32(*exit_code), + } + + writer.into_frame(FrameType::RemoteEvent) + } + + pub fn from_frame(frame: Frame) -> Result { + if !matches!(frame.frame_type, FrameType::RemoteEvent) { + return Err("expected RemoteEvent frame".to_string()); + } + + let mut reader = WireReader::new(&frame.payload); + reader.magic(*b"BBE1")?; + reader.version()?; + let event_kind = reader.u8()?; + reader.zero_reserved()?; + let request_id = RequestId(reader.array16()?); + let kind = match event_kind { + 1 => { + let completed_bytes = reader.u64()?; + let total_bytes = match reader.u8()? { + 0 => None, + 1 => Some(reader.u64()?), + value => return Err(format!("invalid remote progress total flag: {value}")), + }; + RemoteEventKind::SyncProgress { completed_bytes, total_bytes } + } + 2 => RemoteEventKind::Stdout(reader.blob(MAX_FRAME_PAYLOAD, "remote stdout")?), + 3 => RemoteEventKind::Stderr(reader.blob(MAX_FRAME_PAYLOAD, "remote stderr")?), + 4 => RemoteEventKind::Error { + code: RemoteErrorCode::from_u16(reader.u16()?)?, + message: reader.string(MAX_REMOTE_ERROR_BYTES, "remote error")?, + }, + 5 => RemoteEventKind::Cancelled, + 6 => RemoteEventKind::Completed { exit_code: reader.i32()? }, + value => return Err(format!("unknown remote event kind: {value}")), + }; + reader.finish()?; + Ok(Self { request_id, kind }) + } +} + pub fn validate_process_string(field: &str, value: &str) -> Result<(), String> { if value.as_bytes().contains(&0) { return Err(format!("{field} contains a NUL byte")); From 537bddb2d8938dbd631e0c1fd9a9af3c9c14f517 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 01:36:52 +0200 Subject: [PATCH 07/52] Add remote operation proto UT --- src/vscomm/mod_ut.rs | 168 ++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 167 insertions(+), 1 deletion(-) diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index f52e0a9..c9874e2 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -1,4 +1,37 @@ -use super::{validate_exec_request, ExecRequest, FrameType, TOOLCHAIN_PORT, TUI_STATUS_PORT}; +use super::*; + +fn ids() -> (RequestId, WorkspaceSessionId) { + (RequestId([1; 16]), WorkspaceSessionId([2; 16])) +} + +fn build(argv: Vec, env: Vec<(String, String)>) -> Result { + RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new("make").unwrap(), argv, env) +} + +fn raw_build_frame(argv_count: u16, arg: Option<&str>, env_count: u16, env: Option<(&str, &str)>) -> Frame { + let mut writer = WireWriter::new(*b"BBR1"); + writer.u16(REMOTE_PROTOCOL_VERSION); + writer.u8(2); + writer.u8(0); + writer.bytes(&[1; 16]); + writer.bytes(&[2; 16]); + writer.u16(0); + writer.u16(4); + writer.bytes(b"make"); + writer.u16(argv_count); + if let Some(arg) = arg { + writer.u16(arg.len() as u16); + writer.bytes(arg.as_bytes()); + } + writer.u16(env_count); + if let Some((key, value)) = env { + writer.u16(key.len() as u16); + writer.bytes(key.as_bytes()); + writer.u16(value.len() as u16); + writer.bytes(value.as_bytes()); + } + writer.into_frame(FrameType::RemoteRequest).unwrap() +} #[test] fn execution_and_tui_channels_are_distinct() { @@ -33,3 +66,136 @@ fn accept_valid_cargo_request() { assert!(validate_exec_request(&request).is_ok()); } + +#[test] +fn remote_sync_round_trips() { + let (request_id, session_id) = ids(); + let request = RemoteRequest::sync(request_id, session_id); + let decoded = RemoteRequest::from_frame(request.to_frame().unwrap()).unwrap(); + assert_eq!(decoded, request); +} + +#[test] +fn remote_build_round_trips_structured_arguments_and_environment() { + let (request_id, session_id) = ids(); + let request = RemoteRequest::build( + request_id, + session_id, + build(vec!["release mode".into(), "$(not-a-shell-command)".into()], vec![("MODE".into(), "debug value".into())]).unwrap(), + ); + let decoded = RemoteRequest::from_frame(request.to_frame().unwrap()).unwrap(); + assert_eq!(decoded, request); +} + +#[test] +fn every_remote_event_round_trips() { + let request_id = ids().0; + let events = vec![ + RemoteEventKind::SyncProgress { completed_bytes: 4, total_bytes: Some(9) }, + RemoteEventKind::Stdout(b"out".to_vec()), + RemoteEventKind::Stderr(b"err".to_vec()), + RemoteEventKind::Error { code: RemoteErrorCode::Failed, message: "failed".into() }, + RemoteEventKind::Cancelled, + RemoteEventKind::Completed { exit_code: 17 }, + ]; + + for kind in events { + let event = RemoteEvent { request_id, kind }; + assert_eq!(RemoteEvent::from_frame(event.to_frame().unwrap()).unwrap(), event); + } +} + +#[test] +fn supported_remote_version_is_encoded() { + let frame = RemoteRequest::sync(ids().0, ids().1).to_frame().unwrap(); + assert_eq!(u16::from_le_bytes([frame.payload[4], frame.payload[5]]), REMOTE_PROTOCOL_VERSION); +} + +#[test] +fn unknown_remote_version_is_rejected() { + let mut frame = RemoteRequest::sync(ids().0, ids().1).to_frame().unwrap(); + frame.payload[4..6].copy_from_slice(&(REMOTE_PROTOCOL_VERSION + 1).to_le_bytes()); + assert!(RemoteRequest::from_frame(frame).unwrap_err().contains("unsupported remote protocol version")); +} + +#[test] +fn unknown_remote_operation_is_rejected() { + let mut frame = RemoteRequest::sync(ids().0, ids().1).to_frame().unwrap(); + frame.payload[6] = 99; + assert!(RemoteRequest::from_frame(frame).unwrap_err().contains("unknown remote operation")); +} + +#[test] +fn unknown_remote_event_kind_is_rejected() { + let event = RemoteEvent { request_id: ids().0, kind: RemoteEventKind::Cancelled }; + let mut frame = event.to_frame().unwrap(); + frame.payload[6] = 99; + assert!(RemoteEvent::from_frame(frame).unwrap_err().contains("unknown remote event kind")); +} + +#[test] +fn truncated_remote_payload_is_rejected() { + assert!(RemoteRequest::from_frame(Frame::new(FrameType::RemoteRequest, b"BBR1".to_vec())).is_err()); +} + +#[test] +fn oversized_remote_string_is_rejected() { + assert!(RemoteTool::new("x".repeat(MAX_REMOTE_TOOL_BYTES + 1)).is_err()); + let mut writer = WireWriter::new(*b"BBR1"); + writer.u16(REMOTE_PROTOCOL_VERSION); + writer.u8(2); + writer.u8(0); + writer.bytes(&[1; 16]); + writer.bytes(&[2; 16]); + writer.u16((MAX_REMOTE_STRING_BYTES + 1) as u16); + let frame = writer.into_frame(FrameType::RemoteRequest).unwrap(); + assert!(RemoteRequest::from_frame(frame).is_err()); +} + +#[test] +fn excessive_argv_count_is_rejected() { + assert!(build(vec!["arg".into(); MAX_REMOTE_ARG_COUNT + 1], Vec::new()).is_err()); + assert!(RemoteRequest::from_frame(raw_build_frame((MAX_REMOTE_ARG_COUNT + 1) as u16, None, 0, None)).is_err()); +} + +#[test] +fn oversized_individual_argument_is_rejected() { + assert!(build(vec!["x".repeat(MAX_REMOTE_ARG_BYTES + 1)], Vec::new()).is_err()); + assert!(RemoteRequest::from_frame(raw_build_frame(1, Some(&"x".repeat(MAX_REMOTE_ARG_BYTES + 1)), 0, None)).is_err()); +} + +#[test] +fn excessive_environment_count_is_rejected() { + let env = (0..MAX_REMOTE_ENV_COUNT + 1).map(|i| (format!("KEY{i}"), "value".into())).collect(); + assert!(build(Vec::new(), env).is_err()); + assert!(RemoteRequest::from_frame(raw_build_frame(0, None, (MAX_REMOTE_ENV_COUNT + 1) as u16, None)).is_err()); +} + +#[test] +fn oversized_environment_key_and_value_are_rejected() { + assert!(build(Vec::new(), vec![("K".repeat(MAX_REMOTE_ENV_KEY_BYTES + 1), "value".into())]).is_err()); + assert!(build(Vec::new(), vec![("KEY".into(), "V".repeat(MAX_REMOTE_ENV_VALUE_BYTES + 1))]).is_err()); + assert!(RemoteRequest::from_frame(raw_build_frame(0, None, 1, Some((&"K".repeat(MAX_REMOTE_ENV_KEY_BYTES + 1), "value")))).is_err()); + assert!(RemoteRequest::from_frame(raw_build_frame(0, None, 1, Some(("KEY", &"V".repeat(MAX_REMOTE_ENV_VALUE_BYTES + 1))))).is_err()); +} + +#[test] +fn invalid_remote_cwd_is_rejected() { + assert!(WorkspaceRelativePath::new("/absolute").is_err()); + assert!(WorkspaceRelativePath::new("foo/../bar").is_err()); + assert!(WorkspaceRelativePath::new("foo/./bar").is_err()); + assert!(WorkspaceRelativePath::new("foo//bar").is_err()); +} + +#[test] +fn existing_exec_request_wire_format_is_unchanged() { + let request = + ExecRequest { cwd: "/workspace".into(), command: "make".into(), args: vec!["release".into()], env: vec![("MODE".into(), "debug".into())] }; + let encoded = request.serialize(); + assert_eq!(encoded, b"/workspace\0make\0release\0\0MODE=debug\0\0"); + let decoded = ExecRequest::deserialize(&encoded).unwrap(); + assert_eq!(decoded.cwd, request.cwd); + assert_eq!(decoded.command, request.command); + assert_eq!(decoded.args, request.args); + assert_eq!(decoded.env, request.env); +} From 4b2ace81a05793a358d55031a0d833546d5e82ae Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 17:53:58 +0200 Subject: [PATCH 08/52] Add remote auth and backend interfaces --- src/bin/bunkerbox-status.rs | 2 + src/bin/bunkerbox-vscomm.rs | 2 + src/daemon.rs | 74 +++++++++-- src/lib.rs | 1 + src/remote.rs | 258 ++++++++++++++++++++++++++++++++++++ src/vscomm/mod.rs | 30 +++++ 6 files changed, 353 insertions(+), 14 deletions(-) create mode 100644 src/remote.rs diff --git a/src/bin/bunkerbox-status.rs b/src/bin/bunkerbox-status.rs index 51833b9..2024646 100644 --- a/src/bin/bunkerbox-status.rs +++ b/src/bin/bunkerbox-status.rs @@ -1,3 +1,5 @@ +#[path = "../remote.rs"] +mod remote; #[path = "../vscomm/mod.rs"] mod vscomm; diff --git a/src/bin/bunkerbox-vscomm.rs b/src/bin/bunkerbox-vscomm.rs index 3cf906d..77fff63 100644 --- a/src/bin/bunkerbox-vscomm.rs +++ b/src/bin/bunkerbox-vscomm.rs @@ -1,3 +1,5 @@ +#[path = "../remote.rs"] +mod remote; #[path = "../vscomm/mod.rs"] mod vscomm; diff --git a/src/daemon.rs b/src/daemon.rs index 4640315..d95a5f9 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -1,6 +1,9 @@ use crate::cfg::EnvMode; use crate::logging; use crate::proxy::{FilterProxy, UnixProxyHandle}; +use crate::remote::{ + RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionContext, RemoteRequest, +}; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; use crate::vscomm::{validate_exec_request, validate_process_path, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; use crate::workspace::WorkspaceCwd; @@ -17,8 +20,14 @@ use tokio::process::Command; const BWRAP_STATUS_FD: RawFd = 3; +#[derive(Clone, Copy)] +enum ProcessStream { + Stdout, + Stderr, +} + enum ChildEvent { - Output(FrameType, Vec), + Output(ProcessStream, Vec), StreamClosed, LauncherStarted, LauncherFailed(String), @@ -30,6 +39,36 @@ struct SandboxProxyConfig { netrelay_path: PathBuf, } +#[derive(Debug, PartialEq, Eq)] +pub enum RemoteDispatchError { + Unauthorized(RemoteAuthorizationError), + Backend(RemoteBackendError), + EventSinkClosed, +} + +pub struct RemoteBroker { + policy: RemoteAuthorizationPolicy, + context: RemoteExecutionContext, + backend: Arc, +} + +impl RemoteBroker { + pub fn new(policy: RemoteAuthorizationPolicy, context: RemoteExecutionContext, backend: Arc) -> Self { + Self { policy, context, backend } + } + + pub async fn dispatch(&self, request: RemoteRequest, events: tokio::sync::mpsc::Sender) -> Result<(), RemoteDispatchError> { + let authorized = self.policy.authorize(&self.context, request).map_err(RemoteDispatchError::Unauthorized)?; + match self.backend.execute(authorized, events.clone()).await { + Ok(()) => Ok(()), + Err(error) => { + events.send(error.event()).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; + Err(RemoteDispatchError::Backend(error)) + } + } + } +} + struct VsockSession { passthrough: Arc>, env_mode: EnvMode, @@ -209,8 +248,8 @@ async fn execute_request(writer: &mut W, session: &Vso let child_stderr = child.stderr.take().ok_or_else(|| "no stderr".to_string())?; let (event_tx, mut event_rx) = tokio::sync::mpsc::unbounded_channel(); - let stdout_task = tokio::spawn(pump_to_channel(child_stdout, FrameType::Stdout, event_tx.clone())); - let stderr_task = tokio::spawn(pump_to_channel(child_stderr, FrameType::Stderr, event_tx.clone())); + let stdout_task = tokio::spawn(pump_to_channel(child_stdout, ProcessStream::Stdout, event_tx.clone())); + let stderr_task = tokio::spawn(pump_to_channel(child_stderr, ProcessStream::Stderr, event_tx.clone())); let status_task = status_reader.map(|reader| { let status_tx = event_tx.clone(); tokio::task::spawn_blocking(move || monitor_bwrap_status(reader, status_tx)) @@ -224,26 +263,26 @@ async fn execute_request(writer: &mut W, session: &Vso while closed_streams < 2 || (session.merged_profile.is_some() && !launcher_started && !launcher_failed) { let Some(event) = event_rx.recv().await else { break }; match event { - ChildEvent::Output(frame_type, data) if launcher_failed => { - let stream = if matches!(frame_type, FrameType::Stdout) { "bwrap stdout" } else { "bwrap stderr" }; + ChildEvent::Output(stream_kind, data) if launcher_failed => { + let stream = if matches!(stream_kind, ProcessStream::Stdout) { "bwrap stdout" } else { "bwrap stderr" }; logging::diagnostic_bytes(stream, &data); } - ChildEvent::Output(frame_type, data) if launcher_started => { - write_frame(writer, &Frame::new(frame_type, data)).await?; + ChildEvent::Output(stream_kind, data) if launcher_started => { + write_frame(writer, &Frame::new(process_stream_frame_type(stream_kind), data)).await?; } - ChildEvent::Output(frame_type, data) => buffered_output.push((frame_type, data)), + ChildEvent::Output(stream_kind, data) => buffered_output.push((stream_kind, data)), ChildEvent::StreamClosed => closed_streams += 1, ChildEvent::LauncherStarted => { launcher_started = true; - for (frame_type, data) in buffered_output.drain(..) { - write_frame(writer, &Frame::new(frame_type, data)).await?; + for (stream_kind, data) in buffered_output.drain(..) { + write_frame(writer, &Frame::new(process_stream_frame_type(stream_kind), data)).await?; } } ChildEvent::LauncherFailed(err) => { launcher_failed = true; logging::diagnostic(&format!("bwrap setup failed: {err}")); - for (frame_type, data) in buffered_output.drain(..) { - let stream = if matches!(frame_type, FrameType::Stdout) { "bwrap stdout" } else { "bwrap stderr" }; + for (stream_kind, data) in buffered_output.drain(..) { + let stream = if matches!(stream_kind, ProcessStream::Stdout) { "bwrap stdout" } else { "bwrap stderr" }; logging::diagnostic_bytes(stream, &data); } } @@ -514,13 +553,13 @@ fn monitor_bwrap_status(reader: File, tx: tokio::sync::mpsc::UnboundedSender(mut reader: R, frame_type: FrameType, tx: tokio::sync::mpsc::UnboundedSender) { +async fn pump_to_channel(mut reader: R, stream_kind: ProcessStream, tx: tokio::sync::mpsc::UnboundedSender) { let mut buf = [0u8; 8192]; loop { match reader.read(&mut buf).await { Ok(0) => break, Ok(n) => { - if tx.send(ChildEvent::Output(frame_type, buf[..n].to_vec())).is_err() { + if tx.send(ChildEvent::Output(stream_kind, buf[..n].to_vec())).is_err() { return; } } @@ -530,6 +569,13 @@ async fn pump_to_channel(mut reader: R, frame_type: Fra let _ = tx.send(ChildEvent::StreamClosed); } +fn process_stream_frame_type(stream: ProcessStream) -> FrameType { + match stream { + ProcessStream::Stdout => FrameType::Stdout, + ProcessStream::Stderr => FrameType::Stderr, + } +} + async fn write_frame(writer: &mut W, frame: &Frame) -> Result<(), String> { frame.write_async(writer).await.map_err(|e| format!("write frame: {e}")) } diff --git a/src/lib.rs b/src/lib.rs index 2433330..b0e64d5 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -8,6 +8,7 @@ pub mod logging; pub mod netrelay; pub mod overlay; pub mod proxy; +pub mod remote; pub mod sandbox; pub mod tui; pub mod vscomm; diff --git a/src/remote.rs b/src/remote.rs new file mode 100644 index 0000000..4cdfb6c --- /dev/null +++ b/src/remote.rs @@ -0,0 +1,258 @@ +#![allow(dead_code)] + +use std::future::Future; +use std::pin::Pin; +use tokio::sync::mpsc; + +pub const MAX_REMOTE_STRING_BYTES: usize = 4 * 1024; +pub const MAX_REMOTE_TOOL_BYTES: usize = 256; +pub const MAX_REMOTE_ARG_COUNT: usize = 256; +pub const MAX_REMOTE_ARG_BYTES: usize = 4 * 1024; +pub const MAX_REMOTE_ENV_COUNT: usize = 64; +pub const MAX_REMOTE_ENV_KEY_BYTES: usize = 256; +pub const MAX_REMOTE_ENV_VALUE_BYTES: usize = 4 * 1024; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RequestId(pub [u8; 16]); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct WorkspaceSessionId(pub [u8; 16]); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteTargetId(pub [u8; 16]); + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkspaceRelativePath(String); + +impl WorkspaceRelativePath { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + validate_string("remote cwd", &value, MAX_REMOTE_STRING_BYTES)?; + if value.is_empty() { + return Ok(Self(value)); + } + + if value.starts_with('/') || value.split('/').any(|part| part.is_empty() || part == "." || part == "..") { + return Err("remote cwd must be a normalized relative path".to_string()); + } + + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteTool(String); + +impl RemoteTool { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + validate_string("remote tool", &value, MAX_REMOTE_TOOL_BYTES)?; + if value.is_empty() { + return Err("remote tool is empty".to_string()); + } + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteBuild { + cwd: WorkspaceRelativePath, + tool: RemoteTool, + argv: Vec, + env: Vec<(String, String)>, +} + +impl RemoteBuild { + pub fn new(cwd: WorkspaceRelativePath, tool: RemoteTool, argv: Vec, env: Vec<(String, String)>) -> Result { + validate_count(argv.len(), MAX_REMOTE_ARG_COUNT, "remote argv")?; + argv.iter().try_for_each(|arg| validate_string("remote argument", arg, MAX_REMOTE_ARG_BYTES))?; + validate_count(env.len(), MAX_REMOTE_ENV_COUNT, "remote environment")?; + env.iter().try_for_each(|(key, value)| { + validate_string("remote environment key", key, MAX_REMOTE_ENV_KEY_BYTES)?; + if key.is_empty() || key.contains('=') { + return Err("remote environment key is invalid".to_string()); + } + validate_string("remote environment value", value, MAX_REMOTE_ENV_VALUE_BYTES) + })?; + Ok(Self { cwd, tool, argv, env }) + } + + pub fn cwd(&self) -> &WorkspaceRelativePath { + &self.cwd + } + + pub fn tool(&self) -> &RemoteTool { + &self.tool + } + + pub fn argv(&self) -> &[String] { + &self.argv + } + + pub fn env(&self) -> &[(String, String)] { + &self.env + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RemoteOperation { + Sync, + Build(RemoteBuild), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteRequest { + request_id: RequestId, + workspace_session_id: WorkspaceSessionId, + operation: RemoteOperation, +} + +impl RemoteRequest { + pub fn sync(request_id: RequestId, workspace_session_id: WorkspaceSessionId) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Sync } + } + + pub fn build(request_id: RequestId, workspace_session_id: WorkspaceSessionId, build: RemoteBuild) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Build(build) } + } + + pub fn request_id(&self) -> RequestId { + self.request_id + } + + pub fn workspace_session_id(&self) -> WorkspaceSessionId { + self.workspace_session_id + } + + pub fn operation(&self) -> &RemoteOperation { + &self.operation + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteExecutionContext { + pub target: RemoteTargetId, + pub workspace_session_id: WorkspaceSessionId, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AuthorizedRemoteRequest { + request: RemoteRequest, + target: RemoteTargetId, +} + +impl AuthorizedRemoteRequest { + pub fn request_id(&self) -> RequestId { + self.request.request_id + } + + pub fn target(&self) -> RemoteTargetId { + self.target + } + + pub fn request(&self) -> &RemoteRequest { + &self.request + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RemoteAuthorizationError { + SessionMismatch, + TargetNotAllowed, + ToolNotAllowed(String), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteAuthorizationPolicy { + allowed_target: RemoteTargetId, + allowed_session: WorkspaceSessionId, + allowed_tools: Vec, +} + +impl RemoteAuthorizationPolicy { + pub fn new(allowed_target: RemoteTargetId, allowed_session: WorkspaceSessionId, allowed_tools: Vec) -> Self { + Self { allowed_target, allowed_session, allowed_tools } + } + + pub fn authorize(&self, context: &RemoteExecutionContext, request: RemoteRequest) -> Result { + if context.workspace_session_id != self.allowed_session || request.workspace_session_id != context.workspace_session_id { + return Err(RemoteAuthorizationError::SessionMismatch); + } + if context.target != self.allowed_target { + return Err(RemoteAuthorizationError::TargetNotAllowed); + } + + if let RemoteOperation::Build(build) = request.operation() { + if !self.allowed_tools.iter().any(|tool| tool == build.tool().as_str()) { + return Err(RemoteAuthorizationError::ToolNotAllowed(build.tool().as_str().to_string())); + } + } + + Ok(AuthorizedRemoteRequest { request, target: context.target }) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RemoteBackendEvent { + SyncProgress { completed_bytes: u64, total_bytes: Option }, + Stdout(Vec), + Stderr(Vec), + Error { message: String }, + Cancelled, + Completed { exit_code: i32 }, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RemoteBackendError { + Failed(String), + Spawn(String), + Timeout, + Cancelled, +} + +impl RemoteBackendError { + pub fn event(&self) -> RemoteBackendEvent { + match self { + Self::Failed(message) | Self::Spawn(message) => RemoteBackendEvent::Error { message: message.clone() }, + Self::Timeout => RemoteBackendEvent::Error { message: "remote backend timed out".to_string() }, + Self::Cancelled => RemoteBackendEvent::Cancelled, + } + } +} + +pub type RemoteFuture<'a, T> = Pin + Send + 'a>>; + +pub trait RemoteBackend: Send + Sync { + fn execute<'a>( + &'a self, request: AuthorizedRemoteRequest, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>>; +} + +fn validate_string(field: &str, value: &str, max: usize) -> Result<(), String> { + if value.as_bytes().contains(&0) { + return Err(format!("{field} contains a NUL byte")); + } + if value.len() > max { + return Err(format!("{field} exceeds maximum length {max}")); + } + Ok(()) +} + +fn validate_count(count: usize, max: usize, field: &str) -> Result<(), String> { + if count > max { + return Err(format!("{field} exceeds maximum count {max}")); + } + Ok(()) +} + +#[cfg(test)] +#[path = "remote_ut.rs"] +mod tests; diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index be54f54..71f3474 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -5,6 +5,7 @@ use std::ffi::OsStr; use std::io::{self, Read, Write}; use std::path::Path; +use crate::remote as remote_domain; #[cfg(unix)] use std::os::unix::ffi::OsStrExt; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; @@ -189,6 +190,20 @@ impl RemoteRequest { reader.finish()?; Ok(request) } + + pub fn into_domain(self) -> Result { + let request_id = remote_domain::RequestId(self.request_id.0); + let session_id = remote_domain::WorkspaceSessionId(self.workspace_session_id.0); + match self.operation { + RemoteOperation::Sync(_) => Ok(remote_domain::RemoteRequest::sync(request_id, session_id)), + RemoteOperation::Build(build) => { + let cwd = remote_domain::WorkspaceRelativePath::new(build.cwd.as_str())?; + let tool = remote_domain::RemoteTool::new(build.tool.as_str())?; + let build = remote_domain::RemoteBuild::new(cwd, tool, build.argv, build.env)?; + Ok(remote_domain::RemoteRequest::build(request_id, session_id, build)) + } + } + } } fn encode_remote_build(writer: &mut WireWriter, build: &RemoteBuild) -> Result<(), String> { @@ -440,6 +455,21 @@ pub struct RemoteEvent { } impl RemoteEvent { + pub fn from_backend_event(request_id: remote_domain::RequestId, event: remote_domain::RemoteBackendEvent) -> Self { + let request_id = RequestId(request_id.0); + let kind = match event { + remote_domain::RemoteBackendEvent::SyncProgress { completed_bytes, total_bytes } => { + RemoteEventKind::SyncProgress { completed_bytes, total_bytes } + } + remote_domain::RemoteBackendEvent::Stdout(data) => RemoteEventKind::Stdout(data), + remote_domain::RemoteBackendEvent::Stderr(data) => RemoteEventKind::Stderr(data), + remote_domain::RemoteBackendEvent::Error { message } => RemoteEventKind::Error { code: RemoteErrorCode::Failed, message }, + remote_domain::RemoteBackendEvent::Cancelled => RemoteEventKind::Cancelled, + remote_domain::RemoteBackendEvent::Completed { exit_code } => RemoteEventKind::Completed { exit_code }, + }; + Self { request_id, kind } + } + pub fn to_frame(&self) -> Result { let mut writer = WireWriter::new(*b"BBE1"); writer.u16(REMOTE_PROTOCOL_VERSION); From ffecebc7505a1ea5879be0c8a1d6b16befc255a4 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 17:54:08 +0200 Subject: [PATCH 09/52] Add Unit tests for the remote auth --- src/remote_ut.rs | 43 +++++++++++++++++++++++++++++++++++++++++++ src/vscomm/mod_ut.rs | 24 ++++++++++++++++++++++++ 2 files changed, 67 insertions(+) create mode 100644 src/remote_ut.rs diff --git a/src/remote_ut.rs b/src/remote_ut.rs new file mode 100644 index 0000000..4066e22 --- /dev/null +++ b/src/remote_ut.rs @@ -0,0 +1,43 @@ +use super::*; + +fn request(tool: &str) -> RemoteRequest { + RemoteRequest::build( + RequestId([1; 16]), + WorkspaceSessionId([2; 16]), + RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new(tool).unwrap(), + vec!["build".into()], + vec![("MODE".into(), "debug".into())], + ) + .unwrap(), + ) +} + +fn context() -> RemoteExecutionContext { + RemoteExecutionContext { target: RemoteTargetId([3; 16]), workspace_session_id: WorkspaceSessionId([2; 16]) } +} + +#[test] +fn policy_authorizes_typed_request() { + let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let authorized = policy.authorize(&context(), request("make")).unwrap(); + + assert_eq!(authorized.request_id(), RequestId([1; 16])); + assert_eq!(authorized.target(), RemoteTargetId([3; 16])); + assert_eq!(authorized.request().workspace_session_id, WorkspaceSessionId([2; 16])); +} + +#[test] +fn policy_rejects_unapproved_tool() { + let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["cargo".into()]); + + assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::ToolNotAllowed("make".into()))); +} + +#[test] +fn backend_errors_have_typed_events() { + assert_eq!(RemoteBackendError::Spawn("could not start".into()).event(), RemoteBackendEvent::Error { message: "could not start".into() }); + assert_eq!(RemoteBackendError::Timeout.event(), RemoteBackendEvent::Error { message: "remote backend timed out".into() }); + assert_eq!(RemoteBackendError::Cancelled.event(), RemoteBackendEvent::Cancelled); +} diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index c9874e2..a4bac24 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -199,3 +199,27 @@ fn existing_exec_request_wire_format_is_unchanged() { assert_eq!(decoded.args, request.args); assert_eq!(decoded.env, request.env); } + +#[test] +fn protocol_request_converts_to_transport_independent_domain_request() { + let request = RemoteRequest::build(ids().0, ids().1, build(vec!["--release".into()], vec![("MODE".into(), "debug".into())]).unwrap()); + let domain = request.into_domain().unwrap(); + + assert_eq!(domain.request_id(), crate::remote::RequestId([1; 16])); + let crate::remote::RemoteOperation::Build(build) = domain.operation() else { panic!("expected build") }; + assert_eq!(build.cwd().as_str(), "src"); + assert_eq!(build.tool().as_str(), "make"); + assert_eq!(build.argv(), ["--release"]); + assert_eq!(build.env(), [("MODE".into(), "debug".into())]); +} + +#[test] +fn backend_events_convert_to_protocol_events_without_transport_in_backend() { + let request_id = crate::remote::RequestId([7; 16]); + let event = RemoteEvent::from_backend_event(request_id, crate::remote::RemoteBackendEvent::Stdout(b"out".to_vec())); + assert_eq!(event.request_id, RequestId([7; 16])); + assert_eq!(event.kind, RemoteEventKind::Stdout(b"out".to_vec())); + + let event = RemoteEvent::from_backend_event(request_id, crate::remote::RemoteBackendEvent::Completed { exit_code: 3 }); + assert_eq!(event.kind, RemoteEventKind::Completed { exit_code: 3 }); +} From e2c08137c1fa06ecb09b42052742104f03d2a8f3 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 21:02:26 +0200 Subject: [PATCH 10/52] Add fake remote backend and tests --- src/bin/bunkerbox-vscomm.rs | 36 +++++++++++++++++- src/bunkerbox-vscomm_ut.rs | 75 ++++++++++++++++++++++++++++++++++++- src/daemon.rs | 64 +++++++++++++++++++++++++------ src/remote.rs | 6 +++ 4 files changed, 168 insertions(+), 13 deletions(-) diff --git a/src/bin/bunkerbox-vscomm.rs b/src/bin/bunkerbox-vscomm.rs index 77fff63..9ea316e 100644 --- a/src/bin/bunkerbox-vscomm.rs +++ b/src/bin/bunkerbox-vscomm.rs @@ -10,7 +10,10 @@ use std::mem; use std::os::unix::fs::PermissionsExt; use std::path::{Path, PathBuf}; -use vscomm::{encode_ui_payload, validate_exec_request, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT, TUI_STATUS_PORT, VSCOMM_BIN_DIR}; +use vscomm::{ + encode_ui_payload, validate_exec_request, ExecRequest, Frame, FrameType, RemoteBuild, RemoteEvent, RemoteEventKind, RemoteRequest, RemoteTool, + RequestId, WorkspaceRelativePath, WorkspaceSessionId, TOOLCHAIN_PORT, TUI_STATUS_PORT, VSCOMM_BIN_DIR, +}; const HOST_CID: u32 = 2; @@ -85,6 +88,37 @@ fn handle_response_to(response: Frame, stdout: &mut WO } } +pub fn remote_sync_request(request_id: RequestId, session_id: WorkspaceSessionId) -> RemoteRequest { + RemoteRequest::sync(request_id, session_id) +} + +pub fn remote_build_request( + request_id: RequestId, session_id: WorkspaceSessionId, cwd: impl Into, tool: impl Into, argv: Vec, + env: Vec<(String, String)>, +) -> Result { + let build = RemoteBuild::new(WorkspaceRelativePath::new(cwd)?, RemoteTool::new(tool)?, argv, env)?; + Ok(RemoteRequest::build(request_id, session_id, build)) +} + +pub fn execute_remote_request_to( + stream: &mut S, request: RemoteRequest, stdout: &mut WOut, stderr: &mut WErr, +) -> Result { + request.to_frame()?.write(stream).map_err(|e| format!("send remote request: {e}"))?; + + loop { + let frame = Frame::read(stream).map_err(|e| format!("read remote event: {e}"))?; + let event = RemoteEvent::from_frame(frame)?; + match event.kind { + RemoteEventKind::SyncProgress { .. } => {} + RemoteEventKind::Stdout(data) => stdout.write_all(&data).map_err(|e| format!("stdout: {e}"))?, + RemoteEventKind::Stderr(data) => stderr.write_all(&data).map_err(|e| format!("stderr: {e}"))?, + RemoteEventKind::Error { message, .. } => return Err(message), + RemoteEventKind::Cancelled => return Err("remote operation cancelled".to_string()), + RemoteEventKind::Completed { exit_code } => return Ok(exit_code), + } + } +} + fn notify_tui_error(message: &str) { let Ok(mut stream) = vsock_connect(HOST_CID, TUI_STATUS_PORT) else { return; diff --git a/src/bunkerbox-vscomm_ut.rs b/src/bunkerbox-vscomm_ut.rs index 9ddee4e..71a4d76 100644 --- a/src/bunkerbox-vscomm_ut.rs +++ b/src/bunkerbox-vscomm_ut.rs @@ -1,4 +1,38 @@ -use super::{handle_response, handle_response_to, Frame, FrameType}; +use super::vscomm::{RemoteEvent, RemoteEventKind, RequestId, WorkspaceSessionId}; +use super::{execute_remote_request_to, handle_response, handle_response_to, remote_build_request, Frame, FrameType}; +use std::io::{self, Read, Write}; + +struct MemoryStream { + input: io::Cursor>, + output: Vec, +} + +impl MemoryStream { + fn new(frames: Vec) -> Self { + let mut input = Vec::new(); + for frame in frames { + frame.write(&mut input).unwrap(); + } + Self { input: io::Cursor::new(input), output: Vec::new() } + } +} + +impl Read for MemoryStream { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + self.input.read(buf) + } +} + +impl Write for MemoryStream { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.output.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} #[test] fn forward_stdout_and_stderr_frames() { @@ -15,3 +49,42 @@ fn forward_stdout_and_stderr_frames() { fn preserve_exit_status() { assert_eq!(handle_response(Frame::new(FrameType::Exit, (-17i32).to_le_bytes().to_vec())).unwrap(), Some(-17)); } + +#[test] +fn explicit_remote_client_preserves_streams_status_and_request_id() { + let request_id = RequestId([9; 16]); + let request = remote_build_request(request_id, WorkspaceSessionId([8; 16]), "src", "make", vec!["release mode".into()], vec![]).unwrap(); + let responses = vec![ + RemoteEvent { request_id, kind: RemoteEventKind::Stdout(b"out".to_vec()) }.to_frame().unwrap(), + RemoteEvent { request_id, kind: RemoteEventKind::Stderr(b"err".to_vec()) }.to_frame().unwrap(), + RemoteEvent { request_id, kind: RemoteEventKind::Completed { exit_code: 23 } }.to_frame().unwrap(), + ]; + let mut stream = MemoryStream::new(responses); + let mut stdout = Vec::new(); + let mut stderr = Vec::new(); + + assert_eq!(execute_remote_request_to(&mut stream, request, &mut stdout, &mut stderr).unwrap(), 23); + assert_eq!(stdout, b"out"); + assert_eq!(stderr, b"err"); + + let sent = Frame::read(&mut io::Cursor::new(stream.output)).unwrap(); + let decoded = super::vscomm::RemoteRequest::from_frame(sent).unwrap(); + assert_eq!(decoded.request_id, request_id); + let super::vscomm::RemoteOperation::Build(build) = decoded.operation else { panic!("expected build") }; + assert_eq!(build.cwd.as_str(), "src"); + assert_eq!(build.argv, ["release mode"]); +} + +#[test] +fn explicit_remote_client_returns_remote_failure_without_local_fallback() { + let request_id = RequestId([4; 16]); + let request = super::remote_sync_request(request_id, WorkspaceSessionId([5; 16])); + let response = RemoteEvent { + request_id, + kind: RemoteEventKind::Error { code: super::vscomm::RemoteErrorCode::Failed, message: "backend unavailable".into() }, + }; + let mut stream = MemoryStream::new(vec![response.to_frame().unwrap()]); + + let error = execute_remote_request_to(&mut stream, request, &mut Vec::new(), &mut Vec::new()).unwrap_err(); + assert_eq!(error, "backend unavailable"); +} diff --git a/src/daemon.rs b/src/daemon.rs index d95a5f9..e296d8b 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -58,7 +58,13 @@ impl RemoteBroker { } pub async fn dispatch(&self, request: RemoteRequest, events: tokio::sync::mpsc::Sender) -> Result<(), RemoteDispatchError> { - let authorized = self.policy.authorize(&self.context, request).map_err(RemoteDispatchError::Unauthorized)?; + let authorized = match self.policy.authorize(&self.context, request) { + Ok(authorized) => authorized, + Err(error) => { + events.send(error.event()).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; + return Err(RemoteDispatchError::Unauthorized(error)); + } + }; match self.backend.execute(authorized, events.clone()).await { Ok(()) => Ok(()), Err(error) => { @@ -69,12 +75,23 @@ impl RemoteBroker { } } +struct UnavailableRemoteBackend; + +impl RemoteBackend for UnavailableRemoteBackend { + fn execute<'a>( + &'a self, _request: crate::remote::AuthorizedRemoteRequest, _events: tokio::sync::mpsc::Sender, + ) -> crate::remote::RemoteFuture<'a, Result<(), RemoteBackendError>> { + Box::pin(async { Err(RemoteBackendError::Failed("remote backend is unavailable".to_string())) }) + } +} + struct VsockSession { passthrough: Arc>, env_mode: EnvMode, workspace: PathBuf, merged_profile: Option>, proxy_config: Option>, + remote_broker: Arc, } pub struct VsockDaemon { @@ -112,6 +129,7 @@ impl VsockDaemon { if merged_profile.is_some() && !allow.is_empty() { let rt = tokio::runtime::Handle::current(); + let netrelay_path = find_netrelay_binary()?; let netrelay_path = find_netrelay_binary()?; let dir = make_proxy_runtime_dir()?; @@ -125,12 +143,23 @@ impl VsockDaemon { proxy_config = Some(SandboxProxyConfig { socket_path, netrelay_path }); } + let remote_target = crate::remote::RemoteTargetId([0; 16]); + let remote_session = crate::remote::WorkspaceSessionId([0; 16]); + let remote_policy = RemoteAuthorizationPolicy::new(remote_target, remote_session, Vec::new()); + let remote_backend: Arc = Arc::new(UnavailableRemoteBackend); + let remote_broker = Arc::new(RemoteBroker::new( + remote_policy, + RemoteExecutionContext { target: remote_target, workspace_session_id: remote_session }, + remote_backend, + )); + let session = Arc::new(VsockSession { passthrough: Arc::new(passthrough), env_mode, workspace, merged_profile, proxy_config: proxy_config.map(Arc::new), + remote_broker, }); let listener = tokio_vsock::VsockListener::bind(tokio_vsock::VsockAddr::new(libc::VMADDR_CID_ANY, TOOLCHAIN_PORT)).map_err(|e| { @@ -197,7 +226,14 @@ async fn daemon_loop( async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSession) -> Result<(), String> { let (mut reader, mut writer) = tokio::io::split(stream); - let req = read_exec_request(&mut reader).await?; + let frame = Frame::read_async(&mut reader).await.map_err(|e| format!("read frame: {e}"))?; + if matches!(frame.frame_type, FrameType::RemoteRequest) { + return dispatch_remote_frame(frame, &session.remote_broker, &mut writer).await; + } + if !matches!(frame.frame_type, FrameType::ExecReq) { + return Err(format!("expected ExecReq or RemoteRequest, got {:?}", frame.frame_type as u16)); + } + let req = ExecRequest::deserialize(&frame.payload)?; if let Err(err) = validate_exec_request(&req) { logging::diagnostic(&format!("bunkerbox-vscomm: invalid request: {err}")); @@ -221,6 +257,21 @@ async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSess Ok(()) } +pub async fn dispatch_remote_frame(frame: Frame, broker: &RemoteBroker, writer: &mut W) -> Result<(), String> { + let request = crate::vscomm::RemoteRequest::from_frame(frame).map_err(|err| format!("decode remote request: {err}"))?.into_domain()?; + let request_id = request.request_id(); + let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(64); + let dispatch_result = broker.dispatch(request, event_tx).await; + + while let Some(event) = event_rx.recv().await { + let response = + crate::vscomm::RemoteEvent::from_backend_event(request_id, event).to_frame().map_err(|err| format!("encode remote event: {err}"))?; + write_frame(writer, &response).await?; + } + + dispatch_result.map_err(|err| format!("remote dispatch failed: {err:?}")) +} + async fn execute_request(writer: &mut W, session: &VsockSession, req: &ExecRequest) -> Result<(), String> { validate_exec_request(req)?; let cwd = WorkspaceCwd::resolve(&session.workspace, Path::new(&req.cwd))?; @@ -480,15 +531,6 @@ fn is_allowed(passthrough: &[String], command: &str, args: &[String]) -> bool { false } -async fn read_exec_request(reader: &mut R) -> Result { - let frame = Frame::read_async(reader).await.map_err(|e| format!("read frame: {e}"))?; - if !matches!(frame.frame_type, FrameType::ExecReq) { - return Err(format!("expected ExecReq, got {:?}", frame.frame_type as u16)); - } - - ExecRequest::deserialize(&frame.payload) -} - fn bwrap_status_pipe() -> Result<(File, File), String> { let mut fds = [-1; 2]; if unsafe { libc::pipe(fds.as_mut_ptr()) } != 0 { diff --git a/src/remote.rs b/src/remote.rs index 4cdfb6c..eef8337 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -170,6 +170,12 @@ pub enum RemoteAuthorizationError { ToolNotAllowed(String), } +impl RemoteAuthorizationError { + pub fn event(&self) -> RemoteBackendEvent { + RemoteBackendEvent::Error { message: format!("remote authorization rejected: {self:?}") } + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct RemoteAuthorizationPolicy { allowed_target: RemoteTargetId, From bc1f3a8aa3d968f5450791e50c45dbbea84e485a Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 21:10:25 +0200 Subject: [PATCH 11/52] Add concurrent backend streaming --- Cargo.toml | 2 +- src/bin/bunkerbox-vscomm.rs | 14 +++++++++++-- src/bunkerbox-vscomm_ut.rs | 39 +++++++++++++++++++++++++++++++++---- src/daemon.rs | 31 ++++++++++++++++++++++------- 4 files changed, 72 insertions(+), 14 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 4ccabc4..ea873e0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -16,7 +16,7 @@ serde = { version = "1", features = ["derive"] } serde_json = "1" serde_yaml = "0.9" sha2 = "0.10" -tokio = { version = "1", features = ["rt-multi-thread", "net", "macros", "sync", "io-util", "process"] } +tokio = { version = "1", features = ["rt-multi-thread", "net", "macros", "sync", "io-util", "process", "time"] } tokio-vsock = "0.7" ratatui = "0.30" crossterm = "0.28" diff --git a/src/bin/bunkerbox-vscomm.rs b/src/bin/bunkerbox-vscomm.rs index 9ea316e..6382fae 100644 --- a/src/bin/bunkerbox-vscomm.rs +++ b/src/bin/bunkerbox-vscomm.rs @@ -103,15 +103,25 @@ pub fn remote_build_request( pub fn execute_remote_request_to( stream: &mut S, request: RemoteRequest, stdout: &mut WOut, stderr: &mut WErr, ) -> Result { + let request_id = request.request_id; request.to_frame()?.write(stream).map_err(|e| format!("send remote request: {e}"))?; loop { let frame = Frame::read(stream).map_err(|e| format!("read remote event: {e}"))?; let event = RemoteEvent::from_frame(frame)?; + if event.request_id != request_id { + return Err("remote event request ID mismatch".to_string()); + } match event.kind { RemoteEventKind::SyncProgress { .. } => {} - RemoteEventKind::Stdout(data) => stdout.write_all(&data).map_err(|e| format!("stdout: {e}"))?, - RemoteEventKind::Stderr(data) => stderr.write_all(&data).map_err(|e| format!("stderr: {e}"))?, + RemoteEventKind::Stdout(data) => { + stdout.write_all(&data).map_err(|e| format!("stdout: {e}"))?; + stdout.flush().map_err(|e| format!("flush stdout: {e}"))?; + } + RemoteEventKind::Stderr(data) => { + stderr.write_all(&data).map_err(|e| format!("stderr: {e}"))?; + stderr.flush().map_err(|e| format!("flush stderr: {e}"))?; + } RemoteEventKind::Error { message, .. } => return Err(message), RemoteEventKind::Cancelled => return Err("remote operation cancelled".to_string()), RemoteEventKind::Completed { exit_code } => return Ok(exit_code), diff --git a/src/bunkerbox-vscomm_ut.rs b/src/bunkerbox-vscomm_ut.rs index 71a4d76..1b0c815 100644 --- a/src/bunkerbox-vscomm_ut.rs +++ b/src/bunkerbox-vscomm_ut.rs @@ -7,6 +7,23 @@ struct MemoryStream { output: Vec, } +struct FlushWriter { + bytes: Vec, + flushes: usize, +} + +impl Write for FlushWriter { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.bytes.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + self.flushes += 1; + Ok(()) + } +} + impl MemoryStream { fn new(frames: Vec) -> Self { let mut input = Vec::new(); @@ -60,12 +77,14 @@ fn explicit_remote_client_preserves_streams_status_and_request_id() { RemoteEvent { request_id, kind: RemoteEventKind::Completed { exit_code: 23 } }.to_frame().unwrap(), ]; let mut stream = MemoryStream::new(responses); - let mut stdout = Vec::new(); - let mut stderr = Vec::new(); + let mut stdout = FlushWriter { bytes: Vec::new(), flushes: 0 }; + let mut stderr = FlushWriter { bytes: Vec::new(), flushes: 0 }; assert_eq!(execute_remote_request_to(&mut stream, request, &mut stdout, &mut stderr).unwrap(), 23); - assert_eq!(stdout, b"out"); - assert_eq!(stderr, b"err"); + assert_eq!(stdout.bytes, b"out"); + assert_eq!(stderr.bytes, b"err"); + assert_eq!(stdout.flushes, 1); + assert_eq!(stderr.flushes, 1); let sent = Frame::read(&mut io::Cursor::new(stream.output)).unwrap(); let decoded = super::vscomm::RemoteRequest::from_frame(sent).unwrap(); @@ -88,3 +107,15 @@ fn explicit_remote_client_returns_remote_failure_without_local_fallback() { let error = execute_remote_request_to(&mut stream, request, &mut Vec::new(), &mut Vec::new()).unwrap_err(); assert_eq!(error, "backend unavailable"); } + +#[test] +fn explicit_remote_client_rejects_mismatched_request_id() { + let request_id = RequestId([4; 16]); + let response = RemoteEvent { request_id: RequestId([5; 16]), kind: RemoteEventKind::Completed { exit_code: 0 } }; + let mut stream = MemoryStream::new(vec![response.to_frame().unwrap()]); + + let error = + execute_remote_request_to(&mut stream, super::remote_sync_request(request_id, WorkspaceSessionId([5; 16])), &mut Vec::new(), &mut Vec::new()) + .unwrap_err(); + assert_eq!(error, "remote event request ID mismatch"); +} diff --git a/src/daemon.rs b/src/daemon.rs index e296d8b..46ce195 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -261,15 +261,32 @@ pub async fn dispatch_remote_frame(frame: Frame, broke let request = crate::vscomm::RemoteRequest::from_frame(frame).map_err(|err| format!("decode remote request: {err}"))?.into_domain()?; let request_id = request.request_id(); let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(64); - let dispatch_result = broker.dispatch(request, event_tx).await; + let mut dispatch = Box::pin(broker.dispatch(request, event_tx)); + let mut dispatch_result: Option> = None; - while let Some(event) = event_rx.recv().await { - let response = - crate::vscomm::RemoteEvent::from_backend_event(request_id, event).to_frame().map_err(|err| format!("encode remote event: {err}"))?; - write_frame(writer, &response).await?; - } + loop { + if dispatch_result.is_some() { + let Some(event) = event_rx.recv().await else { + let result = dispatch_result.take().expect("dispatch result is present"); + return result.map_err(|err| format!("remote dispatch failed: {err:?}")); + }; + let response = + crate::vscomm::RemoteEvent::from_backend_event(request_id, event).to_frame().map_err(|err| format!("encode remote event: {err}"))?; + write_frame(writer, &response).await?; + continue; + } - dispatch_result.map_err(|err| format!("remote dispatch failed: {err:?}")) + tokio::select! { + result = &mut dispatch => dispatch_result = Some(result), + event = event_rx.recv() => { + let Some(event) = event else { return Err("remote event stream closed before backend completion".to_string()) }; + let response = crate::vscomm::RemoteEvent::from_backend_event(request_id, event) + .to_frame() + .map_err(|err| format!("encode remote event: {err}"))?; + write_frame(writer, &response).await?; + } + } + } } async fn execute_request(writer: &mut W, session: &VsockSession, req: &ExecRequest) -> Result<(), String> { From d36d56721e7f06cd5e3e27661988a8fc1b265b69 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 22:41:32 +0200 Subject: [PATCH 12/52] Add unit tests for loopback remote harness --- src/bunkerbox-remote_ut.rs | 154 +++++++++++++++++++++ src/bunkerbox-vscomm_ut.rs | 17 +-- src/loopback_ut.rs | 116 ++++++++++++++++ src/main_ut.rs | 38 +++++- src/snapshot_ut.rs | 272 +++++++++++++++++++++++++++++++++++++ src/vscomm/mod_ut.rs | 14 ++ 6 files changed, 602 insertions(+), 9 deletions(-) create mode 100644 src/bunkerbox-remote_ut.rs create mode 100644 src/loopback_ut.rs create mode 100644 src/snapshot_ut.rs diff --git a/src/bunkerbox-remote_ut.rs b/src/bunkerbox-remote_ut.rs new file mode 100644 index 0000000..6b1604f --- /dev/null +++ b/src/bunkerbox-remote_ut.rs @@ -0,0 +1,154 @@ +use super::*; +use bunkerbox::vscomm::{Frame, RemoteEvent, RemoteEventKind}; + +struct MemoryStream { + input: io::Cursor>, + output: Vec, +} + +impl MemoryStream { + fn new(events: Vec) -> Self { + let mut input = Vec::new(); + for event in events { + event.to_frame().unwrap().write(&mut input).unwrap(); + } + Self { input: io::Cursor::new(input), output: Vec::new() } + } +} + +impl Read for MemoryStream { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + self.input.read(buf) + } +} + +impl Write for MemoryStream { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.output.extend_from_slice(buf); + Ok(buf.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[test] +fn parses_sync_command() { + assert_eq!(parse_command(&["sync".into()]), Ok(RemoteCommand::Sync)); +} + +#[test] +fn parses_build_tool_and_args_without_joining() { + assert_eq!( + parse_command(&["build".into(), "make".into(), "release mode".into(), "$(literal)".into()]), + Ok(RemoteCommand::Build { tool: "make".into(), args: vec!["release mode".into(), "$(literal)".into()] }) + ); +} + +#[test] +fn build_request_preserves_logical_cwd_and_arguments() { + let request = build_request( + RemoteCommand::Build { tool: "make".into(), args: vec!["release mode".into(), "$(literal)".into()] }, + "src".into(), + RequestId([1; 16]), + WorkspaceSessionId([2; 16]), + ) + .unwrap(); + let frame = request.to_frame().unwrap(); + let decoded = RemoteRequest::from_frame(frame).unwrap(); + let bunkerbox::vscomm::RemoteOperation::Build(build) = decoded.operation else { panic!("expected build") }; + assert_eq!(build.cwd.as_str(), "src"); + assert_eq!(build.tool.as_str(), "make"); + assert_eq!(build.argv, ["release mode", "$(literal)"]); +} + +#[test] +fn sync_success_uses_existing_remote_helper_and_returns_status() { + let request_id = RequestId([3; 16]); + let mut stream = MemoryStream::new(vec![RemoteEvent { request_id, kind: RemoteEventKind::Completed { exit_code: 0 } }]); + let status = execute_remote_request_to( + &mut stream, + build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16])).unwrap(), + &mut Vec::new(), + &mut Vec::new(), + ) + .unwrap(); + assert_eq!(status, 0); + assert!(Frame::read(&mut io::Cursor::new(stream.output)).is_ok()); +} + +#[test] +fn build_success_preserves_output_bytes_and_nonzero_exit_code() { + let request_id = RequestId([6; 16]); + let mut stream = MemoryStream::new(vec![ + RemoteEvent { request_id, kind: RemoteEventKind::Stdout(vec![b'o', b'\n', 0xff]) }, + RemoteEvent { request_id, kind: RemoteEventKind::Stderr(b"err\n".to_vec()) }, + RemoteEvent { request_id, kind: RemoteEventKind::Completed { exit_code: 17 } }, + ]); + let mut stdout = Vec::new(); + let mut stderr = Vec::new(); + + let status = execute_remote_request_to( + &mut stream, + build_request( + RemoteCommand::Build { tool: "make".into(), args: vec!["release".into()] }, + "src".into(), + request_id, + WorkspaceSessionId([2; 16]), + ) + .unwrap(), + &mut stdout, + &mut stderr, + ) + .unwrap(); + + assert_eq!(status, 17); + assert_eq!(stdout, vec![b'o', b'\n', 0xff]); + assert_eq!(stderr, b"err\n"); +} + +#[test] +fn remote_failures_return_errors_without_local_fallback() { + for kind in + [RemoteEventKind::Error { code: bunkerbox::vscomm::RemoteErrorCode::Failed, message: "unavailable".into() }, RemoteEventKind::Cancelled] + { + let request_id = RequestId([4; 16]); + let mut stream = MemoryStream::new(vec![RemoteEvent { request_id, kind }]); + let result = execute_remote_request_to( + &mut stream, + build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16])).unwrap(), + &mut Vec::new(), + &mut Vec::new(), + ); + assert!(result.is_err()); + } +} + +#[test] +fn rejected_tool_and_mismatched_response_fail_closed() { + let request_id = RequestId([7; 16]); + let rejected = RemoteEvent { + request_id, + kind: RemoteEventKind::Error { code: bunkerbox::vscomm::RemoteErrorCode::Failed, message: "authorization rejected".into() }, + }; + let mut stream = MemoryStream::new(vec![rejected]); + let request = build_request( + RemoteCommand::Build { tool: "cargo".into(), args: vec!["build".into()] }, + String::new(), + request_id, + WorkspaceSessionId([2; 16]), + ) + .unwrap(); + assert_eq!(execute_remote_request_to(&mut stream, request, &mut Vec::new(), &mut Vec::new()).unwrap_err(), "authorization rejected"); + + let mut stream = MemoryStream::new(vec![RemoteEvent { request_id: RequestId([8; 16]), kind: RemoteEventKind::Completed { exit_code: 0 } }]); + let request = build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16])).unwrap(); + assert_eq!(execute_remote_request_to(&mut stream, request, &mut Vec::new(), &mut Vec::new()).unwrap_err(), "remote event request ID mismatch"); +} + +#[test] +fn logical_cwd_is_workspace_relative_only() { + assert_eq!(logical_workspace_cwd(Path::new("/workspace/project/src")).unwrap(), "project/src"); + assert!(logical_workspace_cwd(Path::new("/tmp/project")).is_err()); +} diff --git a/src/bunkerbox-vscomm_ut.rs b/src/bunkerbox-vscomm_ut.rs index 1b0c815..b07d66c 100644 --- a/src/bunkerbox-vscomm_ut.rs +++ b/src/bunkerbox-vscomm_ut.rs @@ -1,5 +1,6 @@ -use super::vscomm::{RemoteEvent, RemoteEventKind, RequestId, WorkspaceSessionId}; -use super::{execute_remote_request_to, handle_response, handle_response_to, remote_build_request, Frame, FrameType}; +use super::{handle_response, handle_response_to, Frame, FrameType}; +use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request}; +use bunkerbox::vscomm::{Frame as RemoteFrame, RemoteEvent, RemoteEventKind, RequestId, WorkspaceSessionId}; use std::io::{self, Read, Write}; struct MemoryStream { @@ -25,7 +26,7 @@ impl Write for FlushWriter { } impl MemoryStream { - fn new(frames: Vec) -> Self { + fn new(frames: Vec) -> Self { let mut input = Vec::new(); for frame in frames { frame.write(&mut input).unwrap(); @@ -86,10 +87,10 @@ fn explicit_remote_client_preserves_streams_status_and_request_id() { assert_eq!(stdout.flushes, 1); assert_eq!(stderr.flushes, 1); - let sent = Frame::read(&mut io::Cursor::new(stream.output)).unwrap(); - let decoded = super::vscomm::RemoteRequest::from_frame(sent).unwrap(); + let sent = RemoteFrame::read(&mut io::Cursor::new(stream.output)).unwrap(); + let decoded = bunkerbox::vscomm::RemoteRequest::from_frame(sent).unwrap(); assert_eq!(decoded.request_id, request_id); - let super::vscomm::RemoteOperation::Build(build) = decoded.operation else { panic!("expected build") }; + let bunkerbox::vscomm::RemoteOperation::Build(build) = decoded.operation else { panic!("expected build") }; assert_eq!(build.cwd.as_str(), "src"); assert_eq!(build.argv, ["release mode"]); } @@ -97,10 +98,10 @@ fn explicit_remote_client_preserves_streams_status_and_request_id() { #[test] fn explicit_remote_client_returns_remote_failure_without_local_fallback() { let request_id = RequestId([4; 16]); - let request = super::remote_sync_request(request_id, WorkspaceSessionId([5; 16])); + let request = remote_sync_request(request_id, WorkspaceSessionId([5; 16])); let response = RemoteEvent { request_id, - kind: RemoteEventKind::Error { code: super::vscomm::RemoteErrorCode::Failed, message: "backend unavailable".into() }, + kind: RemoteEventKind::Error { code: bunkerbox::vscomm::RemoteErrorCode::Failed, message: "backend unavailable".into() }, }; let mut stream = MemoryStream::new(vec![response.to_frame().unwrap()]); diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs new file mode 100644 index 0000000..4da344f --- /dev/null +++ b/src/loopback_ut.rs @@ -0,0 +1,116 @@ +use super::*; +use crate::remote::{RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteTool, RequestId, WorkspaceRelativePath}; +use tempfile::TempDir; + +fn fixture() -> (TempDir, Arc, RemoteTargetId, WorkspaceSessionId) { + let temp = tempfile::tempdir().unwrap(); + let workspace = temp.path().join("workspace"); + fs::create_dir(&workspace).unwrap(); + fs::create_dir(workspace.join("src")).unwrap(); + fs::write(workspace.join("src/input.txt"), b"snapshot input\n").unwrap(); + let session_id = WorkspaceSessionId([1; 16]); + let target = RemoteTargetId([2; 16]); + let snapshot_store = SnapshotStore::new(temp.path().join("snapshots")); + let exclusions = crate::snapshot::SnapshotExclusionPolicy::from_patterns(Vec::::new()).unwrap(); + let builder = SnapshotBuilder::new(snapshot_store.clone(), crate::snapshot::SnapshotLimits::default(), exclusions); + let session = Arc::new(RunRemoteSession::new(session_id, target, workspace, snapshot_store, builder, temp.path().join("jobs")).unwrap()); + (temp, session, target, session_id) +} + +fn authorized_build( + target: RemoteTargetId, session: WorkspaceSessionId, tool: &str, args: Vec, env: Vec<(String, String)>, +) -> crate::remote::AuthorizedRemoteRequest { + let build = RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new(tool).unwrap(), args, env).unwrap(); + let request = RemoteRequest::build(RequestId([3; 16]), session, build); + let policy = RemoteAuthorizationPolicy::new(target, session, vec![tool.to_string()]); + policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session }, request).unwrap() +} + +async fn collect_events(mut receiver: mpsc::Receiver) -> Vec { + let mut events = Vec::new(); + while let Some(event) = receiver.recv().await { + events.push(event); + } + events +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn sync_and_build_materialize_a_bound_snapshot() { + let (_temp, session, target, session_id) = fixture(); + let session_state = session.clone(); + let mut tools = BTreeMap::new(); + let printf = ["/usr/local/bin/printf", "/usr/bin/printf", "/bin/printf"].into_iter().map(PathBuf::from).find(|path| path.is_file()).unwrap(); + tools.insert("printf".to_string(), printf); + let backend = LoopbackBackend::new(session, tools); + + let (sync_tx, sync_rx) = mpsc::channel(8); + let sync = RemoteRequest::sync(RequestId([4; 16]), session_id); + let policy = RemoteAuthorizationPolicy::new(target, session_id, Vec::new()); + let authorized = policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, sync).unwrap(); + assert_eq!(backend.execute(authorized, sync_tx).await, Ok(())); + assert_eq!( + collect_events(sync_rx).await, + vec![RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }, RemoteBackendEvent::Completed { exit_code: 0 },] + ); + + let (build_tx, build_rx) = mpsc::channel(8); + let authorized = authorized_build(target, session_id, "printf", vec!["value with spaces:$(literal)".to_string()], Vec::new()); + assert_eq!(backend.execute(authorized, build_tx).await, Ok(())); + assert_eq!( + collect_events(build_rx).await, + vec![RemoteBackendEvent::Stdout(b"value with spaces:$(literal)".to_vec()), RemoteBackendEvent::Completed { exit_code: 0 },] + ); + assert!(fs::read_dir(&session_state.jobs_root).unwrap().next().is_none()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn nonempty_remote_environment_is_rejected() { + let (_temp, session, target, session_id) = fixture(); + let backend = LoopbackBackend::new(session, BTreeMap::new()); + let (events, _receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "printf", Vec::new(), vec![("UNTRUSTED".to_string(), "1".to_string())]); + assert_eq!( + backend.execute(request, events).await, + Err(RemoteBackendError::Failed("remote environment is unavailable until remote environment policy is configured".to_string(),)) + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn build_preserves_argument_bytes_and_reports_nonzero_stderr() { + let (_temp, session, target, session_id) = fixture(); + session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["ls".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session, tools); + let (events, receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "ls", vec!["$(not-a-shell-argument)".to_string()], Vec::new()); + assert_eq!(backend.execute(request, events).await, Ok(())); + let events = collect_events(receiver).await; + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Stderr(bytes) if !bytes.is_empty()))); + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code } if *exit_code != 0))); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn missing_tool_fails_before_execution() { + let (_temp, session, target, session_id) = fixture(); + let backend = LoopbackBackend::new(session, BTreeMap::new()); + let (events, _receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "missing-tool", Vec::new(), Vec::new()); + assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Spawn("loopback tool is not configured: missing-tool".to_string()))); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn timeout_kills_a_direct_child_process() { + let (_temp, session, target, session_id) = fixture(); + session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["sleep".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session, tools).with_timeout(Duration::from_millis(50)); + let (events, _receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "sleep", vec!["5".to_string()], Vec::new()); + assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Timeout)); +} diff --git a/src/main_ut.rs b/src/main_ut.rs index 730facf..b1e0af1 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -1,4 +1,7 @@ -use super::{decode_workspace_handoff, encode_workspace_handoff}; +use super::{decode_workspace_handoff, encode_workspace_handoff, read_run_handoff, remote_tool_names, write_run_handoff}; +use bunkerbox::vscomm::WorkspaceSessionId; +use std::fs::File; +use std::os::fd::FromRawFd; use std::path::Path; #[test] @@ -22,3 +25,36 @@ fn workspace_handoff_rejects_truncated_payload() { assert!(decode_workspace_handoff(&frame[..frame.len() - 1]).is_err()); } + +#[test] +fn run_handoff_round_trips_path_and_session() { + let (parent, child) = unsafe { + let mut fds = [-1; 2]; + assert_eq!(libc::pipe(fds.as_mut_ptr()), 0); + (File::from_raw_fd(fds[0]), File::from_raw_fd(fds[1])) + }; + let mut child = child; + write_run_handoff(&mut child, Path::new("/workspace/project"), WorkspaceSessionId([7; 16])).unwrap(); + drop(child); + let mut parent = parent; + let (path, session) = read_run_handoff(&mut parent).unwrap(); + assert_eq!(path, Path::new("/workspace/project")); + assert_eq!(session, WorkspaceSessionId([7; 16])); +} + +#[test] +fn run_handoff_rejects_zero_session() { + let mut payload = b"/workspace/project".to_vec(); + payload.push(0); + payload.extend_from_slice(&[0; 16]); + let frame = encode_workspace_handoff(&payload).unwrap(); + let path = tempfile::NamedTempFile::new().unwrap(); + std::fs::write(path.path(), frame).unwrap(); + let mut file = File::open(path.path()).unwrap(); + assert!(read_run_handoff(&mut file).is_err()); +} + +#[test] +fn remote_tool_names_reduce_passthrough_entries_to_executables() { + assert_eq!(remote_tool_names(&["make *".into(), "cargo build".into(), "make test".into()]), vec!["cargo", "make"]); +} diff --git a/src/snapshot_ut.rs b/src/snapshot_ut.rs new file mode 100644 index 0000000..173774e --- /dev/null +++ b/src/snapshot_ut.rs @@ -0,0 +1,272 @@ +use super::*; +use crate::cfg::{ProjectConfig, ProjectSection}; +use crate::remote::WorkspaceSessionId; +use std::fs; +use std::os::unix::fs::{symlink, PermissionsExt}; +use std::os::unix::net::UnixListener; +use std::path::Path; +use std::time::Duration; +use tempfile::TempDir; + +fn session(value: u8) -> WorkspaceSessionId { + WorkspaceSessionId([value; 16]) +} + +fn builder(store: &TempDir, limits: SnapshotLimits, patterns: &[&str]) -> SnapshotBuilder { + SnapshotBuilder::new( + SnapshotStore::new(store.path()), + limits, + SnapshotExclusionPolicy::from_patterns(patterns.iter().map(|pattern| (*pattern).to_string())).unwrap(), + ) +} + +fn build_at(source: &TempDir, store: &TempDir, limits: SnapshotLimits, patterns: &[&str]) -> Result { + builder(store, limits, patterns).build_root(source.path(), session(1)) +} + +fn write_file(root: &Path, path: &str, contents: &[u8]) { + let path = root.join(path); + fs::create_dir_all(path.parent().unwrap()).unwrap(); + fs::write(path, contents).unwrap(); +} + +#[test] +fn snapshots_nested_modified_and_untracked_files_with_modes() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "src/main.rs", b"modified"); + write_file(source.path(), "agent/new.txt", b"created"); + let executable = source.path().join("tool.sh"); + fs::write(&executable, b"#!/bin/sh\n").unwrap(); + fs::set_permissions(&executable, fs::Permissions::from_mode(0o755)).unwrap(); + + let snapshot = build_at(&source, &store, SnapshotLimits::default(), &[]).unwrap(); + let paths = snapshot.entries().iter().map(|entry| entry.path().as_str()).collect::>(); + assert_eq!(paths, vec!["agent", "agent/new.txt", "src", "src/main.rs", "tool.sh"]); + assert_eq!(snapshot.total_file_bytes(), 8 + 7 + 10); + let tool = snapshot.entries().iter().find(|entry| entry.path().as_str() == "tool.sh").unwrap(); + assert_eq!(tool.kind(), SnapshotEntryKind::RegularFile); + assert_eq!(tool.mode(), 0o755); + assert!(tool.content_digest().is_some()); + assert!(store.path().to_string_lossy().is_empty() || format!("{:?}", snapshot).contains("SnapshotId")); +} + +#[test] +fn identical_trees_have_deterministic_manifest_identity() { + let source_a = TempDir::new().unwrap(); + let source_b = TempDir::new().unwrap(); + let store_a = TempDir::new().unwrap(); + let store_b = TempDir::new().unwrap(); + for source in [&source_a, &source_b] { + write_file(source.path(), "b/file", b"same"); + write_file(source.path(), "a.txt", b"content"); + } + + let first = build_at(&source_a, &store_a, SnapshotLimits::default(), &[]).unwrap(); + let second = build_at(&source_b, &store_b, SnapshotLimits::default(), &[]).unwrap(); + assert_eq!(first.handle().snapshot_id(), second.handle().snapshot_id()); + assert_eq!(first.entries().iter().map(|entry| entry.path().as_str()).collect::>(), vec!["a.txt", "b", "b/file"]); +} + +#[test] +fn snapshot_store_resolves_only_the_bound_session() { + let source = TempDir::new().unwrap(); + let store_dir = TempDir::new().unwrap(); + write_file(source.path(), "file", b"data"); + let snapshot = build_at(&source, &store_dir, SnapshotLimits::default(), &[]).unwrap(); + let store = SnapshotStore::new(store_dir.path()); + assert_eq!(store.resolve(snapshot.handle()).unwrap(), snapshot); + + let wrong_session = SnapshotHandle { session_id: session(2), snapshot_id: snapshot.handle().snapshot_id() }; + assert!(store.resolve(&wrong_session).is_err()); + assert!(!format!("{:?}", snapshot.handle()).contains(&source.path().to_string_lossy().to_string())); +} + +#[test] +fn exclusions_prune_defaults_basenames_and_anchored_subtrees() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + for path in ["target/out", "nested/target/out", "docs/generated/file", ".git/config", ".bunkerbox/state", ".env", ".ssh/key"] { + write_file(source.path(), path, b"excluded"); + } + write_file(source.path(), "docs/keep/file", b"included"); + let snapshot = build_at(&source, &store, SnapshotLimits::default(), &["target/", "docs/generated/"]).unwrap(); + let paths = snapshot.entries().iter().map(|entry| entry.path().as_str()).collect::>(); + assert_eq!(paths, vec!["docs", "docs/keep", "docs/keep/file", "nested"]); +} + +#[test] +fn config_and_runtime_exclusions_use_explicit_snapshot_semantics() { + let config = ProjectConfig { project: ProjectSection { exclude: vec!["vendor/".into()], ..Default::default() }, ..Default::default() }; + let policy = SnapshotExclusionPolicy::from_config(&config, Some(&["generated/tree/".to_string()])).unwrap(); + assert!(policy.excludes("vendor/file")); + assert!(policy.excludes("generated/tree/file")); + assert!(!policy.excludes("vendorized/file")); +} + +#[test] +fn malformed_exclusions_are_rejected() { + for pattern in ["/absolute", "foo/../bar", "foo//bar", "foo/./bar", ""] { + assert!(SnapshotExclusionPolicy::from_patterns([pattern.to_string()]).is_err(), "{pattern}"); + } +} + +#[test] +fn every_symlink_is_rejected_without_following_it() { + let cases = ["internal-file", "internal-dir", "external", "dangling", "loop-a"]; + for case in cases { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "real/file", b"data"); + match case { + "internal-file" => symlink("real/file", source.path().join(case)).unwrap(), + "internal-dir" => symlink("real", source.path().join(case)).unwrap(), + "external" => { + let outside = TempDir::new().unwrap(); + symlink(outside.path(), source.path().join(case)).unwrap(); + } + "dangling" => symlink("missing", source.path().join(case)).unwrap(), + "loop-a" => { + symlink("loop-b", source.path().join("loop-a")).unwrap(); + symlink("loop-a", source.path().join("loop-b")).unwrap(); + } + _ => unreachable!(), + } + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err(), "{case}"); + } +} + +#[test] +fn special_files_and_hard_links_are_rejected() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "regular", b"data"); + fs::hard_link(source.path().join("regular"), source.path().join("alias")).unwrap(); + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err()); + + let fifo = source.path().join("pipe"); + let fifo_name = std::ffi::CString::new(fifo.as_os_str().as_bytes()).unwrap(); + assert_eq!(unsafe { libc::mkfifo(fifo_name.as_ptr(), 0o600) }, 0); + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err()); + + fs::remove_file(&fifo).unwrap(); + let socket_path = source.path().join("socket"); + let _listener = UnixListener::bind(&socket_path).unwrap(); + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err()); +} + +#[test] +fn trusted_limits_reject_entries_and_cleanup_staging() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "a", b"1234"); + write_file(source.path(), "b", b"5678"); + let limits = SnapshotLimits { max_entries: 1, ..SnapshotLimits::default() }; + assert!(build_at(&source, &store, limits, &[]).is_err()); + assert_eq!(fs::read_dir(store.path().join(".staging")).unwrap().count(), 0); + + let limits = SnapshotLimits { max_file_bytes: 3, ..SnapshotLimits::default() }; + assert!(build_at(&source, &store, limits, &[]).is_err()); + assert_eq!(fs::read_dir(store.path().join(".staging")).unwrap().count(), 0); + + let limits = SnapshotLimits { max_total_bytes: 5, ..SnapshotLimits::default() }; + assert!(build_at(&source, &store, limits, &[]).is_err()); + assert_eq!(fs::read_dir(store.path().join(".staging")).unwrap().count(), 0); +} + +#[test] +fn path_manifest_and_deadline_limits_are_enforced() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "long-name", b"data"); + let limits = SnapshotLimits { max_component_bytes: 4, ..SnapshotLimits::default() }; + assert!(build_at(&source, &store, limits, &[]).is_err()); + + let limits = SnapshotLimits { max_manifest_bytes: 1, ..SnapshotLimits::default() }; + assert!(build_at(&source, &store, limits, &[]).is_err()); + + let limits = SnapshotLimits { max_duration: Duration::from_nanos(1), ..SnapshotLimits::default() }; + assert!(build_at(&source, &store, limits, &[]).is_err()); +} + +#[test] +fn zero_session_is_not_accepted_as_snapshot_authority() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "file", b"data"); + assert!(builder(&store, SnapshotLimits::default(), &[]).build_root(source.path(), WorkspaceSessionId([0; 16])).is_err()); +} + +#[test] +fn unreadable_workspace_entry_fails_when_supported() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + let path = source.path().join("private"); + fs::write(&path, b"secret").unwrap(); + fs::set_permissions(&path, fs::Permissions::from_mode(0o000)).unwrap(); + let result = build_at(&source, &store, SnapshotLimits::default(), &[]); + fs::set_permissions(&path, fs::Permissions::from_mode(0o600)).unwrap(); + if unsafe { libc::geteuid() } != 0 { + assert!(result.is_err()); + } +} + +#[test] +fn source_path_is_not_in_snapshot_metadata() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "file", b"data"); + let snapshot = build_at(&source, &store, SnapshotLimits::default(), &[]).unwrap(); + let debug = format!("{snapshot:?}"); + assert!(!debug.contains(&source.path().to_string_lossy().to_string())); + assert!(!debug.contains(&store.path().to_string_lossy().to_string())); +} + +#[test] +fn materialization_recreates_nested_files_and_modes() { + let source = TempDir::new().unwrap(); + let store_dir = TempDir::new().unwrap(); + write_file(source.path(), "src/main.rs", b"fn main() {}\n"); + let executable = source.path().join("tool.sh"); + fs::write(&executable, b"#!/bin/sh\n").unwrap(); + fs::set_permissions(&executable, fs::Permissions::from_mode(0o755)).unwrap(); + let snapshot = build_at(&source, &store_dir, SnapshotLimits::default(), &[]).unwrap(); + let store = SnapshotStore::new(store_dir.path()); + let destination = store_dir.path().join("materialized"); + + let materialized = store.materialize(snapshot.handle(), &destination).unwrap(); + assert_eq!(materialized.root(), destination); + assert_eq!(fs::read(destination.join("src/main.rs")).unwrap(), b"fn main() {}\n"); + assert_eq!(fs::metadata(destination.join("tool.sh")).unwrap().permissions().mode() & 0o777, 0o755); +} + +#[test] +fn materialization_rejects_existing_destination_and_cleans_digest_failures() { + let source = TempDir::new().unwrap(); + let store_dir = TempDir::new().unwrap(); + write_file(source.path(), "file", b"contents"); + let snapshot = build_at(&source, &store_dir, SnapshotLimits::default(), &[]).unwrap(); + let store = SnapshotStore::new(store_dir.path()); + let existing = store_dir.path().join("existing"); + fs::create_dir(&existing).unwrap(); + assert!(store.materialize(snapshot.handle(), &existing).is_err()); + + let staged_file = store.snapshot_path(snapshot.handle()).join("files/file"); + fs::write(staged_file, b"tampered").unwrap(); + let destination = store_dir.path().join("failed"); + assert!(store.materialize(snapshot.handle(), &destination).is_err()); + assert!(!destination.exists()); +} + +#[test] +fn materialization_rejects_destination_symlink() { + let source = TempDir::new().unwrap(); + let store_dir = TempDir::new().unwrap(); + let outside = TempDir::new().unwrap(); + write_file(source.path(), "file", b"contents"); + let snapshot = build_at(&source, &store_dir, SnapshotLimits::default(), &[]).unwrap(); + let destination = store_dir.path().join("link"); + symlink(outside.path(), &destination).unwrap(); + assert!(SnapshotStore::new(store_dir.path()).materialize(snapshot.handle(), &destination).is_err()); + assert!(outside.path().read_dir().unwrap().next().is_none()); +} diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index a4bac24..00f0f7f 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -4,6 +4,20 @@ fn ids() -> (RequestId, WorkspaceSessionId) { (RequestId([1; 16]), WorkspaceSessionId([2; 16])) } +#[test] +fn workspace_session_hex_round_trips_without_accepting_zero() { + let session = WorkspaceSessionId([0xab; 16]); + assert_eq!(WorkspaceSessionId::from_hex(&session.to_hex()), Ok(session)); + assert!(WorkspaceSessionId::from_hex(&"0".repeat(32)).is_err()); +} + +#[test] +fn workspace_session_hex_rejects_non_ascii_without_panicking() { + let mut value = "0".repeat(30); + value.push('\u{00e9}'); + assert!(WorkspaceSessionId::from_hex(&value).is_err()); +} + fn build(argv: Vec, env: Vec<(String, String)>) -> Result { RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new("make").unwrap(), argv, env) } From b9c97f9a95ed8a9af0fcb1162fae8e451016dec7 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 22:41:49 +0200 Subject: [PATCH 13/52] Implement remote loopback harness --- src/bin/bunkerbox-remote.rs | 134 ++++++ src/bin/bunkerbox-vscomm.rs | 47 +- src/daemon.rs | 81 ++-- src/kata.rs | 12 +- src/lib.rs | 3 + src/loopback.rs | 383 ++++++++++++++++ src/main.rs | 258 +++++++---- src/remote_client.rs | 43 ++ src/snapshot.rs | 892 ++++++++++++++++++++++++++++++++++++ src/vscomm/mod.rs | 31 ++ 10 files changed, 1727 insertions(+), 157 deletions(-) create mode 100644 src/bin/bunkerbox-remote.rs create mode 100644 src/loopback.rs create mode 100644 src/remote_client.rs create mode 100644 src/snapshot.rs diff --git a/src/bin/bunkerbox-remote.rs b/src/bin/bunkerbox-remote.rs new file mode 100644 index 0000000..d8439e0 --- /dev/null +++ b/src/bin/bunkerbox-remote.rs @@ -0,0 +1,134 @@ +use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request}; +use bunkerbox::vscomm::{RemoteRequest, RequestId, WorkspaceSessionId, TOOLCHAIN_PORT}; +use rand::RngCore; +use std::env; +use std::io::{self, Read, Write}; +use std::mem; +use std::path::Path; + +const HOST_CID: u32 = 2; + +#[derive(Debug, PartialEq, Eq)] +enum RemoteCommand { + Sync, + Build { tool: String, args: Vec }, +} + +fn main() { + match run() { + Ok(code) => std::process::exit(code), + Err(error) => { + eprintln!("bunkerbox-remote: remote operation failed: {error}"); + std::process::exit(1); + } + } +} + +fn run() -> Result { + let args = env::args().skip(1).collect::>(); + let command = parse_command(&args)?; + match &command { + RemoteCommand::Sync => eprintln!("bunkerbox-remote: syncing"), + RemoteCommand::Build { tool, .. } => eprintln!("bunkerbox-remote: building {tool}"), + } + let cwd = logical_workspace_cwd(&env::current_dir().map_err(|error| format!("current directory: {error}"))?)?; + let session = env::var("BUNKERBOX_REMOTE_SESSION") + .map_err(|_| "BUNKERBOX_REMOTE_SESSION is missing".to_string()) + .and_then(|value| WorkspaceSessionId::from_hex(&value))?; + let request = build_request(command, cwd, new_request_id(), session)?; + let mut stream = connect_toolchain()?; + let mut stdout = io::stdout(); + let mut stderr = io::stderr(); + execute_remote_request_to(&mut stream, request, &mut stdout, &mut stderr) +} + +fn parse_command(args: &[String]) -> Result { + match args { + [command] if command == "sync" => Ok(RemoteCommand::Sync), + [command, tool, rest @ ..] if command == "build" && !tool.is_empty() => Ok(RemoteCommand::Build { tool: tool.clone(), args: rest.to_vec() }), + [command, ..] if command == "build" => Err("usage: bunkerbox-remote build [args...]".to_string()), + [] => Err("usage: bunkerbox-remote sync | build [args...]".to_string()), + _ => Err("usage: bunkerbox-remote sync | build [args...]".to_string()), + } +} + +fn build_request(command: RemoteCommand, cwd: String, request_id: RequestId, session_id: WorkspaceSessionId) -> Result { + match command { + RemoteCommand::Sync => Ok(remote_sync_request(request_id, session_id)), + RemoteCommand::Build { tool, args } => remote_build_request(request_id, session_id, cwd, tool, args, Vec::new()), + } +} + +fn logical_workspace_cwd(path: &Path) -> Result { + let relative = path.strip_prefix("/workspace").map_err(|_| "current directory must be under /workspace".to_string())?; + let value = relative.to_str().ok_or_else(|| "current directory is not valid UTF-8".to_string())?.to_string(); + bunkerbox::remote::WorkspaceRelativePath::new(&value)?; + Ok(value) +} + +fn new_request_id() -> RequestId { + let mut bytes = [0; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + RequestId(bytes) +} + +fn connect_toolchain() -> Result { + vsock_connect(HOST_CID, TOOLCHAIN_PORT).map_err(|error| format!("toolchain vsock connect: {error}")) +} + +fn vsock_connect(cid: u32, port: u32) -> io::Result { + unsafe { + let fd = libc::socket(libc::AF_VSOCK, libc::SOCK_STREAM, 0); + if fd < 0 { + return Err(io::Error::last_os_error()); + } + + let addr = libc::sockaddr_vm { svm_family: libc::AF_VSOCK as u16, svm_reserved1: 0, svm_port: port, svm_cid: cid, svm_zero: [0u8; 4] }; + let addr_ptr = &addr as *const libc::sockaddr_vm as *const libc::sockaddr; + let addr_len = mem::size_of::() as libc::socklen_t; + if libc::connect(fd, addr_ptr, addr_len) < 0 { + let error = io::Error::last_os_error(); + libc::close(fd); + return Err(error); + } + Ok(VsockStream { fd }) + } +} + +struct VsockStream { + fd: libc::c_int, +} + +impl Read for VsockStream { + fn read(&mut self, buf: &mut [u8]) -> io::Result { + let result = unsafe { libc::read(self.fd, buf.as_mut_ptr() as *mut libc::c_void, buf.len()) }; + if result < 0 { + return Err(io::Error::last_os_error()); + } + Ok(result as usize) + } +} + +impl Write for VsockStream { + fn write(&mut self, buf: &[u8]) -> io::Result { + let result = unsafe { libc::write(self.fd, buf.as_ptr() as *const libc::c_void, buf.len()) }; + if result < 0 { + return Err(io::Error::last_os_error()); + } + Ok(result as usize) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +impl Drop for VsockStream { + fn drop(&mut self) { + unsafe { libc::close(self.fd) }; + } +} + +#[cfg(test)] +#[path = "../bunkerbox-remote_ut.rs"] +mod tests; diff --git a/src/bin/bunkerbox-vscomm.rs b/src/bin/bunkerbox-vscomm.rs index 6382fae..aea40b5 100644 --- a/src/bin/bunkerbox-vscomm.rs +++ b/src/bin/bunkerbox-vscomm.rs @@ -10,10 +10,8 @@ use std::mem; use std::os::unix::fs::PermissionsExt; use std::path::{Path, PathBuf}; -use vscomm::{ - encode_ui_payload, validate_exec_request, ExecRequest, Frame, FrameType, RemoteBuild, RemoteEvent, RemoteEventKind, RemoteRequest, RemoteTool, - RequestId, WorkspaceRelativePath, WorkspaceSessionId, TOOLCHAIN_PORT, TUI_STATUS_PORT, VSCOMM_BIN_DIR, -}; +pub use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request}; +use vscomm::{encode_ui_payload, validate_exec_request, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT, TUI_STATUS_PORT, VSCOMM_BIN_DIR}; const HOST_CID: u32 = 2; @@ -88,47 +86,6 @@ fn handle_response_to(response: Frame, stdout: &mut WO } } -pub fn remote_sync_request(request_id: RequestId, session_id: WorkspaceSessionId) -> RemoteRequest { - RemoteRequest::sync(request_id, session_id) -} - -pub fn remote_build_request( - request_id: RequestId, session_id: WorkspaceSessionId, cwd: impl Into, tool: impl Into, argv: Vec, - env: Vec<(String, String)>, -) -> Result { - let build = RemoteBuild::new(WorkspaceRelativePath::new(cwd)?, RemoteTool::new(tool)?, argv, env)?; - Ok(RemoteRequest::build(request_id, session_id, build)) -} - -pub fn execute_remote_request_to( - stream: &mut S, request: RemoteRequest, stdout: &mut WOut, stderr: &mut WErr, -) -> Result { - let request_id = request.request_id; - request.to_frame()?.write(stream).map_err(|e| format!("send remote request: {e}"))?; - - loop { - let frame = Frame::read(stream).map_err(|e| format!("read remote event: {e}"))?; - let event = RemoteEvent::from_frame(frame)?; - if event.request_id != request_id { - return Err("remote event request ID mismatch".to_string()); - } - match event.kind { - RemoteEventKind::SyncProgress { .. } => {} - RemoteEventKind::Stdout(data) => { - stdout.write_all(&data).map_err(|e| format!("stdout: {e}"))?; - stdout.flush().map_err(|e| format!("flush stdout: {e}"))?; - } - RemoteEventKind::Stderr(data) => { - stderr.write_all(&data).map_err(|e| format!("stderr: {e}"))?; - stderr.flush().map_err(|e| format!("flush stderr: {e}"))?; - } - RemoteEventKind::Error { message, .. } => return Err(message), - RemoteEventKind::Cancelled => return Err("remote operation cancelled".to_string()), - RemoteEventKind::Completed { exit_code } => return Ok(exit_code), - } - } -} - fn notify_tui_error(message: &str) { let Ok(mut stream) = vsock_connect(HOST_CID, TUI_STATUS_PORT) else { return; diff --git a/src/daemon.rs b/src/daemon.rs index 46ce195..1d35fc2 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -1,5 +1,6 @@ use crate::cfg::EnvMode; use crate::logging; +use crate::loopback::{LoopbackBackend, RunRemoteSession}; use crate::proxy::{FilterProxy, UnixProxyHandle}; use crate::remote::{ RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionContext, RemoteRequest, @@ -14,7 +15,7 @@ use std::os::fd::{AsRawFd, FromRawFd, RawFd}; use std::os::unix::fs::DirBuilderExt; use std::path::{Path, PathBuf}; use std::process::Stdio; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::process::Command; @@ -75,16 +76,6 @@ impl RemoteBroker { } } -struct UnavailableRemoteBackend; - -impl RemoteBackend for UnavailableRemoteBackend { - fn execute<'a>( - &'a self, _request: crate::remote::AuthorizedRemoteRequest, _events: tokio::sync::mpsc::Sender, - ) -> crate::remote::RemoteFuture<'a, Result<(), RemoteBackendError>> { - Box::pin(async { Err(RemoteBackendError::Failed("remote backend is unavailable".to_string())) }) - } -} - struct VsockSession { passthrough: Arc>, env_mode: EnvMode, @@ -101,9 +92,42 @@ pub struct VsockDaemon { sandbox_proxy_dir: Option, } +pub struct RemoteDaemonConfig { + session: Arc, + allowed_tools: Vec, + tools: std::collections::BTreeMap, +} + +impl RemoteDaemonConfig { + pub fn new(session: Arc, allowed_tools: Vec, tools: std::collections::BTreeMap) -> Self { + Self { session, allowed_tools, tools } + } +} + +struct RemoteComponents { + context: RemoteExecutionContext, + policy: RemoteAuthorizationPolicy, + backend: Arc, +} + impl VsockDaemon { - pub fn start( + pub fn start_with_remote( + passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, + remote: RemoteDaemonConfig, + ) -> Result { + let remote_policy = RemoteAuthorizationPolicy::new(remote.session.target(), remote.session.session_id(), remote.allowed_tools); + let remote_context = RemoteExecutionContext { target: remote.session.target(), workspace_session_id: remote.session.session_id() }; + let remote_components = RemoteComponents { + context: remote_context, + policy: remote_policy, + backend: Arc::new(LoopbackBackend::new(remote.session, remote.tools)), + }; + Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, remote_components) + } + + fn start_inner( passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, + remote: RemoteComponents, ) -> Result { let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); @@ -129,7 +153,6 @@ impl VsockDaemon { if merged_profile.is_some() && !allow.is_empty() { let rt = tokio::runtime::Handle::current(); - let netrelay_path = find_netrelay_binary()?; let netrelay_path = find_netrelay_binary()?; let dir = make_proxy_runtime_dir()?; @@ -143,16 +166,8 @@ impl VsockDaemon { proxy_config = Some(SandboxProxyConfig { socket_path, netrelay_path }); } - let remote_target = crate::remote::RemoteTargetId([0; 16]); - let remote_session = crate::remote::WorkspaceSessionId([0; 16]); - let remote_policy = RemoteAuthorizationPolicy::new(remote_target, remote_session, Vec::new()); - let remote_backend: Arc = Arc::new(UnavailableRemoteBackend); - let remote_broker = Arc::new(RemoteBroker::new( - remote_policy, - RemoteExecutionContext { target: remote_target, workspace_session_id: remote_session }, - remote_backend, - )); - + let connections = Arc::new(Mutex::new(Vec::new())); + let remote_broker = Arc::new(RemoteBroker::new(remote.policy, remote.context, remote.backend)); let session = Arc::new(VsockSession { passthrough: Arc::new(passthrough), env_mode, @@ -162,18 +177,18 @@ impl VsockDaemon { remote_broker, }); - let listener = tokio_vsock::VsockListener::bind(tokio_vsock::VsockAddr::new(libc::VMADDR_CID_ANY, TOOLCHAIN_PORT)).map_err(|e| { + let listener = tokio_vsock::VsockListener::bind(tokio_vsock::VsockAddr::new(libc::VMADDR_CID_ANY, TOOLCHAIN_PORT)).map_err(|error| { if let Some(h) = sandbox_proxy.take() { h.stop(); } if let Some(d) = sandbox_proxy_dir.take() { let _ = std::fs::remove_dir_all(&d); } - format!("failed to bind toolchain vsock port {TOOLCHAIN_PORT}: {e}") + format!("failed to bind toolchain vsock port {TOOLCHAIN_PORT}: {error}") })?; let join_handle = tokio::spawn(async move { - let result = daemon_loop(session, listener, shutdown_rx).await; + let result = daemon_loop(session, listener, shutdown_rx, connections).await; if let Err(err) = result { logging::diagnostic(&format!("bunkerbox: vsock daemon: {err}")); } @@ -196,6 +211,7 @@ impl VsockDaemon { async fn daemon_loop( session: Arc, listener: tokio_vsock::VsockListener, mut shutdown_rx: tokio::sync::oneshot::Receiver<()>, + connections: Arc>>>, ) -> Result<(), String> { loop { tokio::select! { @@ -203,11 +219,14 @@ async fn daemon_loop( match result { Ok((stream, _peer)) => { let session = session.clone(); - tokio::spawn(async move { + let connection = tokio::spawn(async move { if let Err(err) = handle_connection(stream, &session).await { logging::diagnostic(&format!("bunkerbox: toolchain vsock session failed: {err}")); } }); + let mut active = connections.lock().map_err(|_| "connection task lock poisoned".to_string())?; + active.retain(|task| !task.is_finished()); + active.push(connection); } Err(e) => { logging::diagnostic(&format!("bunkerbox: vsock accept error: {e}")); @@ -220,6 +239,14 @@ async fn daemon_loop( } } + let tasks = connections.lock().map_err(|_| "connection task lock poisoned".to_string())?.drain(..).collect::>(); + for task in &tasks { + task.abort(); + } + for task in tasks { + let _ = task.await; + } + Ok(()) } diff --git a/src/kata.rs b/src/kata.rs index 2cac0da..2ac635d 100644 --- a/src/kata.rs +++ b/src/kata.rs @@ -1,6 +1,6 @@ use crate::cfg::{HomeMode, NetworkMode, RuntimeConfig}; +use crate::vscomm::WorkspaceSessionId; use crate::vscomm::TOOLCHAIN_PORT; -use crate::workspace::WorkspaceHandle; use aes_gcm::aead::consts::U12; use aes_gcm::aead::Aead; use aes_gcm::{Aes256Gcm, KeyInit, Nonce}; @@ -17,6 +17,11 @@ use std::path::{Path, PathBuf}; use std::process::{Command, Stdio}; use std::thread; +pub struct WorkspaceBinding<'a> { + pub path: &'a Path, + pub remote_session: WorkspaceSessionId, +} + const BRIDGE_SUBNET: &str = "10.247.0.0/24"; const BRIDGE_NAME: &str = "bunkerbox0"; @@ -40,7 +45,7 @@ fn cleanup_partial_session(session_dir: Option<&PathBuf>, home_path: Option<&Pat } pub fn run( - config: &RuntimeConfig, workspace: WorkspaceHandle, container_name: &str, _share_dir: &Path, app_name: &str, vsock_enabled: bool, + config: &RuntimeConfig, workspace: WorkspaceBinding<'_>, container_name: &str, _share_dir: &Path, app_name: &str, vsock_enabled: bool, _status_fd: RawFd, ) -> Result<(), String> { if !config.oci.is_file() { @@ -168,7 +173,7 @@ pub fn run( ensure_bridge_egress_firewall(config, resolv_conf.as_deref())?; } let resolv_conf_mount = resolv_conf.as_ref().map(|path| format!("type=bind,src={},dst=/etc/resolv.conf,options=rbind:ro", path.display())); - let workspace_mount = format!("type=bind,src={},dst=/workspace,options=rbind:rw", workspace.path().display()); + let workspace_mount = format!("type=bind,src={},dst=/workspace,options=rbind:rw", workspace.path.display()); let mut container_env = Vec::new(); let mut tools_mount: Option = None; let init_cmd = String::from("/bunkerbox-tools/init.sh"); @@ -191,6 +196,7 @@ pub fn run( if vsock_enabled { container_env.push(format!("BUNKERBOX_TOOLCHAIN_PORT={TOOLCHAIN_PORT}")); + container_env.push(format!("BUNKERBOX_REMOTE_SESSION={}", workspace.remote_session.to_hex())); } if let Some(ref cmds) = config.command { diff --git a/src/lib.rs b/src/lib.rs index b0e64d5..c7b0059 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -6,10 +6,13 @@ pub mod daemon; pub mod kata; pub mod logging; pub mod netrelay; +pub mod loopback; pub mod overlay; pub mod proxy; pub mod remote; +pub mod remote_client; pub mod sandbox; +pub mod snapshot; pub mod tui; pub mod vscomm; pub mod workspace; diff --git a/src/loopback.rs b/src/loopback.rs new file mode 100644 index 0000000..5c938e7 --- /dev/null +++ b/src/loopback.rs @@ -0,0 +1,383 @@ +use crate::remote::{ + AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteTargetId, WorkspaceSessionId, +}; +use crate::snapshot::{SnapshotBuilder, SnapshotHandle, SnapshotStore}; +use std::collections::BTreeMap; +use std::fs; +use std::io; +use std::os::unix::fs::{MetadataExt, PermissionsExt}; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; +use tokio::io::AsyncReadExt; +use tokio::process::Command; +use tokio::sync::mpsc; +use tokio::time::sleep; + +pub const LOOPBACK_BUILD_TIMEOUT: Duration = Duration::from_secs(30); +const LOOPBACK_PATH: &str = "/usr/local/bin:/usr/bin:/bin"; +static NEXT_JOB_ID: AtomicU64 = AtomicU64::new(1); + +pub struct RunRemoteSession { + session_id: WorkspaceSessionId, + target: RemoteTargetId, + workspace_root: PathBuf, + snapshot_store: SnapshotStore, + snapshot_builder: SnapshotBuilder, + current_snapshot: Mutex>, + snapshot_operation: Mutex<()>, + jobs_root: PathBuf, +} + +impl RunRemoteSession { + pub fn new( + session_id: WorkspaceSessionId, target: RemoteTargetId, workspace_root: PathBuf, snapshot_store: SnapshotStore, + snapshot_builder: SnapshotBuilder, jobs_root: PathBuf, + ) -> Result { + if session_id.0 == [0; 16] { + return Err("remote session ID must be nonzero".to_string()); + } + if target.0 == [0; 16] { + return Err("remote target ID must be nonzero".to_string()); + } + let snapshot_root = snapshot_store.root_for_cleanup(); + create_private_root(&snapshot_root, "snapshot store")?; + if let Err(error) = create_private_root(&jobs_root, "loopback jobs") { + let _ = fs::remove_dir_all(&snapshot_root); + return Err(error); + } + Ok(Self { + session_id, + target, + workspace_root, + snapshot_store, + snapshot_builder, + current_snapshot: Mutex::new(None), + snapshot_operation: Mutex::new(()), + jobs_root, + }) + } + + pub fn session_id(&self) -> WorkspaceSessionId { + self.session_id + } + + pub fn target(&self) -> RemoteTargetId { + self.target + } + + pub fn workspace_root(&self) -> &Path { + &self.workspace_root + } + + pub fn snapshot_store(&self) -> SnapshotStore { + self.snapshot_store.clone() + } + + pub fn current_snapshot(&self) -> Result, String> { + self.current_snapshot.lock().map(|current| current.clone()).map_err(|_| "remote session state lock poisoned".to_string()) + } + + pub fn sync_snapshot(&self) -> Result { + let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; + let snapshot = self.snapshot_builder.build_root(&self.workspace_root, self.session_id)?; + let new_handle = snapshot.handle().clone(); + let old_handle = self.current_snapshot()?.clone(); + if let Some(old_handle) = old_handle.filter(|old| old != &new_handle) { + self.snapshot_store.remove(&old_handle)?; + } + self.current_snapshot.lock().map_err(|_| "remote session state lock poisoned".to_string())?.replace(new_handle.clone()); + Ok(new_handle) + } + + fn materialize_current_snapshot(&self, destination: &Path) -> Result<(), String> { + let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; + let handle = self.current_snapshot()?.ok_or_else(|| "remote build requires a successful sync".to_string())?; + self.snapshot_store.materialize(&handle, destination).map(|_| ()) + } + + fn new_job_path(&self) -> Result { + let path = self.jobs_root.join(format!("job-{}", NEXT_JOB_ID.fetch_add(1, Ordering::Relaxed))); + if path.exists() { + return Err("loopback job path already exists".to_string()); + } + Ok(path) + } +} + +impl Drop for RunRemoteSession { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.jobs_root); + let _ = fs::remove_dir_all(self.snapshot_store_root()); + } +} + +impl RunRemoteSession { + fn snapshot_store_root(&self) -> PathBuf { + // SnapshotStore deliberately exposes no public root path; this private + // cleanup path is kept alongside the run-owned workspace state. + self.snapshot_store.root_for_cleanup() + } +} + +pub struct LoopbackBackend { + session: Arc, + tools: Arc>, + timeout: Duration, +} + +impl LoopbackBackend { + pub fn new(session: Arc, tools: BTreeMap) -> Self { + Self { session, tools: Arc::new(tools), timeout: LOOPBACK_BUILD_TIMEOUT } + } + + pub fn with_timeout(mut self, timeout: Duration) -> Self { + self.timeout = timeout; + self + } +} + +impl RemoteBackend for LoopbackBackend { + fn execute<'a>( + &'a self, request: AuthorizedRemoteRequest, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + let session = self.session.clone(); + let tools = self.tools.clone(); + let timeout = self.timeout; + Box::pin(async move { + match request.request().operation() { + RemoteOperation::Sync => execute_sync(session, events).await, + RemoteOperation::Build(build) => execute_build(session, tools, timeout, build, events).await, + } + }) + } +} + +async fn execute_sync(session: Arc, events: mpsc::Sender) -> Result<(), RemoteBackendError> { + send_event(&events, RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }).await?; + let sync = tokio::task::spawn_blocking(move || session.sync_snapshot()) + .await + .map_err(|error| RemoteBackendError::Failed(format!("snapshot worker failed: {error}")))?; + sync.map_err(RemoteBackendError::Failed)?; + send_event(&events, RemoteBackendEvent::Completed { exit_code: 0 }).await +} + +async fn execute_build( + session: Arc, tools: Arc>, timeout_duration: Duration, build: &crate::remote::RemoteBuild, + events: mpsc::Sender, +) -> Result<(), RemoteBackendError> { + if !build.env().is_empty() { + return Err(RemoteBackendError::Failed("remote environment is unavailable until remote environment policy is configured".to_string())); + } + let executable = tools + .get(build.tool().as_str()) + .ok_or_else(|| RemoteBackendError::Spawn(format!("loopback tool is not configured: {}", build.tool().as_str())))?; + let job_path = session.new_job_path().map_err(RemoteBackendError::Failed)?; + let _job = JobGuard { path: job_path.clone() }; + let destination = job_path.clone(); + tokio::task::spawn_blocking({ + let session = session.clone(); + move || session.materialize_current_snapshot(&destination) + }) + .await + .map_err(|error| RemoteBackendError::Failed(format!("materialization worker failed: {error}")))? + .map_err(RemoteBackendError::Failed)?; + + let cwd = job_path.join(build.cwd().as_str()); + let cwd_metadata = fs::symlink_metadata(&cwd).map_err(|error| RemoteBackendError::Failed(format!("remote cwd is unavailable: {error}")))?; + if !cwd_metadata.file_type().is_dir() { + return Err(RemoteBackendError::Failed("remote cwd is not a directory".to_string())); + } + + let mut command = Command::new(executable); + command + .args(build.argv()) + .current_dir(&cwd) + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()); + command.env_clear().env("PATH", LOOPBACK_PATH); + unsafe { + command.pre_exec(|| { + if libc::setpgid(0, 0) != 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(()) + }); + } + command.kill_on_drop(true); + let mut child = command.spawn().map_err(|error| RemoteBackendError::Spawn(format!("spawn loopback tool: {error}")))?; + let process_group = child.id().map(|pid| ProcessGroupGuard { pgid: pid as i32, active: true }); + let stdout = child.stdout.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stdout".to_string()))?; + let stderr = child.stderr.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stderr".to_string()))?; + let mut stdout_task = Box::pin(tokio::spawn(pump(stdout, RemoteStream::Stdout, events.clone()))); + let mut stderr_task = Box::pin(tokio::spawn(pump(stderr, RemoteStream::Stderr, events.clone()))); + let mut child_wait = Box::pin(child.wait()); + let mut timeout_sleep = Box::pin(sleep(timeout_duration)); + let mut child_status = None; + let mut stdout_done = false; + let mut stderr_done = false; + let mut failure = None; + + while child_status.is_none() || !stdout_done || !stderr_done { + tokio::select! { + status = &mut child_wait, if child_status.is_none() => { + child_status = Some(status.map_err(|error| RemoteBackendError::Failed(format!("wait for loopback tool: {error}")))); + } + result = &mut stdout_task, if !stdout_done => { + stdout_done = true; + if let Err(error) = join_pump(result).await { failure.get_or_insert(error); kill_process_group(process_group.as_ref()); } + } + result = &mut stderr_task, if !stderr_done => { + stderr_done = true; + if let Err(error) = join_pump(result).await { failure.get_or_insert(error); kill_process_group(process_group.as_ref()); } + } + _ = &mut timeout_sleep, if child_status.is_none() => { + failure.get_or_insert(RemoteBackendError::Timeout); + kill_process_group(process_group.as_ref()); + } + } + } + + if failure.is_none() { + kill_process_group(process_group.as_ref()); + } + if let Some(mut process_group) = process_group { + process_group.active = false; + } + if let Some(error) = failure { + return Err(error); + } + let status = child_status.unwrap()?; + send_event(&events, RemoteBackendEvent::Completed { exit_code: status.code().unwrap_or(-1) }).await +} + +#[derive(Clone, Copy)] +enum RemoteStream { + Stdout, + Stderr, +} + +async fn pump( + mut reader: R, stream: RemoteStream, events: mpsc::Sender, +) -> Result<(), RemoteBackendError> { + let mut buffer = [0u8; 8192]; + loop { + let count = reader.read(&mut buffer).await.map_err(|error| RemoteBackendError::Failed(format!("read loopback output: {error}")))?; + if count == 0 { + return Ok(()); + } + let event = match stream { + RemoteStream::Stdout => RemoteBackendEvent::Stdout(buffer[..count].to_vec()), + RemoteStream::Stderr => RemoteBackendEvent::Stderr(buffer[..count].to_vec()), + }; + events.send(event).await.map_err(|_| RemoteBackendError::Cancelled)?; + } +} + +async fn join_pump(result: Result, tokio::task::JoinError>) -> Result<(), RemoteBackendError> { + result.map_err(|error| RemoteBackendError::Failed(format!("loopback output task failed: {error}")))? +} + +async fn send_event(events: &mpsc::Sender, event: RemoteBackendEvent) -> Result<(), RemoteBackendError> { + events.send(event).await.map_err(|_| RemoteBackendError::Cancelled) +} + +struct JobGuard { + path: PathBuf, +} + +impl Drop for JobGuard { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.path); + } +} + +struct ProcessGroupGuard { + pgid: i32, + active: bool, +} + +impl Drop for ProcessGroupGuard { + fn drop(&mut self) { + if self.active { + kill_process_group(Some(self)); + } + } +} + +fn kill_process_group(group: Option<&ProcessGroupGuard>) { + if let Some(group) = group { + unsafe { + libc::kill(-group.pgid, libc::SIGTERM); + libc::kill(-group.pgid, libc::SIGKILL); + } + } +} + +pub fn resolve_fixed_tools(names: impl IntoIterator) -> BTreeMap { + let mut tools = BTreeMap::new(); + for name in names { + if name.is_empty() || name.contains('/') || name.as_bytes().contains(&0) { + continue; + } + for directory in ["/usr/local/bin", "/usr/bin", "/bin"] { + let path = Path::new(directory).join(&name); + if path.is_file() { + if let Ok(metadata) = fs::metadata(&path) { + use std::os::unix::fs::PermissionsExt; + if metadata.permissions().mode() & 0o111 != 0 { + tools.insert(name.clone(), path); + break; + } + } + } + } + } + tools +} + +pub fn cleanup_stale_roots(parent: &Path) -> Result<(), String> { + let current_uid = unsafe { libc::geteuid() }; + for entry in fs::read_dir(parent).map_err(|error| format!("read temporary runtime roots: {error}"))? { + let entry = entry.map_err(|error| format!("read temporary runtime root: {error}"))?; + let name = entry.file_name(); + let name = name.to_string_lossy(); + if !name.starts_with("bunkerbox-loopback-") && !name.starts_with("bunkerbox-snapshots-") { + continue; + } + let Some(pid) = name.split('-').nth(2).and_then(|value| value.parse::().ok()) else { + continue; + }; + if process_is_alive(pid) { + continue; + } + let metadata = fs::symlink_metadata(entry.path()).map_err(|error| format!("inspect stale runtime root: {error}"))?; + if metadata.file_type().is_dir() && metadata.uid() == current_uid { + fs::remove_dir_all(entry.path()).map_err(|error| format!("remove stale runtime root: {error}"))?; + } + } + Ok(()) +} + +fn process_is_alive(pid: libc::pid_t) -> bool { + if pid <= 0 { + return false; + } + let result = unsafe { libc::kill(pid, 0) }; + result == 0 || io::Error::last_os_error().raw_os_error() == Some(libc::EPERM) +} + +fn create_private_root(path: &Path, label: &str) -> Result<(), String> { + fs::create_dir(path).map_err(|error| format!("create {label}: {error}"))?; + if let Err(error) = fs::set_permissions(path, fs::Permissions::from_mode(0o700)) { + let _ = fs::remove_dir(path); + return Err(format!("set private {label} mode: {error}")); + } + Ok(()) +} + +#[cfg(test)] +#[path = "loopback_ut.rs"] +mod tests; diff --git a/src/main.rs b/src/main.rs index a7ec69c..79a9087 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,5 +1,6 @@ use bunkerbox::cfg::{ProjectConfig, WorkspaceMode}; -use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, overlay, tui, vscomm, workspace}; +use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, loopback, overlay, snapshot, tui, vscomm, workspace}; +use rand::RngCore; use std::ffi::OsString; use std::fs::File; use std::io; @@ -28,9 +29,11 @@ fn run() -> Result<(), String> { if cfg::RuntimeConfig::invoked_name()? != clidef::APPNAME { let share_dir = share_dir_from_args()?; if let Some(config) = cfg::RuntimeConfig::for_invoked_name(&share_dir)? { - let rt = tokio::runtime::Runtime::new().map_err(|e| format!("tokio: {e}"))?; - let _guard = rt.enter(); - return run_packaged_runtime(config, workspace_override, &share_dir); + return tokio::runtime::Runtime::new().map_err(|e| format!("tokio: {e}"))?.block_on(async move { + tokio::task::spawn_blocking(move || run_packaged_runtime(config, workspace_override, &share_dir)) + .await + .map_err(|error| format!("packaged runtime thread failed: {error}"))? + }); } } @@ -192,7 +195,8 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option>> = Arc::new(Mutex::new(None)); + + ensure_sudo()?; let mut sock_fds = [-1i32, -1]; if unsafe { libc::socketpair(libc::AF_UNIX, libc::SOCK_STREAM, 0, sock_fds.as_mut_ptr()) } != 0 { @@ -225,52 +229,17 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option 0, Err(e) => { eprintln!("bunkerbox: {e}"); @@ -295,23 +264,60 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option> = Arc::new(Mutex::new(tui::OverlayState::new())); let status_listener = start_status_listener(overlay.clone())?; - let setup_handle = tokio::runtime::Handle::current().clone(); - let daemon_slot = daemon_holder.clone(); - let setup_thread = std::thread::spawn(move || -> Result<(), String> { - let workspace = read_workspace_handoff(setup_parent_fd)?; - if passthrough.is_empty() { - return Ok(()); + let setup_result = (|| -> Result<(workspace::WorkspaceHandle, Arc, daemon::VsockDaemon), String> { + let workspace = workspace::resolve(workspace_mode, quota, exclude.as_deref(), &name)?; + let remote_session = new_session_id(); + let target = new_target_id(); + loopback::cleanup_stale_roots(&std::env::temp_dir())?; + let snapshot_root = std::env::temp_dir().join(format!("bunkerbox-snapshots-{}-{}", std::process::id(), remote_session.to_hex())); + let jobs_root = std::env::temp_dir().join(format!("bunkerbox-loopback-{}-{}", std::process::id(), remote_session.to_hex())); + let snapshot_store = snapshot::SnapshotStore::new(&snapshot_root); + let exclusion_policy = snapshot::SnapshotExclusionPolicy::from_config(&env, exclude.as_deref())?; + let snapshot_builder = snapshot::SnapshotBuilder::new(snapshot_store.clone(), snapshot::SnapshotLimits::default(), exclusion_policy); + let session = Arc::new(loopback::RunRemoteSession::new( + bunkerbox::remote::WorkspaceSessionId(remote_session.0), + target, + workspace.path().to_path_buf(), + snapshot_store, + snapshot_builder, + jobs_root, + )?); + let allowed_tools = remote_tool_names(&passthrough); + let tools = loopback::resolve_fixed_tools(allowed_tools.clone()); + let daemon = daemon::VsockDaemon::start_with_remote( + passthrough, + env_mode, + workspace.path().to_path_buf(), + profiles, + share_dir_owned, + merged_allow, + daemon::RemoteDaemonConfig::new(session.clone(), allowed_tools, tools), + )?; + if let Err(error) = write_run_handoff(&mut setup_parent, workspace.path(), remote_session) { + tokio::runtime::Handle::current().block_on(daemon.shutdown()); + return Err(error); } - - let _guard = setup_handle.enter(); - let daemon = daemon::VsockDaemon::start(passthrough, env_mode, workspace, profiles, share_dir_owned, merged_allow)?; - *daemon_slot.lock().map_err(|_| "daemon state lock poisoned".to_string())? = Some(daemon); - Ok(()) - }); + Ok((workspace, session, daemon)) + })(); + + let (workspace, remote_session, daemon) = match setup_result { + Ok(value) => value, + Err(error) => { + unsafe { + libc::kill(pid, libc::SIGTERM); + libc::waitpid(pid, std::ptr::null_mut(), 0); + libc::close(master); + } + tokio::runtime::Handle::current().block_on(status_listener.shutdown()); + return Err(error); + } + }; + drop(setup_parent); let tui_result = tui::event_loop(master, rows, cols, parent_fd, overlay); @@ -319,20 +325,13 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option result, - Err(_) => Err("workspace setup thread panicked".to_string()), - }; - - if let Some(d) = daemon_holder.lock().map_err(|_| "daemon state lock poisoned".to_string())?.take() { - tokio::runtime::Handle::current().block_on(d.shutdown()); - } + tokio::runtime::Handle::current().block_on(daemon.shutdown()); + drop(remote_session); + drop(workspace); tokio::runtime::Handle::current().block_on(status_listener.shutdown()); tui_result?; - setup_result?; - if status != 0 { return Err(format!("child exited with status {status}")); } @@ -340,17 +339,112 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option Result<(), String> { - let bytes = path.as_os_str().as_bytes(); - let frame = encode_workspace_handoff(bytes)?; - let mut file = unsafe { File::from_raw_fd(fd) }; - io::Write::write_all(&mut file, &frame).map_err(|err| format!("write workspace handoff: {err}")) +fn new_session_id() -> vscomm::WorkspaceSessionId { + loop { + let mut bytes = [0u8; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + if bytes != [0; 16] { + return vscomm::WorkspaceSessionId(bytes); + } + } +} + +fn new_target_id() -> bunkerbox::remote::RemoteTargetId { + loop { + let mut bytes = [0u8; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + if bytes != [0; 16] { + return bunkerbox::remote::RemoteTargetId(bytes); + } + } +} + +fn remote_tool_names(entries: &[String]) -> Vec { + let mut names = std::collections::BTreeSet::new(); + for entry in entries { + let command = entry.trim().strip_suffix(" *").unwrap_or(entry.trim()); + if let Some(tool) = command.split_whitespace().next().filter(|tool| !tool.is_empty() && !tool.contains('/')) { + names.insert(tool.to_string()); + } + } + names.into_iter().collect() +} + +fn write_run_handoff(file: &mut File, path: &Path, session_id: vscomm::WorkspaceSessionId) -> Result<(), String> { + if !path.is_absolute() { + return Err("workspace handoff path must be absolute".to_string()); + } + let mut payload = path.as_os_str().as_bytes().to_vec(); + if payload.contains(&0) { + return Err("workspace path contains NUL".to_string()); + } + payload.push(0); + payload.extend_from_slice(&session_id.0); + let frame = encode_workspace_handoff(&payload)?; + io::Write::write_all(file, &frame).map_err(|err| format!("write run handoff: {err}")) +} + +fn ensure_sudo() -> Result<(), String> { + if !std::process::Command::new("sudo") + .arg("-n") + .arg("true") + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .status() + .map(|status| status.success()) + .unwrap_or(false) + { + let pass = bunkerbox::logging::prompt_password("Sudo password", "Enter your sudo password")?; + let mut child = std::process::Command::new("sudo") + .arg("-S") + .arg("-v") + .stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .spawn() + .map_err(|error| format!("failed to run sudo: {error}"))?; + io::Write::write_all(child.stdin.as_mut().ok_or_else(|| "failed to open sudo stdin".to_string())?, pass.as_bytes()) + .map_err(|error| format!("failed to write sudo password: {error}"))?; + drop(child.stdin.take()); + if !child.wait().map_err(|error| format!("sudo failed: {error}"))?.success() { + return Err("sudo: authentication failed".to_string()); + } + } + + std::thread::spawn(|| loop { + std::thread::sleep(std::time::Duration::from_secs(240)); + let _ = std::process::Command::new("sudo") + .arg("-n") + .arg("-v") + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .status(); + }); + Ok(()) +} + +fn read_run_handoff(file: &mut File) -> Result<(PathBuf, vscomm::WorkspaceSessionId), String> { + let payload = read_workspace_handoff(file)?; + let bytes = payload.as_os_str().as_bytes(); + if bytes.len() < 17 || bytes[bytes.len() - 17] != 0 { + return Err("run handoff is malformed".to_string()); + } + let path = PathBuf::from(OsString::from_vec(bytes[..bytes.len() - 17].to_vec())); + if !path.is_absolute() { + return Err("run handoff path must be absolute".to_string()); + } + let mut session = [0u8; 16]; + session.copy_from_slice(&bytes[bytes.len() - 16..]); + if session == [0; 16] { + return Err("run handoff session is zero".to_string()); + } + Ok((path, vscomm::WorkspaceSessionId(session))) } -fn read_workspace_handoff(fd: RawFd) -> Result { - let mut file = unsafe { File::from_raw_fd(fd) }; +fn read_workspace_handoff(file: &mut File) -> Result { let mut header = [0u8; 8]; - io::Read::read_exact(&mut file, &mut header).map_err(|err| format!("read workspace handoff header: {err}"))?; + io::Read::read_exact(file, &mut header).map_err(|err| format!("read workspace handoff header: {err}"))?; let payload_len = u32::from_le_bytes([header[4], header[5], header[6], header[7]]) as usize; if payload_len > MAX_WORKSPACE_HANDOFF_BYTES { @@ -358,7 +452,7 @@ fn read_workspace_handoff(fd: RawFd) -> Result { } let mut payload = vec![0u8; payload_len]; - io::Read::read_exact(&mut file, &mut payload).map_err(|err| format!("read workspace handoff: {err}"))?; + io::Read::read_exact(file, &mut payload).map_err(|err| format!("read workspace handoff: {err}"))?; let mut frame = header.to_vec(); frame.extend_from_slice(&payload); diff --git a/src/remote_client.rs b/src/remote_client.rs new file mode 100644 index 0000000..16f226c --- /dev/null +++ b/src/remote_client.rs @@ -0,0 +1,43 @@ +use crate::vscomm::{RemoteBuild, RemoteEvent, RemoteEventKind, RemoteRequest, RemoteTool, RequestId, WorkspaceRelativePath, WorkspaceSessionId}; +use std::io::{Read, Write}; + +pub fn remote_sync_request(request_id: RequestId, session_id: WorkspaceSessionId) -> RemoteRequest { + RemoteRequest::sync(request_id, session_id) +} + +pub fn remote_build_request( + request_id: RequestId, session_id: WorkspaceSessionId, cwd: impl Into, tool: impl Into, argv: Vec, + env: Vec<(String, String)>, +) -> Result { + let build = RemoteBuild::new(WorkspaceRelativePath::new(cwd)?, RemoteTool::new(tool)?, argv, env)?; + Ok(RemoteRequest::build(request_id, session_id, build)) +} + +pub fn execute_remote_request_to( + stream: &mut S, request: RemoteRequest, stdout: &mut WOut, stderr: &mut WErr, +) -> Result { + let request_id = request.request_id; + request.to_frame()?.write(stream).map_err(|e| format!("send remote request: {e}"))?; + + loop { + let frame = crate::vscomm::Frame::read(stream).map_err(|e| format!("read remote event: {e}"))?; + let event = RemoteEvent::from_frame(frame)?; + if event.request_id != request_id { + return Err("remote event request ID mismatch".to_string()); + } + match event.kind { + RemoteEventKind::SyncProgress { .. } => {} + RemoteEventKind::Stdout(data) => { + stdout.write_all(&data).map_err(|e| format!("stdout: {e}"))?; + stdout.flush().map_err(|e| format!("flush stdout: {e}"))?; + } + RemoteEventKind::Stderr(data) => { + stderr.write_all(&data).map_err(|e| format!("stderr: {e}"))?; + stderr.flush().map_err(|e| format!("flush stderr: {e}"))?; + } + RemoteEventKind::Error { message, .. } => return Err(message), + RemoteEventKind::Cancelled => return Err("remote operation cancelled".to_string()), + RemoteEventKind::Completed { exit_code } => return Ok(exit_code), + } + } +} diff --git a/src/snapshot.rs b/src/snapshot.rs new file mode 100644 index 0000000..bda717a --- /dev/null +++ b/src/snapshot.rs @@ -0,0 +1,892 @@ +use crate::cfg::ProjectConfig; +use crate::remote::WorkspaceSessionId; +use crate::workspace::WorkspaceHandle; +use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; +use std::collections::BTreeSet; +use std::ffi::{CStr, CString, OsStr, OsString}; +use std::fs::{self, File, OpenOptions}; +use std::io::{self, Read, Write}; +use std::os::fd::{AsRawFd, FromRawFd, RawFd}; +use std::os::unix::ffi::{OsStrExt, OsStringExt}; +use std::os::unix::fs::OpenOptionsExt; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{Duration, Instant}; + +pub const MAX_SNAPSHOT_ENTRIES: usize = 10_000; +pub const MAX_SNAPSHOT_TOTAL_BYTES: u64 = 512 * 1024 * 1024; +pub const MAX_SNAPSHOT_FILE_BYTES: u64 = 64 * 1024 * 1024; +pub const MAX_SNAPSHOT_PATH_BYTES: usize = 4 * 1024; +pub const MAX_SNAPSHOT_COMPONENT_BYTES: usize = 255; +pub const MAX_SNAPSHOT_DEPTH: usize = 64; +pub const MAX_SNAPSHOT_MANIFEST_BYTES: usize = 16 * 1024 * 1024; +pub const SNAPSHOT_COPY_BUFFER_BYTES: usize = 64 * 1024; + +static NEXT_STAGING_ID: AtomicU64 = AtomicU64::new(1); + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct SnapshotId([u8; 32]); + +impl SnapshotId { + pub fn as_bytes(&self) -> &[u8; 32] { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SnapshotRelativePath(String); + +impl SnapshotRelativePath { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + validate_relative_path(&value, MAX_SNAPSHOT_PATH_BYTES, MAX_SNAPSHOT_COMPONENT_BYTES, MAX_SNAPSHOT_DEPTH)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum SnapshotEntryKind { + Directory, + RegularFile, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SnapshotEntry { + path: SnapshotRelativePath, + kind: SnapshotEntryKind, + mode: u16, + size: u64, + content_digest: Option<[u8; 32]>, +} + +impl SnapshotEntry { + pub fn path(&self) -> &SnapshotRelativePath { + &self.path + } + + pub fn kind(&self) -> SnapshotEntryKind { + self.kind + } + + pub fn mode(&self) -> u16 { + self.mode + } + + pub fn size(&self) -> u64 { + self.size + } + + pub fn content_digest(&self) -> Option<&[u8; 32]> { + self.content_digest.as_ref() + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct SnapshotLimits { + pub max_entries: usize, + pub max_total_bytes: u64, + pub max_file_bytes: u64, + pub max_path_bytes: usize, + pub max_component_bytes: usize, + pub max_depth: usize, + pub max_manifest_bytes: usize, + pub max_duration: Duration, +} + +impl Default for SnapshotLimits { + fn default() -> Self { + Self { + max_entries: MAX_SNAPSHOT_ENTRIES, + max_total_bytes: MAX_SNAPSHOT_TOTAL_BYTES, + max_file_bytes: MAX_SNAPSHOT_FILE_BYTES, + max_path_bytes: MAX_SNAPSHOT_PATH_BYTES, + max_component_bytes: MAX_SNAPSHOT_COMPONENT_BYTES, + max_depth: MAX_SNAPSHOT_DEPTH, + max_manifest_bytes: MAX_SNAPSHOT_MANIFEST_BYTES, + max_duration: Duration::from_secs(60), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SnapshotExclusionPolicy { + basename_prunes: BTreeSet, + anchored_prunes: BTreeSet, +} + +impl SnapshotExclusionPolicy { + pub fn from_config(config: &ProjectConfig, runtime_exclude: Option<&[String]>) -> Result { + Self::from_patterns(config.effective_exclude(runtime_exclude)) + } + + pub fn from_patterns(patterns: impl IntoIterator) -> Result { + let mut policy = Self { basename_prunes: BTreeSet::new(), anchored_prunes: BTreeSet::new() }; + for name in [".git", ".bunker", ".bunkerbox", ".env", ".envrc", ".ssh"] { + policy.basename_prunes.insert(name.to_string()); + } + for pattern in patterns { + policy.add_pattern(&pattern)?; + } + Ok(policy) + } + + pub fn excludes(&self, path: &str) -> bool { + let components = path.split('/'); + if components.clone().any(|component| self.basename_prunes.contains(component)) { + return true; + } + self.anchored_prunes.iter().any(|prefix| path == prefix || path.starts_with(&format!("{prefix}/"))) + } + + fn add_pattern(&mut self, raw: &str) -> Result<(), String> { + let pattern = raw.trim().trim_end_matches('/'); + if pattern.is_empty() || pattern.starts_with('/') || pattern.contains('\\') { + return Err(format!("invalid snapshot exclusion: {raw}")); + } + let components = pattern.split('/').collect::>(); + if components.iter().any(|component| component.is_empty() || *component == "." || *component == "..") { + return Err(format!("invalid snapshot exclusion: {raw}")); + } + for component in &components { + if component.len() > MAX_SNAPSHOT_COMPONENT_BYTES || component.as_bytes().contains(&0) { + return Err(format!("snapshot exclusion component is too long: {raw}")); + } + } + if components.len() == 1 { + self.basename_prunes.insert(components[0].to_string()); + } else { + let normalized = components.join("/"); + if normalized.len() > MAX_SNAPSHOT_PATH_BYTES { + return Err(format!("snapshot exclusion is too long: {raw}")); + } + self.anchored_prunes.insert(normalized); + } + Ok(()) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SnapshotHandle { + session_id: WorkspaceSessionId, + snapshot_id: SnapshotId, +} + +impl SnapshotHandle { + pub fn session_id(&self) -> WorkspaceSessionId { + self.session_id + } + + pub fn snapshot_id(&self) -> SnapshotId { + self.snapshot_id + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkspaceSnapshot { + handle: SnapshotHandle, + entries: Vec, + total_file_bytes: u64, +} + +impl WorkspaceSnapshot { + pub fn handle(&self) -> &SnapshotHandle { + &self.handle + } + + pub fn entries(&self) -> &[SnapshotEntry] { + &self.entries + } + + pub fn total_file_bytes(&self) -> u64 { + self.total_file_bytes + } +} + +#[derive(Debug, Clone)] +pub struct SnapshotStore { + root: PathBuf, +} + +impl SnapshotStore { + pub fn new(root: impl Into) -> Self { + Self { root: root.into() } + } + + pub fn resolve(&self, handle: &SnapshotHandle) -> Result { + let manifest_path = self.manifest_path(handle); + let metadata = fs::symlink_metadata(&manifest_path).map_err(|error| format!("snapshot is unavailable: {error}"))?; + if !metadata.file_type().is_file() { + return Err("snapshot manifest is not a regular file".to_string()); + } + if metadata.len() > MAX_SNAPSHOT_MANIFEST_BYTES as u64 { + return Err("stored snapshot manifest exceeds limit".to_string()); + } + let stored: StoredSnapshot = serde_json::from_slice(&fs::read(&manifest_path).map_err(|error| format!("read snapshot manifest: {error}"))?) + .map_err(|error| format!("decode snapshot manifest: {error}"))?; + if stored.session_id != handle.session_id.0 { + return Err("snapshot session mismatch".to_string()); + } + let (entries, total_file_bytes) = stored.into_entries()?; + let snapshot_id = snapshot_id(&entries); + if snapshot_id != handle.snapshot_id { + return Err("snapshot manifest identity mismatch".to_string()); + } + Ok(WorkspaceSnapshot { handle: handle.clone(), entries, total_file_bytes }) + } + + pub fn remove(&self, handle: &SnapshotHandle) -> Result<(), String> { + let path = self.snapshot_path(handle); + match fs::symlink_metadata(&path) { + Ok(metadata) if metadata.file_type().is_dir() => fs::remove_dir_all(&path).map_err(|error| format!("remove snapshot: {error}")), + Ok(_) => Err("snapshot publication is not a directory".to_string()), + Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()), + Err(error) => Err(format!("inspect snapshot for removal: {error}")), + } + } + + pub fn materialize(&self, handle: &SnapshotHandle, destination: &Path) -> Result { + let snapshot = self.resolve(handle)?; + if fs::symlink_metadata(destination).is_ok() { + return Err("materialization destination already exists".to_string()); + } + if let Some(parent) = destination.parent() { + fs::create_dir_all(parent).map_err(|error| format!("create materialization parent: {error}"))?; + } + fs::create_dir(destination).map_err(|error| format!("create materialization destination: {error}"))?; + let mut cleanup = MaterializationGuard { path: destination.to_path_buf(), committed: false }; + set_mode(destination, 0o700)?; + let source_root = open_directory(&self.files_path(handle))?; + let destination_root = open_directory(destination)?; + + for entry in snapshot.entries() { + match entry.kind { + SnapshotEntryKind::Directory => ensure_destination_directory(&destination_root, entry.path.as_str(), entry.mode)?, + SnapshotEntryKind::RegularFile => { + let source = open_relative_file(&source_root, entry.path.as_str(), libc::O_RDONLY)?; + let destination_file = create_relative_file(&destination_root, entry.path.as_str(), entry.mode)?; + copy_materialized_file( + &source, + &destination_file, + entry.size, + entry.content_digest.ok_or_else(|| "regular file has no digest".to_string())?, + entry.path.as_str(), + )?; + } + } + } + + cleanup.committed = true; + Ok(MaterializedWorkspace { root: destination.to_path_buf() }) + } + + fn manifest_path(&self, handle: &SnapshotHandle) -> PathBuf { + self.snapshot_path(handle).join("manifest.json") + } + + fn snapshot_path(&self, handle: &SnapshotHandle) -> PathBuf { + self.root.join(hex(&handle.session_id.0)).join(hex(&handle.snapshot_id.0)) + } + + fn files_path(&self, handle: &SnapshotHandle) -> PathBuf { + self.snapshot_path(handle).join("files") + } + + pub(crate) fn root_for_cleanup(&self) -> PathBuf { + self.root.clone() + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MaterializedWorkspace { + root: PathBuf, +} + +impl MaterializedWorkspace { + pub fn root(&self) -> &Path { + &self.root + } +} + +struct MaterializationGuard { + path: PathBuf, + committed: bool, +} + +impl Drop for MaterializationGuard { + fn drop(&mut self) { + if !self.committed { + let _ = fs::remove_dir_all(&self.path); + } + } +} + +pub struct SnapshotBuilder { + store: SnapshotStore, + limits: SnapshotLimits, + exclusions: SnapshotExclusionPolicy, +} + +impl SnapshotBuilder { + pub fn new(store: SnapshotStore, limits: SnapshotLimits, exclusions: SnapshotExclusionPolicy) -> Self { + Self { store, limits, exclusions } + } + + pub fn build(&self, workspace: &WorkspaceHandle, session_id: WorkspaceSessionId) -> Result { + self.build_root(workspace.path(), session_id) + } + + pub(crate) fn build_root(&self, workspace_root: &Path, session_id: WorkspaceSessionId) -> Result { + validate_limits(&self.limits)?; + if session_id.0 == [0; 16] { + return Err("snapshot requires an authoritative nonzero workspace session".to_string()); + } + let started = Instant::now(); + let canonical_root = fs::canonicalize(workspace_root).map_err(|error| format!("resolve snapshot workspace: {error}"))?; + let root = open_directory(&canonical_root)?; + let root_stat = stat_fd(root.as_raw_fd())?; + let root_device = root_stat.st_dev; + let store_root = prepare_store_root(&self.store.root)?; + let staging_root = store_root.join(".staging"); + fs::create_dir_all(&staging_root).map_err(|error| format!("create snapshot staging root: {error}"))?; + set_mode(&staging_root, 0o700)?; + let stage = staging_root.join(format!("{}-{}", hex(&session_id.0), NEXT_STAGING_ID.fetch_add(1, Ordering::Relaxed))); + fs::create_dir(&stage).map_err(|error| format!("create snapshot staging directory: {error}"))?; + set_mode(&stage, 0o700)?; + let mut cleanup = StagingGuard { path: stage.clone(), committed: false }; + let files_root = stage.join("files"); + fs::create_dir(&files_root).map_err(|error| format!("create snapshot content directory: {error}"))?; + + let mut state = WalkState { + entries: Vec::new(), + total_file_bytes: 0, + started, + root_device, + stage_files: files_root, + next_buffer: vec![0; SNAPSHOT_COPY_BUFFER_BYTES], + }; + walk_directory(root.as_raw_fd(), "", 0, &self.limits, &self.exclusions, &mut state)?; + state.entries.sort_by(|left, right| left.path.as_str().cmp(right.path.as_str())); + let id = snapshot_id(&state.entries); + let stored = StoredSnapshot::from_entries(session_id, &state.entries, state.total_file_bytes); + let manifest = serde_json::to_vec(&stored).map_err(|error| format!("encode snapshot manifest: {error}"))?; + if manifest.len() > self.limits.max_manifest_bytes { + return Err(format!("snapshot manifest exceeds maximum size {}", self.limits.max_manifest_bytes)); + } + fs::write(stage.join("manifest.json"), manifest).map_err(|error| format!("write snapshot manifest: {error}"))?; + + let final_session = store_root.join(hex(&session_id.0)); + fs::create_dir_all(&final_session).map_err(|error| format!("create snapshot session directory: {error}"))?; + set_mode(&final_session, 0o700)?; + let final_path = final_session.join(hex(&id.0)); + if !final_path.exists() { + fs::rename(&stage, &final_path).map_err(|error| format!("publish snapshot: {error}"))?; + } else { + if !fs::symlink_metadata(&final_path).map_err(|error| format!("inspect existing snapshot: {error}"))?.file_type().is_dir() { + return Err("existing snapshot publication is not a directory".to_string()); + } + fs::remove_dir_all(&stage).map_err(|error| format!("discard duplicate snapshot staging: {error}"))?; + } + cleanup.committed = true; + Ok(WorkspaceSnapshot { + handle: SnapshotHandle { session_id, snapshot_id: id }, + entries: state.entries, + total_file_bytes: state.total_file_bytes, + }) + } +} + +struct StagingGuard { + path: PathBuf, + committed: bool, +} + +impl Drop for StagingGuard { + fn drop(&mut self) { + if !self.committed { + let _ = fs::remove_dir_all(&self.path); + } + } +} + +struct WalkState { + entries: Vec, + total_file_bytes: u64, + started: Instant, + root_device: libc::dev_t, + stage_files: PathBuf, + next_buffer: Vec, +} + +fn walk_directory( + directory_fd: RawFd, parent: &str, depth: usize, limits: &SnapshotLimits, exclusions: &SnapshotExclusionPolicy, state: &mut WalkState, +) -> Result<(), String> { + check_deadline(state.started, limits)?; + if depth > limits.max_depth { + return Err(format!("snapshot exceeds maximum depth {}", limits.max_depth)); + } + for name in read_directory_names(directory_fd)? { + check_deadline(state.started, limits)?; + let component = name.to_str().ok_or_else(|| "snapshot contains a non-UTF-8 path component".to_string())?; + validate_component(component, limits.max_component_bytes)?; + let relative = if parent.is_empty() { component.to_string() } else { format!("{parent}/{component}") }; + validate_relative_path(&relative, limits.max_path_bytes, limits.max_component_bytes, limits.max_depth)?; + if exclusions.excludes(&relative) { + continue; + } + let child_stat = stat_at(directory_fd, &name)?; + if child_stat.st_dev != state.root_device { + return Err(format!("snapshot entry crosses filesystem boundary: {relative}")); + } + let entry_kind = child_kind(&child_stat, &relative)?; + match entry_kind { + SnapshotEntryKind::Directory => { + add_entry_limit(state.entries.len(), limits)?; + let child = open_child_directory(directory_fd, &name, &relative)?; + let mode = normalized_mode(child_stat.st_mode); + state.entries.push(SnapshotEntry { + path: SnapshotRelativePath::new(relative.clone())?, + kind: SnapshotEntryKind::Directory, + mode, + size: 0, + content_digest: None, + }); + walk_directory(child.as_raw_fd(), &relative, depth + 1, limits, exclusions, state)?; + } + SnapshotEntryKind::RegularFile => { + add_entry_limit(state.entries.len(), limits)?; + if child_stat.st_nlink > 1 { + return Err(format!("snapshot rejects hard-linked file: {relative}")); + } + let size = checked_file_size(child_stat.st_size, limits.max_file_bytes, &relative)?; + let new_total = state.total_file_bytes.checked_add(size).ok_or_else(|| "snapshot total size overflow".to_string())?; + if new_total > limits.max_total_bytes { + return Err(format!("snapshot exceeds maximum total size {}", limits.max_total_bytes)); + } + let (digest, bytes_read) = copy_and_hash_file(directory_fd, &name, &relative, size, &child_stat, limits, state)?; + if bytes_read != size { + return Err(format!("file changed while snapshotting: {relative}")); + } + state.total_file_bytes = new_total; + state.entries.push(SnapshotEntry { + path: SnapshotRelativePath::new(relative)?, + kind: SnapshotEntryKind::RegularFile, + mode: normalized_mode(child_stat.st_mode), + size, + content_digest: Some(digest), + }); + } + } + } + Ok(()) +} + +fn copy_and_hash_file( + parent_fd: RawFd, name: &OsStr, relative: &str, expected_size: u64, expected_stat: &libc::stat, limits: &SnapshotLimits, state: &mut WalkState, +) -> Result<([u8; 32], u64), String> { + let file = open_child_file(parent_fd, name, relative)?; + let opened_stat = stat_fd(file.as_raw_fd())?; + compare_file_stat(expected_stat, &opened_stat, relative)?; + let destination = state.stage_files.join(relative); + if let Some(parent) = destination.parent() { + fs::create_dir_all(parent).map_err(|error| format!("create staged parent for {relative}: {error}"))?; + } + let mut staged = File::create(&destination).map_err(|error| format!("create staged file {relative}: {error}"))?; + let mut hasher = Sha256::new(); + let mut read_bytes = 0u64; + loop { + check_deadline(state.started, limits)?; + let count = (&file).read(&mut state.next_buffer).map_err(|error| format!("read snapshot file {relative}: {error}"))?; + if count == 0 { + break; + } + read_bytes = read_bytes.checked_add(count as u64).ok_or_else(|| format!("snapshot file size overflow: {relative}"))?; + if read_bytes > expected_size || read_bytes > limits.max_file_bytes { + return Err(format!("file changed beyond snapshot limit: {relative}")); + } + hasher.update(&state.next_buffer[..count]); + staged.write_all(&state.next_buffer[..count]).map_err(|error| format!("stage snapshot file {relative}: {error}"))?; + } + staged.sync_all().map_err(|error| format!("flush staged file {relative}: {error}"))?; + let final_stat = stat_fd(file.as_raw_fd())?; + compare_file_stat(expected_stat, &final_stat, relative)?; + if read_bytes != expected_size { + return Err(format!("file changed while snapshotting: {relative}")); + } + set_mode(&destination, normalized_mode(expected_stat.st_mode))?; + Ok((hasher.finalize().into(), read_bytes)) +} + +fn check_deadline(started: Instant, limits: &SnapshotLimits) -> Result<(), String> { + if started.elapsed() >= limits.max_duration { + return Err("snapshot creation deadline exceeded".to_string()); + } + Ok(()) +} + +fn add_entry_limit(count: usize, limits: &SnapshotLimits) -> Result<(), String> { + if count >= limits.max_entries { + return Err(format!("snapshot exceeds maximum entry count {}", limits.max_entries)); + } + Ok(()) +} + +fn validate_limits(limits: &SnapshotLimits) -> Result<(), String> { + if limits.max_entries == 0 + || limits.max_total_bytes == 0 + || limits.max_file_bytes == 0 + || limits.max_path_bytes == 0 + || limits.max_path_bytes > MAX_SNAPSHOT_PATH_BYTES + || limits.max_component_bytes == 0 + || limits.max_component_bytes > MAX_SNAPSHOT_COMPONENT_BYTES + || limits.max_depth == 0 + || limits.max_manifest_bytes == 0 + || limits.max_duration.is_zero() + { + return Err("invalid snapshot limits".to_string()); + } + Ok(()) +} + +fn checked_file_size(size: libc::off_t, max: u64, path: &str) -> Result { + if size < 0 { + return Err(format!("snapshot file has invalid size: {path}")); + } + let size = size as u64; + if size > max { + return Err(format!("snapshot file exceeds maximum size {max}: {path}")); + } + Ok(size) +} + +fn validate_relative_path(value: &str, max_path: usize, max_component: usize, max_depth: usize) -> Result<(), String> { + if value.is_empty() || value.starts_with('/') || value.contains('\\') || value.len() > max_path { + return Err(format!("invalid snapshot relative path: {value}")); + } + let components = value.split('/').collect::>(); + if components.len() > max_depth || components.iter().any(|part| part.is_empty() || *part == "." || *part == "..") { + return Err(format!("invalid snapshot relative path: {value}")); + } + components.iter().try_for_each(|part| validate_component(part, max_component)) +} + +fn validate_component(value: &str, max: usize) -> Result<(), String> { + if value.is_empty() || value == "." || value == ".." || value.len() > max || value.as_bytes().contains(&0) { + return Err(format!("invalid snapshot path component: {value}")); + } + Ok(()) +} + +fn child_kind(stat: &libc::stat, path: &str) -> Result { + match stat.st_mode & libc::S_IFMT { + libc::S_IFDIR => Ok(SnapshotEntryKind::Directory), + libc::S_IFREG => Ok(SnapshotEntryKind::RegularFile), + libc::S_IFLNK => Err(format!("snapshot rejects symlink: {path}")), + _ => Err(format!("snapshot rejects special file: {path}")), + } +} + +fn normalized_mode(mode: libc::mode_t) -> u16 { + (mode & 0o777) as u16 +} + +fn compare_file_stat(expected: &libc::stat, actual: &libc::stat, path: &str) -> Result<(), String> { + if expected.st_dev != actual.st_dev + || expected.st_ino != actual.st_ino + || expected.st_mode & libc::S_IFMT != actual.st_mode & libc::S_IFMT + || expected.st_size != actual.st_size + || expected.st_nlink != actual.st_nlink + || expected.st_mode & 0o777 != actual.st_mode & 0o777 + { + return Err(format!("file changed while snapshotting: {path}")); + } + Ok(()) +} + +fn set_mode(path: &Path, mode: u16) -> Result<(), String> { + use std::os::unix::fs::PermissionsExt; + fs::set_permissions(path, fs::Permissions::from_mode(mode as u32)).map_err(|error| format!("set staged file mode: {error}")) +} + +fn prepare_store_root(root: &Path) -> Result { + fs::create_dir_all(root).map_err(|error| format!("create snapshot store: {error}"))?; + fs::canonicalize(root).map_err(|error| format!("resolve snapshot store: {error}")) +} + +fn open_directory(path: &Path) -> Result { + OpenOptions::new() + .read(true) + .custom_flags(libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .open(path) + .map_err(|error| format!("open snapshot workspace: {error}")) +} + +fn open_child_directory(parent: RawFd, name: &OsStr, path: &str) -> Result { + open_at(parent, name, libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot directory {path}: {error}")) +} + +fn open_child_file(parent: RawFd, name: &OsStr, path: &str) -> Result { + open_at(parent, name, libc::O_RDONLY | libc::O_NOFOLLOW | libc::O_CLOEXEC).map_err(|error| format!("open snapshot file {path}: {error}")) +} + +fn open_at(parent: RawFd, name: &OsStr, flags: i32) -> io::Result { + let name = CString::new(name.as_bytes()).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "NUL in snapshot path"))?; + let fd = unsafe { libc::openat(parent, name.as_ptr(), flags, 0) }; + if fd < 0 { + return Err(io::Error::last_os_error()); + } + Ok(unsafe { File::from_raw_fd(fd) }) +} + +fn open_at_mode(parent: RawFd, name: &OsStr, flags: i32, mode: u32) -> io::Result { + let name = CString::new(name.as_bytes()).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "NUL in snapshot path"))?; + let fd = unsafe { libc::openat(parent, name.as_ptr(), flags, mode) }; + if fd < 0 { + return Err(io::Error::last_os_error()); + } + Ok(unsafe { File::from_raw_fd(fd) }) +} + +fn ensure_destination_directory(root: &File, relative: &str, mode: u16) -> Result<(), String> { + let mut current = root.try_clone().map_err(|error| format!("clone materialization root: {error}"))?; + for component in relative.split('/') { + let name = OsStr::new(component); + let name = CString::new(name.as_bytes()).map_err(|_| "NUL in materialization path".to_string())?; + let result = unsafe { libc::mkdirat(current.as_raw_fd(), name.as_ptr(), 0o700) }; + if result != 0 { + let error = io::Error::last_os_error(); + if error.kind() != io::ErrorKind::AlreadyExists { + return Err(format!("create materialized directory {relative}: {error}")); + } + } + current = open_at(current.as_raw_fd(), OsStr::new(component), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open materialized directory {relative}: {error}"))?; + } + if unsafe { libc::fchmod(current.as_raw_fd(), mode as libc::mode_t) } != 0 { + return Err(format!("set materialized directory mode {relative}: {}", io::Error::last_os_error())); + } + Ok(()) +} + +fn create_relative_file(root: &File, relative: &str, mode: u16) -> Result { + let mut components = relative.split('/').collect::>(); + let file_name = components.pop().ok_or_else(|| "empty materialization path".to_string())?; + let mut parent = root.try_clone().map_err(|error| format!("clone materialization root: {error}"))?; + for component in components { + parent = open_at(parent.as_raw_fd(), OsStr::new(component), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open materialized parent {relative}: {error}"))?; + } + open_at_mode( + parent.as_raw_fd(), + OsStr::new(file_name), + libc::O_WRONLY | libc::O_CREAT | libc::O_EXCL | libc::O_NOFOLLOW | libc::O_CLOEXEC, + (mode & 0o777) as u32, + ) + .map_err(|error| format!("create materialized file {relative}: {error}")) +} + +fn open_relative_file(root: &File, relative: &str, flags: i32) -> Result { + let mut components = relative.split('/').collect::>(); + let file_name = components.pop().ok_or_else(|| "empty snapshot content path".to_string())?; + let mut parent = root.try_clone().map_err(|error| format!("clone snapshot content root: {error}"))?; + for component in components { + parent = open_at(parent.as_raw_fd(), OsStr::new(component), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot content parent {relative}: {error}"))?; + } + open_at(parent.as_raw_fd(), OsStr::new(file_name), flags | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot content {relative}: {error}")) +} + +fn copy_materialized_file(source: &File, destination: &File, expected_size: u64, expected_digest: [u8; 32], path: &str) -> Result<(), String> { + let mut source = source.try_clone().map_err(|error| format!("clone snapshot content {path}: {error}"))?; + let mut destination = destination.try_clone().map_err(|error| format!("clone materialized file {path}: {error}"))?; + let mut buffer = vec![0u8; SNAPSHOT_COPY_BUFFER_BYTES]; + let mut hasher = Sha256::new(); + let mut copied = 0u64; + loop { + let count = source.read(&mut buffer).map_err(|error| format!("read snapshot content {path}: {error}"))?; + if count == 0 { + break; + } + copied = copied.checked_add(count as u64).ok_or_else(|| format!("materialized size overflow: {path}"))?; + if copied > expected_size { + return Err(format!("snapshot content is larger than manifest: {path}")); + } + hasher.update(&buffer[..count]); + destination.write_all(&buffer[..count]).map_err(|error| format!("write materialized file {path}: {error}"))?; + } + if copied != expected_size || hasher.finalize().as_slice() != expected_digest { + return Err(format!("snapshot content digest mismatch: {path}")); + } + destination.sync_all().map_err(|error| format!("flush materialized file {path}: {error}"))?; + Ok(()) +} + +fn stat_at(parent: RawFd, name: &OsStr) -> Result { + let name = CString::new(name.as_bytes()).map_err(|_| "NUL in snapshot path".to_string())?; + let mut stat = unsafe { std::mem::zeroed::() }; + if unsafe { libc::fstatat(parent, name.as_ptr(), &mut stat, libc::AT_SYMLINK_NOFOLLOW) } != 0 { + return Err(format!("stat snapshot entry: {}", io::Error::last_os_error())); + } + Ok(stat) +} + +fn stat_fd(fd: RawFd) -> Result { + let mut stat = unsafe { std::mem::zeroed::() }; + if unsafe { libc::fstat(fd, &mut stat) } != 0 { + return Err(format!("stat snapshot descriptor: {}", io::Error::last_os_error())); + } + Ok(stat) +} + +fn read_directory_names(fd: RawFd) -> Result, String> { + let duplicate = unsafe { libc::dup(fd) }; + if duplicate < 0 { + return Err(format!("duplicate snapshot directory: {}", io::Error::last_os_error())); + } + let stream = unsafe { libc::fdopendir(duplicate) }; + if stream.is_null() { + unsafe { libc::close(duplicate) }; + return Err(format!("open snapshot directory stream: {}", io::Error::last_os_error())); + } + let stream = DirectoryStream(stream); + let mut names = Vec::new(); + loop { + set_errno(0); + let entry = unsafe { libc::readdir(stream.0) }; + if entry.is_null() { + let errno = get_errno(); + if errno != 0 { + return Err(format!("read snapshot directory: {errno}")); + } + break; + } + let name = unsafe { CStr::from_ptr((*entry).d_name.as_ptr()) }.to_bytes(); + if name != b"." && name != b".." { + names.push(OsString::from_vec(name.to_vec())); + } + } + names.sort_by(|left, right| left.as_bytes().cmp(right.as_bytes())); + Ok(names) +} + +struct DirectoryStream(*mut libc::DIR); + +impl Drop for DirectoryStream { + fn drop(&mut self) { + unsafe { libc::closedir(self.0) }; + } +} + +fn set_errno(value: i32) { + unsafe { *libc::__errno_location() = value }; +} + +fn get_errno() -> i32 { + unsafe { *libc::__errno_location() } +} + +#[derive(Serialize, Deserialize)] +struct StoredSnapshot { + session_id: [u8; 16], + entries: Vec, + total_file_bytes: u64, +} + +#[derive(Serialize, Deserialize)] +struct StoredEntry { + path: String, + kind: SnapshotEntryKind, + mode: u16, + size: u64, + content_digest: Option<[u8; 32]>, +} + +impl StoredSnapshot { + fn from_entries(session_id: WorkspaceSessionId, entries: &[SnapshotEntry], total_file_bytes: u64) -> Self { + Self { + session_id: session_id.0, + entries: entries + .iter() + .map(|entry| StoredEntry { + path: entry.path.as_str().to_string(), + kind: entry.kind, + mode: entry.mode, + size: entry.size, + content_digest: entry.content_digest, + }) + .collect(), + total_file_bytes, + } + } + + fn into_entries(self) -> Result<(Vec, u64), String> { + if self.entries.len() > MAX_SNAPSHOT_ENTRIES { + return Err("stored snapshot exceeds maximum entry count".to_string()); + } + let mut total_file_bytes = 0u64; + let mut previous_path = None; + let entries = self + .entries + .into_iter() + .map(|entry| { + let path = SnapshotRelativePath::new(entry.path)?; + if previous_path.as_deref().is_some_and(|previous: &str| previous >= path.as_str()) { + return Err("snapshot manifest entries are not strictly ordered".to_string()); + } + previous_path = Some(path.as_str().to_string()); + if matches!(entry.kind, SnapshotEntryKind::Directory) && (entry.size != 0 || entry.content_digest.is_some()) { + return Err("invalid directory snapshot entry".to_string()); + } + if matches!(entry.kind, SnapshotEntryKind::RegularFile) && entry.content_digest.is_none() { + return Err("invalid regular-file snapshot entry".to_string()); + } + if matches!(entry.kind, SnapshotEntryKind::RegularFile) { + if entry.size > MAX_SNAPSHOT_FILE_BYTES { + return Err("stored snapshot file exceeds maximum size".to_string()); + } + total_file_bytes = total_file_bytes.checked_add(entry.size).ok_or_else(|| "snapshot manifest size overflow".to_string())?; + if total_file_bytes > MAX_SNAPSHOT_TOTAL_BYTES { + return Err("stored snapshot exceeds maximum total size".to_string()); + } + } + Ok(SnapshotEntry { path, kind: entry.kind, mode: entry.mode & 0o777, size: entry.size, content_digest: entry.content_digest }) + }) + .collect::, String>>()?; + if total_file_bytes != self.total_file_bytes { + return Err("snapshot manifest total size mismatch".to_string()); + } + Ok((entries, total_file_bytes)) + } +} + +fn snapshot_id(entries: &[SnapshotEntry]) -> SnapshotId { + let mut canonical = Vec::new(); + for entry in entries { + canonical.push(match entry.kind { + SnapshotEntryKind::Directory => 0, + SnapshotEntryKind::RegularFile => 1, + }); + canonical.extend_from_slice(&(entry.path.as_str().len() as u32).to_le_bytes()); + canonical.extend_from_slice(entry.path.as_str().as_bytes()); + canonical.extend_from_slice(&entry.mode.to_le_bytes()); + canonical.extend_from_slice(&entry.size.to_le_bytes()); + if let Some(digest) = entry.content_digest { + canonical.extend_from_slice(&digest); + } + } + SnapshotId(Sha256::digest(canonical).into()) +} + +fn hex(bytes: &[u8]) -> String { + bytes.iter().map(|byte| format!("{byte:02x}")).collect() +} + +#[cfg(test)] +#[path = "snapshot_ut.rs"] +mod tests; diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index 71f3474..50c1b07 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -69,6 +69,37 @@ pub struct RequestId(pub [u8; 16]); #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct WorkspaceSessionId(pub [u8; 16]); +impl WorkspaceSessionId { + pub fn from_hex(value: &str) -> Result { + if value.len() != 32 { + return Err("remote session ID must contain 32 hexadecimal characters".to_string()); + } + let mut bytes = [0u8; 16]; + for (index, byte) in bytes.iter_mut().enumerate() { + let high = hex_digit(value.as_bytes()[index * 2]).ok_or_else(|| "remote session ID is not hexadecimal".to_string())?; + let low = hex_digit(value.as_bytes()[index * 2 + 1]).ok_or_else(|| "remote session ID is not hexadecimal".to_string())?; + *byte = (high << 4) | low; + } + if bytes == [0; 16] { + return Err("remote session ID must be nonzero".to_string()); + } + Ok(Self(bytes)) + } + + pub fn to_hex(self) -> String { + self.0.iter().map(|byte| format!("{byte:02x}")).collect() + } +} + +fn hex_digit(value: u8) -> Option { + match value { + b'0'..=b'9' => Some(value - b'0'), + b'a'..=b'f' => Some(value - b'a' + 10), + b'A'..=b'F' => Some(value - b'A' + 10), + _ => None, + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct WorkspaceRelativePath(String); From a32c90d19247cf9ffaa02f2ce3f37da69c01d080 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 23:00:40 +0200 Subject: [PATCH 14/52] Fix regression: TUI rendering on a wrong channel before the first boot phase --- src/logging.rs | 32 +++++---- src/main.rs | 169 +++++++++++++++++++++++++++++++++-------------- src/tui.rs | 69 ++++++++++++------- src/tui/popup.rs | 3 + 4 files changed, 187 insertions(+), 86 deletions(-) diff --git a/src/logging.rs b/src/logging.rs index 868b6f4..e243936 100644 --- a/src/logging.rs +++ b/src/logging.rs @@ -56,6 +56,10 @@ pub fn set_status_fd(fd: RawFd) { STATUS_FD.with(|f| *f.borrow_mut() = Some(fd)); } +pub fn clear_status_fd() { + STATUS_FD.with(|f| *f.borrow_mut() = None); +} + /// Sends a password prompt to the TUI, blocks reading the response from the status fd. pub fn prompt_password(title: &str, prompt: &str) -> Result { let fd = STATUS_FD.with(|f| f.borrow().ok_or("status fd not set".to_string()))?; @@ -68,18 +72,22 @@ pub fn prompt_password(title: &str, prompt: &str) -> Result { libc::write(fd, buf.as_ptr() as *const libc::c_void, buf.len()); } - let mut response = Vec::new(); - let mut byte = [0u8; 1]; - loop { - let n = unsafe { libc::read(fd, byte.as_mut_ptr() as *mut libc::c_void, 1) }; - if n <= 0 { - return Err("failed to read password response".to_string()); - } - if byte[0] == b'\n' { - break; + let result = (|| { + let mut response = Vec::new(); + let mut byte = [0u8; 1]; + loop { + let n = unsafe { libc::read(fd, byte.as_mut_ptr() as *mut libc::c_void, 1) }; + if n <= 0 { + return Err("failed to read password response".to_string()); + } + if byte[0] == b'\n' { + break; + } + response.push(byte[0]); } - response.push(byte[0]); - } + + String::from_utf8(response).map_err(|e| format!("invalid password encoding: {e}")) + })(); let payload = crate::vscomm::encode_ui_payload("password", "hide", "", ""); let mut buf = b"@".to_vec(); @@ -89,7 +97,7 @@ pub fn prompt_password(title: &str, prompt: &str) -> Result { libc::write(fd, buf.as_ptr() as *const libc::c_void, buf.len()); } - String::from_utf8(response).map_err(|e| format!("invalid password encoding: {e}")) + result } pub fn log(msg: &str) { diff --git a/src/main.rs b/src/main.rs index 79a9087..d664b23 100644 --- a/src/main.rs +++ b/src/main.rs @@ -4,12 +4,13 @@ use rand::RngCore; use std::ffi::OsString; use std::fs::File; use std::io; -use std::os::fd::{FromRawFd, RawFd}; +use std::os::fd::{AsRawFd, FromRawFd, RawFd}; use std::os::unix::ffi::{OsStrExt, OsStringExt}; use std::path::{Path, PathBuf}; use std::sync::{Arc, Mutex}; const WORKSPACE_HANDOFF_MAGIC: &[u8; 4] = b"WS01"; +const STARTUP_READY_MAGIC: &[u8; 4] = b"RDY1"; const MAX_WORKSPACE_HANDOFF_BYTES: usize = 64 * 1024; fn main() { @@ -196,8 +197,6 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option> = Arc::new(Mutex::new(tui::OverlayState::new())); - let status_listener = start_status_listener(overlay.clone())?; - - let setup_result = (|| -> Result<(workspace::WorkspaceHandle, Arc, daemon::VsockDaemon), String> { - let workspace = workspace::resolve(workspace_mode, quota, exclude.as_deref(), &name)?; - let remote_session = new_session_id(); - let target = new_target_id(); - loopback::cleanup_stale_roots(&std::env::temp_dir())?; - let snapshot_root = std::env::temp_dir().join(format!("bunkerbox-snapshots-{}-{}", std::process::id(), remote_session.to_hex())); - let jobs_root = std::env::temp_dir().join(format!("bunkerbox-loopback-{}-{}", std::process::id(), remote_session.to_hex())); - let snapshot_store = snapshot::SnapshotStore::new(&snapshot_root); - let exclusion_policy = snapshot::SnapshotExclusionPolicy::from_config(&env, exclude.as_deref())?; - let snapshot_builder = snapshot::SnapshotBuilder::new(snapshot_store.clone(), snapshot::SnapshotLimits::default(), exclusion_policy); - let session = Arc::new(loopback::RunRemoteSession::new( - bunkerbox::remote::WorkspaceSessionId(remote_session.0), - target, - workspace.path().to_path_buf(), - snapshot_store, - snapshot_builder, - jobs_root, - )?); - let allowed_tools = remote_tool_names(&passthrough); - let tools = loopback::resolve_fixed_tools(allowed_tools.clone()); - let daemon = daemon::VsockDaemon::start_with_remote( - passthrough, - env_mode, - workspace.path().to_path_buf(), - profiles, - share_dir_owned, - merged_allow, - daemon::RemoteDaemonConfig::new(session.clone(), allowed_tools, tools), - )?; - if let Err(error) = write_run_handoff(&mut setup_parent, workspace.path(), remote_session) { - tokio::runtime::Handle::current().block_on(daemon.shutdown()); - return Err(error); + let mut startup_fds = [-1i32, -1]; + if unsafe { libc::pipe(startup_fds.as_mut_ptr()) } != 0 { + unsafe { + libc::kill(pid, libc::SIGTERM); + libc::waitpid(pid, std::ptr::null_mut(), 0); + libc::close(master); } - Ok((workspace, session, daemon)) - })(); + return Err(format!("startup status pipe: {}", std::io::Error::last_os_error())); + } + let (startup_status_read, startup_status_write) = (startup_fds[0], startup_fds[1]); - let (workspace, remote_session, daemon) = match setup_result { - Ok(value) => value, + let overlay: Arc> = Arc::new(Mutex::new(tui::OverlayState::new())); + let status_listener = match start_status_listener(overlay.clone()) { + Ok(listener) => listener, Err(error) => { unsafe { + libc::close(startup_status_read); + libc::close(startup_status_write); libc::kill(pid, libc::SIGTERM); libc::waitpid(pid, std::ptr::null_mut(), 0); libc::close(master); } - tokio::runtime::Handle::current().block_on(status_listener.shutdown()); return Err(error); } }; - drop(setup_parent); - let tui_result = tui::event_loop(master, rows, cols, parent_fd, overlay); + let setup_handle = tokio::runtime::Handle::current().clone(); + let setup_thread = + std::thread::spawn(move || -> Result<(workspace::WorkspaceHandle, Arc, daemon::VsockDaemon), String> { + let mut setup_parent = unsafe { File::from_raw_fd(setup_parent_fd) }; + let startup_status = unsafe { File::from_raw_fd(startup_status_write) }; + logging::set_status_fd(startup_status.as_raw_fd()); + let _runtime_guard = setup_handle.enter(); + + if let Err(error) = read_startup_ready(&mut setup_parent) { + logging::log(&format!("Startup failed: {error}")); + return Err(error); + } + + let setup_result = (|| -> Result<(workspace::WorkspaceHandle, Arc, daemon::VsockDaemon), String> { + logging::log("Preparing workspace..."); + let workspace = workspace::resolve(workspace_mode, quota, exclude.as_deref(), &name)?; + let remote_session = new_session_id(); + let target = new_target_id(); + logging::log("Preparing remote session..."); + loopback::cleanup_stale_roots(&std::env::temp_dir())?; + let snapshot_root = std::env::temp_dir().join(format!("bunkerbox-snapshots-{}-{}", std::process::id(), remote_session.to_hex())); + let jobs_root = std::env::temp_dir().join(format!("bunkerbox-loopback-{}-{}", std::process::id(), remote_session.to_hex())); + let snapshot_store = snapshot::SnapshotStore::new(&snapshot_root); + let exclusion_policy = snapshot::SnapshotExclusionPolicy::from_config(&env, exclude.as_deref())?; + let snapshot_builder = snapshot::SnapshotBuilder::new(snapshot_store.clone(), snapshot::SnapshotLimits::default(), exclusion_policy); + let session = Arc::new(loopback::RunRemoteSession::new( + bunkerbox::remote::WorkspaceSessionId(remote_session.0), + target, + workspace.path().to_path_buf(), + snapshot_store, + snapshot_builder, + jobs_root, + )?); + let allowed_tools = remote_tool_names(&passthrough); + let tools = loopback::resolve_fixed_tools(allowed_tools.clone()); + logging::log("Starting remote daemon..."); + let daemon = daemon::VsockDaemon::start_with_remote( + passthrough, + env_mode, + workspace.path().to_path_buf(), + profiles, + share_dir_owned, + merged_allow, + daemon::RemoteDaemonConfig::new(session.clone(), allowed_tools, tools), + )?; + if let Err(error) = write_run_handoff(&mut setup_parent, workspace.path(), remote_session) { + tokio::runtime::Handle::current().block_on(daemon.shutdown()); + return Err(error); + } + Ok((workspace, session, daemon)) + })(); + + if let Err(error) = &setup_result { + logging::log(&format!("Startup failed: {error}")); + } + setup_result + }); + + let tui_result = tui::event_loop(master, rows, cols, parent_fd, startup_status_read, overlay); + let tui_error = tui_result.err(); + if tui_error.is_some() { + unsafe { libc::kill(pid, libc::SIGTERM) }; + } + + let setup_result = match setup_thread.join() { + Ok(result) => result, + Err(_) => Err("workspace setup thread panicked".to_string()), + }; let mut status: i32 = 0; unsafe { libc::waitpid(pid, &mut status, 0) }; unsafe { libc::close(master) }; + unsafe { libc::close(startup_status_read) }; - tokio::runtime::Handle::current().block_on(daemon.shutdown()); - drop(remote_session); - drop(workspace); + let (setup_state, setup_error) = match setup_result { + Ok((workspace, remote_session, daemon)) => { + tokio::runtime::Handle::current().block_on(daemon.shutdown()); + drop(remote_session); + drop(workspace); + (true, None) + } + Err(error) => (false, Some(error)), + }; tokio::runtime::Handle::current().block_on(status_listener.shutdown()); - tui_result?; + if let Some(error) = tui_error { + return Err(error); + } + if let Some(error) = setup_error { + return Err(error); + } + debug_assert!(setup_state); if status != 0 { return Err(format!("child exited with status {status}")); } @@ -384,6 +440,19 @@ fn write_run_handoff(file: &mut File, path: &Path, session_id: vscomm::Workspace io::Write::write_all(file, &frame).map_err(|err| format!("write run handoff: {err}")) } +fn write_startup_ready(file: &mut File) -> Result<(), String> { + io::Write::write_all(file, STARTUP_READY_MAGIC).map_err(|err| format!("write startup readiness: {err}")) +} + +fn read_startup_ready(file: &mut File) -> Result<(), String> { + let mut ready = [0u8; STARTUP_READY_MAGIC.len()]; + io::Read::read_exact(file, &mut ready).map_err(|err| format!("read startup readiness: {err}"))?; + if &ready != STARTUP_READY_MAGIC { + return Err("startup readiness has an invalid type".to_string()); + } + Ok(()) +} + fn ensure_sudo() -> Result<(), String> { if !std::process::Command::new("sudo") .arg("-n") diff --git a/src/tui.rs b/src/tui.rs index 2aefed8..930b574 100644 --- a/src/tui.rs +++ b/src/tui.rs @@ -550,7 +550,9 @@ fn screen_has_ascii_alphanumeric(screen: &vt100::Screen) -> bool { /// /// `overlay` is shared with the VSOCK status listener so VM-originated /// UI commands can update popups, progress bars, and status text. -pub fn event_loop(master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, overlay: Arc>) -> Result<(), String> { +pub fn event_loop( + master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, startup_status_fd: RawFd, overlay: Arc>, +) -> Result<(), String> { let stdin_fd = io::stdin().as_raw_fd(); let mut stdout = io::stdout(); @@ -569,6 +571,8 @@ pub fn event_loop(master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, over let mut last_cols = cols; let mut status_buf = Vec::new(); + let mut startup_status_buf = Vec::new(); + let mut startup_status_fd = startup_status_fd; let mut mouse_capture_enabled = false; unsafe { @@ -594,9 +598,10 @@ pub fn event_loop(master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, over libc::pollfd { fd: master_fd, events: libc::POLLIN, revents: 0 }, libc::pollfd { fd: stdin_fd, events: libc::POLLIN, revents: 0 }, libc::pollfd { fd: status_fd, events: libc::POLLIN, revents: 0 }, + libc::pollfd { fd: startup_status_fd, events: libc::POLLIN, revents: 0 }, ]; - let ret = unsafe { libc::poll(fds.as_mut_ptr(), 3, 16) }; + let ret = unsafe { libc::poll(fds.as_mut_ptr(), 4, 16) }; if ret == -1 { let err = io::Error::last_os_error(); @@ -638,6 +643,16 @@ pub fn event_loop(master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, over } } + if fds[2].revents & (libc::POLLIN | libc::POLLHUP) != 0 { + read_status_messages(status_fd, &mut status_buf, &overlay); + } + if startup_status_fd >= 0 + && fds[3].revents & (libc::POLLIN | libc::POLLHUP) != 0 + && read_status_messages(startup_status_fd, &mut startup_status_buf, &overlay) == 0 + { + startup_status_fd = -1; + } + if fds[1].revents & libc::POLLIN != 0 { if let Ok(input_event) = event::read() { match input_event { @@ -667,28 +682,6 @@ pub fn event_loop(master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, over } } - if fds[2].revents & (libc::POLLIN | libc::POLLHUP) != 0 { - let mut chunk = [0u8; 256]; - let n = unsafe { libc::read(status_fd, chunk.as_mut_ptr() as *mut libc::c_void, chunk.len()) }; - if n > 0 { - status_buf.extend_from_slice(&chunk[..n as usize]); - } - while let Some(pos) = status_buf.iter().position(|&b| b == b'\n') { - let line = String::from_utf8_lossy(&status_buf[..pos]).into_owned(); - status_buf.drain(..=pos); - if let Some(cmd) = line.strip_prefix('@') { - if let Some((widget, cmd, opts, val)) = vscomm::decode_ui_payload(cmd.as_bytes()) { - let mut state = overlay.lock().unwrap(); - dispatch_ui_command(&mut state, widget, cmd, opts, val); - } - } else { - let mut state = overlay.lock().unwrap(); - let title = state.popup_title.clone(); - state.popup.show_info(title, &line, Some(palette::FG), Some(palette::ACCENT)); - } - } - } - { let mut state = overlay.lock().unwrap(); let now = Instant::now(); @@ -756,6 +749,34 @@ pub fn event_loop(master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, over Ok(()) } +fn read_status_messages(fd: RawFd, buffer: &mut Vec, overlay: &Arc>) -> isize { + let mut chunk = [0u8; 256]; + let n = unsafe { libc::read(fd, chunk.as_mut_ptr() as *mut libc::c_void, chunk.len()) }; + if n <= 0 { + return 0; + } + process_status_bytes(buffer, &chunk[..n as usize], overlay); + n +} + +fn process_status_bytes(buffer: &mut Vec, bytes: &[u8], overlay: &Arc>) { + buffer.extend_from_slice(bytes); + while let Some(pos) = buffer.iter().position(|&byte| byte == b'\n') { + let line = String::from_utf8_lossy(&buffer[..pos]).into_owned(); + buffer.drain(..=pos); + if let Some(command) = line.strip_prefix('@') { + if let Some((widget, command, options, value)) = vscomm::decode_ui_payload(command.as_bytes()) { + let mut state = overlay.lock().unwrap(); + dispatch_ui_command(&mut state, widget, command, options, value); + } + } else { + let mut state = overlay.lock().unwrap(); + let title = state.popup_title.clone(); + state.popup.show_info(title, &line, Some(palette::FG), Some(palette::ACCENT)); + } + } +} + fn cleanup_terminal(terminal: &mut Terminal>, mouse_capture_enabled: bool) { if mouse_capture_enabled { terminal.backend_mut().execute(DisableMouseCapture).ok(); diff --git a/src/tui/popup.rs b/src/tui/popup.rs index c1f27b4..fea682e 100644 --- a/src/tui/popup.rs +++ b/src/tui/popup.rs @@ -93,6 +93,9 @@ impl PopupWidget { pub fn hide(&mut self) { self.visible = false; + if matches!(&self.content, PopupContent::Password { .. }) { + self.content = PopupContent::Info { title: None, message: String::new(), fg: palette::FG }; + } } pub fn handle_password_key(&mut self, key: &crossterm::event::KeyEvent) { From d0080dbe2069e812080689cf74d9593050bf6490 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Sun, 2 Aug 2026 23:00:55 +0200 Subject: [PATCH 15/52] Add UT for TUI rendering --- src/logging_ut.rs | 30 +++++++++++++++++++++++++++++- src/main_ut.rs | 27 ++++++++++++++++++++++++++- src/tui_ut.rs | 27 ++++++++++++++++++++++++++- 3 files changed, 81 insertions(+), 3 deletions(-) diff --git a/src/logging_ut.rs b/src/logging_ut.rs index e216451..de3eef6 100644 --- a/src/logging_ut.rs +++ b/src/logging_ut.rs @@ -1,5 +1,8 @@ -use super::{configure, diagnostic, diagnostic_bytes}; +use super::{clear_status_fd, configure, diagnostic, diagnostic_bytes, prompt_password, set_status_fd}; use std::fs; +use std::io::{BufRead, BufReader, Write}; +use std::os::fd::{AsRawFd, FromRawFd}; +use std::thread; #[test] fn diagnostics_write_to_file_without_terminal_output() { @@ -16,3 +19,28 @@ fn diagnostics_write_to_file_without_terminal_output() { configure(false, None); } + +#[test] +fn password_prompt_uses_one_response_channel_and_hides_on_completion() { + let mut fds = [-1; 2]; + assert_eq!(unsafe { libc::socketpair(libc::AF_UNIX, libc::SOCK_STREAM, 0, fds.as_mut_ptr()) }, 0); + let input = unsafe { std::fs::File::from_raw_fd(fds[0]) }; + let mut peer = unsafe { std::fs::File::from_raw_fd(fds[1]) }; + let peer_thread = thread::spawn(move || { + let mut reader = BufReader::new(peer.try_clone().unwrap()); + let mut show = String::new(); + reader.read_line(&mut show).unwrap(); + assert!(show.contains("password")); + peer.write_all(b"secret\n").unwrap(); + let mut hide = String::new(); + reader.read_line(&mut hide).unwrap(); + assert!(hide.contains("password")); + assert!(!hide.contains("secret")); + }); + + set_status_fd(input.as_raw_fd()); + assert_eq!(prompt_password("Password", "Enter password").unwrap(), "secret"); + clear_status_fd(); + drop(input); + peer_thread.join().unwrap(); +} diff --git a/src/main_ut.rs b/src/main_ut.rs index b1e0af1..cb5108a 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -1,4 +1,7 @@ -use super::{decode_workspace_handoff, encode_workspace_handoff, read_run_handoff, remote_tool_names, write_run_handoff}; +use super::{ + decode_workspace_handoff, encode_workspace_handoff, read_run_handoff, read_startup_ready, remote_tool_names, write_run_handoff, + write_startup_ready, +}; use bunkerbox::vscomm::WorkspaceSessionId; use std::fs::File; use std::os::fd::FromRawFd; @@ -58,3 +61,25 @@ fn run_handoff_rejects_zero_session() { fn remote_tool_names_reduce_passthrough_entries_to_executables() { assert_eq!(remote_tool_names(&["make *".into(), "cargo build".into(), "make test".into()]), vec!["cargo", "make"]); } + +#[test] +fn startup_ready_handoff_round_trips() { + let (parent, child) = unsafe { + let mut fds = [-1; 2]; + assert_eq!(libc::pipe(fds.as_mut_ptr()), 0); + (File::from_raw_fd(fds[0]), File::from_raw_fd(fds[1])) + }; + let mut child = child; + write_startup_ready(&mut child).unwrap(); + drop(child); + let mut parent = parent; + read_startup_ready(&mut parent).unwrap(); +} + +#[test] +fn startup_ready_handoff_rejects_wrong_type() { + let file = tempfile::NamedTempFile::new().unwrap(); + std::fs::write(file.path(), b"BAD!").unwrap(); + let mut file = File::open(file.path()).unwrap(); + assert!(read_startup_ready(&mut file).is_err()); +} diff --git a/src/tui_ut.rs b/src/tui_ut.rs index e2cc37e..22c7e4f 100644 --- a/src/tui_ut.rs +++ b/src/tui_ut.rs @@ -1,5 +1,6 @@ -use super::{dispatch_ui_command, mouse_to_bytes, MouseEncoding, MouseTracking, OverlayState, Term}; +use super::{dispatch_ui_command, mouse_to_bytes, process_status_bytes, MouseEncoding, MouseTracking, OverlayState, Term}; use crossterm::event::{KeyModifiers, MouseButton, MouseEvent, MouseEventKind}; +use std::sync::{Arc, Mutex}; #[test] fn internal_error_creates_a_non_modal_toast() { @@ -14,6 +15,30 @@ fn internal_error_creates_a_non_modal_toast() { assert_eq!(toast.message, "connection failed"); } +#[test] +fn startup_status_is_processed_before_child_status() { + let overlay = Arc::new(Mutex::new(OverlayState::new())); + let mut buffer = Vec::new(); + let mut message = b"@".to_vec(); + message.extend_from_slice(&crate::vscomm::encode_ui_payload("status", "set", "", "Preparing workspace...")); + message.push(b'\n'); + + let split = message.len() / 2; + process_status_bytes(&mut buffer, &message[..split], &overlay); + assert!(!overlay.lock().unwrap().popup.visible); + process_status_bytes(&mut buffer, &message[split..], &overlay); + assert!(overlay.lock().unwrap().popup.visible); +} + +#[test] +fn hiding_password_clears_sensitive_popup_state() { + let mut state = OverlayState::new(); + dispatch_ui_command(&mut state, "password", "show", "Password", "Enter password"); + assert!(state.popup.password_value().is_some()); + state.popup.hide(); + assert!(state.popup.password_value().is_none()); +} + #[test] fn cursor_report_uses_position_after_prior_bytes_in_same_chunk() { let mut term = Term::new(24, 80); From 93eb8e71c4048bdd20974283ad88bf7c2be1ef12 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 01:18:03 +0200 Subject: [PATCH 16/52] ADd unit tests for strict remote env and configurable policies --- src/cfg_ut.rs | 19 +++++++++++++ src/loopback_ut.rs | 53 +++++++++++++++++++++++++++++++------ src/main_ut.rs | 8 ++++-- src/remote_ut.rs | 63 +++++++++++++++++++++++++++++++++++++++++++- src/vscomm/mod_ut.rs | 6 +++++ 5 files changed, 138 insertions(+), 11 deletions(-) diff --git a/src/cfg_ut.rs b/src/cfg_ut.rs index e31f48b..d932fbf 100644 --- a/src/cfg_ut.rs +++ b/src/cfg_ut.rs @@ -487,3 +487,22 @@ fn load_or_create_accepts_paranoid_exact() { write_project_conf(root.path(), "project:\n env: paranoid\n passthrough:\n - \"make\"\n - \"cargo\"\n"); assert!(ProjectConfig::load_or_create(root.path()).is_ok()); } + +#[test] +fn load_or_create_validates_remote_policy_configuration() { + let root = TempDir::new().unwrap(); + write_project_conf( + root.path(), + "project:\n remote:\n environment:\n - PROJECT_MODE\n tools:\n - name: make\n allow-args: false\n", + ); + let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); + assert_eq!(cfg.project.remote.environment, vec!["PROJECT_MODE"]); + assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "make".into(), allow_args: false }]); +} + +#[test] +fn load_or_create_rejects_forbidden_remote_environment() { + let root = TempDir::new().unwrap(); + write_project_conf(root.path(), "project:\n remote:\n environment:\n - SSH_AUTH_SOCK\n"); + assert!(ProjectConfig::load_or_create(root.path()).is_err()); +} diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 4da344f..21cb637 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -63,15 +63,21 @@ async fn sync_and_build_materialize_a_bound_snapshot() { assert!(fs::read_dir(&session_state.jobs_root).unwrap().next().is_none()); } -#[tokio::test(flavor = "multi_thread", worker_threads = 2)] -async fn nonempty_remote_environment_is_rejected() { - let (_temp, session, target, session_id) = fixture(); - let backend = LoopbackBackend::new(session, BTreeMap::new()); - let (events, _receiver) = mpsc::channel(8); - let request = authorized_build(target, session_id, "printf", Vec::new(), vec![("UNTRUSTED".to_string(), "1".to_string())]); +#[test] +fn unapproved_remote_environment_is_rejected_before_backend_execution() { + let (_temp, _session, target, session_id) = fixture(); + let build = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("printf").unwrap(), + Vec::new(), + vec![("UNTRUSTED".to_string(), "1".to_string())], + ) + .unwrap(); + let request = RemoteRequest::build(RequestId([3; 16]), session_id, build); + let policy = RemoteAuthorizationPolicy::new(target, session_id, vec!["printf".to_string()]); assert_eq!( - backend.execute(request, events).await, - Err(RemoteBackendError::Failed("remote environment is unavailable until remote environment policy is configured".to_string(),)) + policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, request), + Err(crate::remote::RemoteAuthorizationError::EnvironmentNotAllowed("UNTRUSTED".to_string())) ); } @@ -92,6 +98,23 @@ async fn build_preserves_argument_bytes_and_reports_nonzero_stderr() { assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code } if *exit_code != 0))); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn child_receives_guest_environment_but_trusted_target_wins() { + let (_temp, session, target, session_id) = fixture(); + session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["printenv".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session, tools).with_target_environment(BTreeMap::from([("CC".into(), "trusted-target".into())])); + let (events, receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "printenv", vec!["CC".into()], vec![("CC".into(), "guest-value".into())]); + assert_eq!(backend.execute(request, events).await, Ok(())); + let events = collect_events(receiver).await; + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Stdout(bytes) if bytes == b"trusted-target\n"))); + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn missing_tool_fails_before_execution() { let (_temp, session, target, session_id) = fixture(); @@ -114,3 +137,17 @@ async fn timeout_kills_a_direct_child_process() { let request = authorized_build(target, session_id, "sleep", vec!["5".to_string()], Vec::new()); assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Timeout)); } + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn output_limit_kills_a_flooding_direct_child() { + let (_temp, session, target, session_id) = fixture(); + session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["printf".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session, tools).with_output_limit(8); + let (events, _receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "printf", vec!["0123456789".into()], Vec::new()); + assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::OutputLimit { limit: 8 })); +} diff --git a/src/main_ut.rs b/src/main_ut.rs index cb5108a..2affc20 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -2,6 +2,7 @@ use super::{ decode_workspace_handoff, encode_workspace_handoff, read_run_handoff, read_startup_ready, remote_tool_names, write_run_handoff, write_startup_ready, }; +use bunkerbox::cfg::RemoteToolSpec; use bunkerbox::vscomm::WorkspaceSessionId; use std::fs::File; use std::os::fd::FromRawFd; @@ -58,8 +59,11 @@ fn run_handoff_rejects_zero_session() { } #[test] -fn remote_tool_names_reduce_passthrough_entries_to_executables() { - assert_eq!(remote_tool_names(&["make *".into(), "cargo build".into(), "make test".into()]), vec!["cargo", "make"]); +fn remote_tool_names_preserve_configured_order() { + assert_eq!( + remote_tool_names(&[RemoteToolSpec { name: "make".into(), allow_args: true }, RemoteToolSpec { name: "cargo".into(), allow_args: false },]), + vec!["make", "cargo"] + ); } #[test] diff --git a/src/remote_ut.rs b/src/remote_ut.rs index 4066e22..8e061f5 100644 --- a/src/remote_ut.rs +++ b/src/remote_ut.rs @@ -8,7 +8,7 @@ fn request(tool: &str) -> RemoteRequest { WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new(tool).unwrap(), vec!["build".into()], - vec![("MODE".into(), "debug".into())], + vec![("CC".into(), "cc".into())], ) .unwrap(), ) @@ -41,3 +41,64 @@ fn backend_errors_have_typed_events() { assert_eq!(RemoteBackendError::Timeout.event(), RemoteBackendEvent::Error { message: "remote backend timed out".into() }); assert_eq!(RemoteBackendError::Cancelled.event(), RemoteBackendEvent::Cancelled); } + +#[test] +fn environment_policy_preserves_allowed_entries() { + let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let authorized = policy.authorize(&context(), request("make")).unwrap(); + let RemoteOperation::Build(build) = authorized.request().operation() else { panic!("expected build") }; + assert_eq!(build.env(), [("CC".into(), "cc".into())]); +} + +#[test] +fn environment_policy_rejects_forbidden_and_unlisted_entries() { + for (name, expected) in [ + ("SSH_AUTH_SOCK", RemoteAuthorizationError::ForbiddenEnvironment("SSH_AUTH_SOCK".into())), + ("UNTRUSTED", RemoteAuthorizationError::EnvironmentNotAllowed("UNTRUSTED".into())), + ] { + let build = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("make").unwrap(), + Vec::new(), + vec![(name.into(), "value".into())], + ) + .unwrap(); + let request = RemoteRequest::build(RequestId([1; 16]), WorkspaceSessionId([2; 16]), build); + let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + assert_eq!(policy.authorize(&context(), request), Err(expected)); + } +} + +#[test] +fn environment_policy_rejects_duplicates_and_control_data() { + let duplicate = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("make").unwrap(), + Vec::new(), + vec![("CC".into(), "one".into()), ("CC".into(), "two".into())], + ) + .unwrap(); + let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let request = RemoteRequest::build(RequestId([1; 16]), WorkspaceSessionId([2; 16]), duplicate); + assert_eq!(policy.authorize(&context(), request), Err(RemoteAuthorizationError::DuplicateEnvironment("CC".into()))); + + assert!(RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("make").unwrap(), + Vec::new(), + vec![("CC".into(), "bad\nvalue".into())], + ) + .is_err()); +} + +#[test] +fn command_policy_rejects_unapproved_arguments() { + let policy = RemoteAuthorizationPolicy::from_policies( + RemoteTargetId([3; 16]), + WorkspaceSessionId([2; 16]), + [("make".into(), RemoteToolPolicy::new(false))], + RemoteEnvironmentPolicy::default(), + ) + .unwrap(); + assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::ToolArgumentsNotAllowed("make".into()))); +} diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index 00f0f7f..e7e021e 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -193,6 +193,12 @@ fn oversized_environment_key_and_value_are_rejected() { assert!(RemoteRequest::from_frame(raw_build_frame(0, None, 1, Some(("KEY", &"V".repeat(MAX_REMOTE_ENV_VALUE_BYTES + 1))))).is_err()); } +#[test] +fn control_data_in_remote_environment_value_is_rejected() { + assert!(build(Vec::new(), vec![("CC".into(), "bad\nvalue".into())]).is_err()); + assert!(RemoteRequest::from_frame(raw_build_frame(0, None, 1, Some(("CC", "bad\nvalue")))).is_err()); +} + #[test] fn invalid_remote_cwd_is_rejected() { assert!(WorkspaceRelativePath::new("/absolute").is_err()); From 8c866fd02a661fb62187a6df7da549e618653a88 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 01:18:14 +0200 Subject: [PATCH 17/52] Add strict remote env and configurable policies --- src/cfg.rs | 55 ++++++++++- src/cfgsetup.rs | 7 +- src/daemon.rs | 52 +++++++++-- src/loopback.rs | 67 +++++++++---- src/main.rs | 23 ++--- src/remote.rs | 232 +++++++++++++++++++++++++++++++++++++++++++--- src/vscomm/mod.rs | 18 +++- 7 files changed, 397 insertions(+), 57 deletions(-) diff --git a/src/cfg.rs b/src/cfg.rs index 89b5e9c..b308273 100644 --- a/src/cfg.rs +++ b/src/cfg.rs @@ -4,6 +4,7 @@ use std::path::{Path, PathBuf}; use serde::Deserialize; +use crate::remote::{RemoteEnvironmentPolicy, RemoteTool}; use crate::vscomm::buildsys::{self, PassthroughMode}; pub const DEFAULT_SHARE_DIR: &str = "/usr/share/bunkerbox"; @@ -196,6 +197,23 @@ pub struct ProjectSection { pub exclude: Vec, #[serde(default)] pub passthrough: Vec, + #[serde(default)] + pub remote: RemoteSection, +} + +#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)] +pub struct RemoteSection { + #[serde(default)] + pub environment: Vec, + #[serde(default)] + pub tools: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +pub struct RemoteToolSpec { + pub name: String, + #[serde(default, rename = "allow-args")] + pub allow_args: bool, } #[derive(Debug, Default, serde::Serialize, serde::Deserialize)] @@ -238,6 +256,7 @@ impl ProjectConfig { quota: Some("auto".into()), exclude: Vec::new(), passthrough: buildsys::scan(repo_root, PassthroughMode::Relaxed), + remote: RemoteSection::default(), }, image: ImageOverrides::default(), profiles: Vec::new(), @@ -258,6 +277,14 @@ impl ProjectConfig { } } } + RemoteEnvironmentPolicy::from_names(self.project.remote.environment.clone())?; + let mut tools = std::collections::BTreeSet::new(); + for tool in &self.project.remote.tools { + RemoteTool::new(tool.name.clone())?; + if !tools.insert(tool.name.clone()) { + return Err(format!("duplicate remote tool: {}", tool.name)); + } + } Ok(()) } @@ -277,7 +304,13 @@ impl ProjectConfig { .map_err(|e| format!("failed to parse legacy {}: {e}", legacy_path.display()))?; let cfg = ProjectConfig { - project: ProjectSection { env: EnvMode::default(), quota: old.quota, exclude: old.exclude, passthrough: old.passthrough }, + project: ProjectSection { + env: EnvMode::default(), + quota: old.quota, + exclude: old.exclude, + passthrough: old.passthrough, + remote: RemoteSection::default(), + }, image: ImageOverrides::default(), profiles: Vec::new(), }; @@ -348,6 +381,26 @@ impl ProjectConfig { } } + if !self.project.remote.environment.is_empty() || !self.project.remote.tools.is_empty() { + y.push_str(" remote:\n"); + y.push_str(" environment:\n"); + if self.project.remote.environment.is_empty() { + y.push_str(" []\n"); + } else { + for name in &self.project.remote.environment { + y.push_str(&format!(" - \"{name}\"\n")); + } + } + y.push_str(" tools:\n"); + if self.project.remote.tools.is_empty() { + y.push_str(" []\n"); + } else { + for tool in &self.project.remote.tools { + y.push_str(&format!(" - name: \"{}\"\n allow-args: {}\n", tool.name, tool.allow_args)); + } + } + } + if self.image.has_override() { y.push('\n'); y.push_str("# Override shared runtime defaults:\n"); diff --git a/src/cfgsetup.rs b/src/cfgsetup.rs index aea70d5..bf2650a 100644 --- a/src/cfgsetup.rs +++ b/src/cfgsetup.rs @@ -29,8 +29,11 @@ pub fn run(runtime: Option<&RuntimeConfig>) -> Result<(), String> { let profiles = pick_profiles(&detected_relaxed)?; let overrides = pick_overrides(runtime)?; - let cfg = - ProjectConfig { project: ProjectSection { env: env_mode, quota: Some(quota), exclude: Vec::new(), passthrough }, image: overrides, profiles }; + let cfg = ProjectConfig { + project: ProjectSection { env: env_mode, quota: Some(quota), exclude: Vec::new(), passthrough, remote: Default::default() }, + image: overrides, + profiles, + }; let path = repo_root.join(ProjectConfig::PATH); std::fs::create_dir_all(path.parent().unwrap()).map_err(|e| format!("failed to create {}: {e}", path.parent().unwrap().display()))?; diff --git a/src/daemon.rs b/src/daemon.rs index 1d35fc2..34a8449 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -3,7 +3,8 @@ use crate::logging; use crate::loopback::{LoopbackBackend, RunRemoteSession}; use crate::proxy::{FilterProxy, UnixProxyHandle}; use crate::remote::{ - RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionContext, RemoteRequest, + RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteEnvironmentPolicy, + RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteToolPolicy, }; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; use crate::vscomm::{validate_exec_request, validate_process_path, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; @@ -95,12 +96,40 @@ pub struct VsockDaemon { pub struct RemoteDaemonConfig { session: Arc, allowed_tools: Vec, + tool_policies: Option>, + environment: Option, tools: std::collections::BTreeMap, + target_environment: std::collections::BTreeMap, + resources: RemoteResourcePolicy, } impl RemoteDaemonConfig { pub fn new(session: Arc, allowed_tools: Vec, tools: std::collections::BTreeMap) -> Self { - Self { session, allowed_tools, tools } + Self { + session, + allowed_tools, + tool_policies: None, + environment: None, + tools, + target_environment: std::collections::BTreeMap::new(), + resources: RemoteResourcePolicy::default(), + } + } + + pub fn with_policy(mut self, tools: Vec<(String, RemoteToolPolicy)>, environment: RemoteEnvironmentPolicy) -> Self { + self.tool_policies = Some(tools); + self.environment = Some(environment); + self + } + + pub fn with_target_environment(mut self, environment: std::collections::BTreeMap) -> Self { + self.target_environment = environment; + self + } + + pub fn with_resources(mut self, resources: RemoteResourcePolicy) -> Self { + self.resources = resources; + self } } @@ -115,13 +144,20 @@ impl VsockDaemon { passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, remote: RemoteDaemonConfig, ) -> Result { - let remote_policy = RemoteAuthorizationPolicy::new(remote.session.target(), remote.session.session_id(), remote.allowed_tools); - let remote_context = RemoteExecutionContext { target: remote.session.target(), workspace_session_id: remote.session.session_id() }; - let remote_components = RemoteComponents { - context: remote_context, - policy: remote_policy, - backend: Arc::new(LoopbackBackend::new(remote.session, remote.tools)), + let RemoteDaemonConfig { session, allowed_tools, tool_policies, environment, tools, target_environment, resources } = remote; + let remote_policy = match (tool_policies, environment) { + (Some(tool_policies), Some(environment)) => { + RemoteAuthorizationPolicy::from_policies(session.target(), session.session_id(), tool_policies, environment)? + } + (None, None) => RemoteAuthorizationPolicy::new(session.target(), session.session_id(), allowed_tools), + _ => return Err("remote tool and environment policies must be configured together".to_string()), }; + let remote_context = RemoteExecutionContext { target: session.target(), workspace_session_id: session.session_id() }; + let backend = LoopbackBackend::new(session, tools) + .with_target_environment(target_environment) + .with_timeout(resources.build_timeout) + .with_output_limit(resources.max_output_bytes); + let remote_components = RemoteComponents { context: remote_context, policy: remote_policy, backend: Arc::new(backend) }; Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, remote_components) } diff --git a/src/loopback.rs b/src/loopback.rs index 5c938e7..39362ae 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -1,5 +1,6 @@ use crate::remote::{ - AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteTargetId, WorkspaceSessionId, + AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteResourcePolicy, + RemoteTargetId, WorkspaceSessionId, }; use crate::snapshot::{SnapshotBuilder, SnapshotHandle, SnapshotStore}; use std::collections::BTreeMap; @@ -15,7 +16,7 @@ use tokio::process::Command; use tokio::sync::mpsc; use tokio::time::sleep; -pub const LOOPBACK_BUILD_TIMEOUT: Duration = Duration::from_secs(30); +pub const LOOPBACK_BUILD_TIMEOUT: Duration = crate::remote::DEFAULT_REMOTE_BUILD_TIMEOUT; const LOOPBACK_PATH: &str = "/usr/local/bin:/usr/bin:/bin"; static NEXT_JOB_ID: AtomicU64 = AtomicU64::new(1); @@ -124,16 +125,34 @@ impl RunRemoteSession { pub struct LoopbackBackend { session: Arc, tools: Arc>, - timeout: Duration, + target_environment: Arc>, + resources: RemoteResourcePolicy, } impl LoopbackBackend { pub fn new(session: Arc, tools: BTreeMap) -> Self { - Self { session, tools: Arc::new(tools), timeout: LOOPBACK_BUILD_TIMEOUT } + Self { + session, + tools: Arc::new(tools), + target_environment: Arc::new(trusted_target_environment()), + resources: RemoteResourcePolicy::default(), + } } pub fn with_timeout(mut self, timeout: Duration) -> Self { - self.timeout = timeout; + self.resources.build_timeout = timeout; + self + } + + pub fn with_output_limit(mut self, max_output_bytes: u64) -> Self { + self.resources.max_output_bytes = max_output_bytes; + self + } + + pub fn with_target_environment(mut self, environment: BTreeMap) -> Self { + let mut trusted = trusted_target_environment(); + trusted.extend(environment); + self.target_environment = Arc::new(trusted); self } } @@ -144,11 +163,12 @@ impl RemoteBackend for LoopbackBackend { ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { let session = self.session.clone(); let tools = self.tools.clone(); - let timeout = self.timeout; + let target_environment = self.target_environment.clone(); + let resources = self.resources; Box::pin(async move { match request.request().operation() { RemoteOperation::Sync => execute_sync(session, events).await, - RemoteOperation::Build(build) => execute_build(session, tools, timeout, build, events).await, + RemoteOperation::Build(build) => execute_build(session, tools, target_environment, resources, build, events).await, } }) } @@ -164,12 +184,9 @@ async fn execute_sync(session: Arc, events: mpsc::Sender, tools: Arc>, timeout_duration: Duration, build: &crate::remote::RemoteBuild, - events: mpsc::Sender, + session: Arc, tools: Arc>, target_environment: Arc>, + resources: RemoteResourcePolicy, build: &crate::remote::RemoteBuild, events: mpsc::Sender, ) -> Result<(), RemoteBackendError> { - if !build.env().is_empty() { - return Err(RemoteBackendError::Failed("remote environment is unavailable until remote environment policy is configured".to_string())); - } let executable = tools .get(build.tool().as_str()) .ok_or_else(|| RemoteBackendError::Spawn(format!("loopback tool is not configured: {}", build.tool().as_str())))?; @@ -197,7 +214,13 @@ async fn execute_build( .stdin(std::process::Stdio::null()) .stdout(std::process::Stdio::piped()) .stderr(std::process::Stdio::piped()); - command.env_clear().env("PATH", LOOPBACK_PATH); + command.env_clear(); + for (key, value) in build.env() { + command.env(key, value); + } + for (key, value) in target_environment.iter() { + command.env(key, value); + } unsafe { command.pre_exec(|| { if libc::setpgid(0, 0) != 0 { @@ -211,10 +234,12 @@ async fn execute_build( let process_group = child.id().map(|pid| ProcessGroupGuard { pgid: pid as i32, active: true }); let stdout = child.stdout.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stdout".to_string()))?; let stderr = child.stderr.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stderr".to_string()))?; - let mut stdout_task = Box::pin(tokio::spawn(pump(stdout, RemoteStream::Stdout, events.clone()))); - let mut stderr_task = Box::pin(tokio::spawn(pump(stderr, RemoteStream::Stderr, events.clone()))); + let output_bytes = Arc::new(AtomicU64::new(0)); + let mut stdout_task = + Box::pin(tokio::spawn(pump(stdout, RemoteStream::Stdout, events.clone(), output_bytes.clone(), resources.max_output_bytes))); + let mut stderr_task = Box::pin(tokio::spawn(pump(stderr, RemoteStream::Stderr, events.clone(), output_bytes, resources.max_output_bytes))); let mut child_wait = Box::pin(child.wait()); - let mut timeout_sleep = Box::pin(sleep(timeout_duration)); + let mut timeout_sleep = Box::pin(sleep(resources.build_timeout)); let mut child_status = None; let mut stdout_done = false; let mut stderr_done = false; @@ -260,7 +285,7 @@ enum RemoteStream { } async fn pump( - mut reader: R, stream: RemoteStream, events: mpsc::Sender, + mut reader: R, stream: RemoteStream, events: mpsc::Sender, output_bytes: Arc, max_output_bytes: u64, ) -> Result<(), RemoteBackendError> { let mut buffer = [0u8; 8192]; loop { @@ -268,6 +293,10 @@ async fn pump( if count == 0 { return Ok(()); } + let total = output_bytes.fetch_add(count as u64, Ordering::Relaxed).saturating_add(count as u64); + if total > max_output_bytes { + return Err(RemoteBackendError::OutputLimit { limit: max_output_bytes }); + } let event = match stream { RemoteStream::Stdout => RemoteBackendEvent::Stdout(buffer[..count].to_vec()), RemoteStream::Stderr => RemoteBackendEvent::Stderr(buffer[..count].to_vec()), @@ -276,6 +305,10 @@ async fn pump( } } +fn trusted_target_environment() -> BTreeMap { + BTreeMap::from([(String::from("PATH"), String::from(LOOPBACK_PATH))]) +} + async fn join_pump(result: Result, tokio::task::JoinError>) -> Result<(), RemoteBackendError> { result.map_err(|error| RemoteBackendError::Failed(format!("loopback output task failed: {error}")))? } diff --git a/src/main.rs b/src/main.rs index d664b23..137a494 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,4 +1,5 @@ -use bunkerbox::cfg::{ProjectConfig, WorkspaceMode}; +use bunkerbox::cfg::{ProjectConfig, RemoteToolSpec, WorkspaceMode}; +use bunkerbox::remote::{RemoteEnvironmentPolicy, RemoteToolPolicy}; use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, loopback, overlay, snapshot, tui, vscomm, workspace}; use rand::RngCore; use std::ffi::OsString; @@ -195,6 +196,10 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option>(); + let remote_tool_names = remote_tool_names(&env.project.remote.tools); let share_dir_owned = share_dir.to_path_buf(); let mut sock_fds = [-1i32, -1]; @@ -328,8 +333,7 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option bunkerbox::remote::RemoteTargetId { } } -fn remote_tool_names(entries: &[String]) -> Vec { - let mut names = std::collections::BTreeSet::new(); - for entry in entries { - let command = entry.trim().strip_suffix(" *").unwrap_or(entry.trim()); - if let Some(tool) = command.split_whitespace().next().filter(|tool| !tool.is_empty() && !tool.contains('/')) { - names.insert(tool.to_string()); - } - } - names.into_iter().collect() +fn remote_tool_names(entries: &[RemoteToolSpec]) -> Vec { + entries.iter().map(|tool| tool.name.clone()).collect() } fn write_run_handoff(file: &mut File, path: &Path, session_id: vscomm::WorkspaceSessionId) -> Result<(), String> { diff --git a/src/remote.rs b/src/remote.rs index eef8337..1c5c869 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -1,7 +1,9 @@ #![allow(dead_code)] +use std::collections::{BTreeMap, BTreeSet}; use std::future::Future; use std::pin::Pin; +use std::time::Duration; use tokio::sync::mpsc; pub const MAX_REMOTE_STRING_BYTES: usize = 4 * 1024; @@ -11,6 +13,9 @@ pub const MAX_REMOTE_ARG_BYTES: usize = 4 * 1024; pub const MAX_REMOTE_ENV_COUNT: usize = 64; pub const MAX_REMOTE_ENV_KEY_BYTES: usize = 256; pub const MAX_REMOTE_ENV_VALUE_BYTES: usize = 4 * 1024; +pub const DEFAULT_REMOTE_BUILD_TIMEOUT: Duration = Duration::from_secs(30); +pub const DEFAULT_REMOTE_OUTPUT_BYTES: u64 = 64 * 1024 * 1024; +pub const DEFAULT_REMOTE_ENVIRONMENT: &[&str] = &["CC", "CXX", "AR", "RUSTFLAGS", "CFLAGS", "CXXFLAGS", "MAKEFLAGS"]; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RequestId(pub [u8; 16]); @@ -51,8 +56,12 @@ impl RemoteTool { pub fn new(value: impl Into) -> Result { let value = value.into(); validate_string("remote tool", &value, MAX_REMOTE_TOOL_BYTES)?; - if value.is_empty() { - return Err("remote tool is empty".to_string()); + if value.is_empty() + || value == "." + || value == ".." + || !value.bytes().all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b'+')) + { + return Err("remote tool must be a single executable identity".to_string()); } Ok(Self(value)) } @@ -76,11 +85,8 @@ impl RemoteBuild { argv.iter().try_for_each(|arg| validate_string("remote argument", arg, MAX_REMOTE_ARG_BYTES))?; validate_count(env.len(), MAX_REMOTE_ENV_COUNT, "remote environment")?; env.iter().try_for_each(|(key, value)| { - validate_string("remote environment key", key, MAX_REMOTE_ENV_KEY_BYTES)?; - if key.is_empty() || key.contains('=') { - return Err("remote environment key is invalid".to_string()); - } - validate_string("remote environment value", value, MAX_REMOTE_ENV_VALUE_BYTES) + validate_environment_key(key)?; + validate_environment_value(value) })?; Ok(Self { cwd, tool, argv, env }) } @@ -102,6 +108,90 @@ impl RemoteBuild { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteResourcePolicy { + pub build_timeout: Duration, + pub max_output_bytes: u64, +} + +impl Default for RemoteResourcePolicy { + fn default() -> Self { + Self { build_timeout: DEFAULT_REMOTE_BUILD_TIMEOUT, max_output_bytes: DEFAULT_REMOTE_OUTPUT_BYTES } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RemoteEnvironmentPolicy { + allowed: BTreeSet, +} + +impl Default for RemoteEnvironmentPolicy { + fn default() -> Self { + Self { allowed: DEFAULT_REMOTE_ENVIRONMENT.iter().map(|name| (*name).to_string()).collect() } + } +} + +impl RemoteEnvironmentPolicy { + pub fn from_names(names: impl IntoIterator) -> Result { + let mut policy = Self::default(); + let mut configured = BTreeSet::new(); + for name in names { + validate_environment_key(&name)?; + if forbidden_environment_name(&name) { + return Err(format!("remote environment variable is forbidden: {name}")); + } + if !configured.insert(name.clone()) { + return Err(format!("duplicate remote environment variable: {name}")); + } + policy.allowed.insert(name); + } + Ok(policy) + } + + pub fn allows(&self, name: &str) -> bool { + self.allowed.contains(name) + } + + pub fn allowed_names(&self) -> impl Iterator { + self.allowed.iter().map(String::as_str) + } + + fn filter(&self, environment: &[(String, String)]) -> Result, RemoteAuthorizationError> { + let mut seen = BTreeSet::new(); + let mut filtered = Vec::with_capacity(environment.len()); + for (key, value) in environment { + validate_environment_key(key).map_err(RemoteAuthorizationError::InvalidEnvironment)?; + validate_environment_value(value).map_err(RemoteAuthorizationError::InvalidEnvironment)?; + if forbidden_environment_name(key) { + return Err(RemoteAuthorizationError::ForbiddenEnvironment(key.clone())); + } + if !self.allowed.contains(key) { + return Err(RemoteAuthorizationError::EnvironmentNotAllowed(key.clone())); + } + if !seen.insert(key.clone()) { + return Err(RemoteAuthorizationError::DuplicateEnvironment(key.clone())); + } + filtered.push((key.clone(), value.clone())); + } + Ok(filtered) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteToolPolicy { + allow_arbitrary_argv: bool, +} + +impl RemoteToolPolicy { + pub fn new(allow_arbitrary_argv: bool) -> Self { + Self { allow_arbitrary_argv } + } + + pub fn allows_arbitrary_argv(self) -> bool { + self.allow_arbitrary_argv + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum RemoteOperation { Sync, @@ -168,6 +258,11 @@ pub enum RemoteAuthorizationError { SessionMismatch, TargetNotAllowed, ToolNotAllowed(String), + ToolArgumentsNotAllowed(String), + InvalidEnvironment(String), + ForbiddenEnvironment(String), + EnvironmentNotAllowed(String), + DuplicateEnvironment(String), } impl RemoteAuthorizationError { @@ -180,12 +275,32 @@ impl RemoteAuthorizationError { pub struct RemoteAuthorizationPolicy { allowed_target: RemoteTargetId, allowed_session: WorkspaceSessionId, - allowed_tools: Vec, + allowed_tools: BTreeMap, + environment: RemoteEnvironmentPolicy, } impl RemoteAuthorizationPolicy { pub fn new(allowed_target: RemoteTargetId, allowed_session: WorkspaceSessionId, allowed_tools: Vec) -> Self { - Self { allowed_target, allowed_session, allowed_tools } + let allowed_tools = allowed_tools.into_iter().map(|tool| (tool, RemoteToolPolicy::new(true))).collect(); + Self { allowed_target, allowed_session, allowed_tools, environment: RemoteEnvironmentPolicy::default() } + } + + pub fn from_policies( + allowed_target: RemoteTargetId, allowed_session: WorkspaceSessionId, tools: impl IntoIterator, + environment: RemoteEnvironmentPolicy, + ) -> Result { + let mut allowed_tools = BTreeMap::new(); + for (tool, policy) in tools { + RemoteTool::new(tool.clone())?; + if allowed_tools.insert(tool.clone(), policy).is_some() { + return Err(format!("duplicate remote tool policy: {tool}")); + } + } + Ok(Self { allowed_target, allowed_session, allowed_tools, environment }) + } + + pub fn allowed_tools(&self) -> impl Iterator { + self.allowed_tools.keys().map(String::as_str) } pub fn authorize(&self, context: &RemoteExecutionContext, request: RemoteRequest) -> Result { @@ -196,11 +311,21 @@ impl RemoteAuthorizationPolicy { return Err(RemoteAuthorizationError::TargetNotAllowed); } - if let RemoteOperation::Build(build) = request.operation() { - if !self.allowed_tools.iter().any(|tool| tool == build.tool().as_str()) { - return Err(RemoteAuthorizationError::ToolNotAllowed(build.tool().as_str().to_string())); + let request = match request.operation() { + RemoteOperation::Sync => request, + RemoteOperation::Build(build) => { + let Some(tool_policy) = self.allowed_tools.get(build.tool().as_str()).copied() else { + return Err(RemoteAuthorizationError::ToolNotAllowed(build.tool().as_str().to_string())); + }; + if !tool_policy.allows_arbitrary_argv() && !build.argv().is_empty() { + return Err(RemoteAuthorizationError::ToolArgumentsNotAllowed(build.tool().as_str().to_string())); + } + let environment = self.environment.filter(build.env())?; + let filtered = RemoteBuild::new(build.cwd.clone(), build.tool.clone(), build.argv.clone(), environment) + .map_err(RemoteAuthorizationError::InvalidEnvironment)?; + RemoteRequest::build(request.request_id, request.workspace_session_id, filtered) } - } + }; Ok(AuthorizedRemoteRequest { request, target: context.target }) } @@ -221,6 +346,7 @@ pub enum RemoteBackendError { Failed(String), Spawn(String), Timeout, + OutputLimit { limit: u64 }, Cancelled, } @@ -229,6 +355,7 @@ impl RemoteBackendError { match self { Self::Failed(message) | Self::Spawn(message) => RemoteBackendEvent::Error { message: message.clone() }, Self::Timeout => RemoteBackendEvent::Error { message: "remote backend timed out".to_string() }, + Self::OutputLimit { limit } => RemoteBackendEvent::Error { message: format!("remote output exceeded limit of {limit} bytes") }, Self::Cancelled => RemoteBackendEvent::Cancelled, } } @@ -252,6 +379,85 @@ fn validate_string(field: &str, value: &str, max: usize) -> Result<(), String> { Ok(()) } +fn validate_environment_key(key: &str) -> Result<(), String> { + validate_string("remote environment key", key, MAX_REMOTE_ENV_KEY_BYTES)?; + let mut characters = key.bytes(); + let Some(first) = characters.next() else { + return Err("remote environment key is empty".to_string()); + }; + if !(first == b'_' || first.is_ascii_alphabetic()) || !characters.all(|byte| byte == b'_' || byte.is_ascii_alphanumeric()) { + return Err("remote environment key is not a valid variable name".to_string()); + } + Ok(()) +} + +fn validate_environment_value(value: &str) -> Result<(), String> { + validate_string("remote environment value", value, MAX_REMOTE_ENV_VALUE_BYTES)?; + if value.bytes().any(|byte| byte < 0x20 || byte == 0x7f) { + return Err("remote environment value contains control data".to_string()); + } + Ok(()) +} + +fn forbidden_environment_name(name: &str) -> bool { + let upper = name.to_ascii_uppercase(); + matches!( + upper.as_str(), + "SSH_AUTH_SOCK" + | "SSH_AGENT_PID" + | "GITHUB_TOKEN" + | "GITLAB_TOKEN" + | "NPM_TOKEN" + | "KUBECONFIG" + | "HOME" + | "CARGO_HOME" + | "RUSTUP_HOME" + | "XDG_CONFIG_HOME" + | "XDG_DATA_HOME" + | "PATH" + | "PWD" + | "OLDPWD" + | "TMP" + | "TMPDIR" + | "TEMP" + | "USER" + | "LOGNAME" + | "SHELL" + | "BASH_ENV" + | "ENV" + | "CDPATH" + | "LD_PRELOAD" + | "LD_LIBRARY_PATH" + | "PYTHONPATH" + | "PERL5LIB" + | "RUBYLIB" + | "NODE_PATH" + | "GOPATH" + | "GOMODCACHE" + | "TOKEN" + | "PASSWORD" + | "PASS" + | "SECRET" + | "KEY" + | "GIT_SSH_COMMAND" + ) || upper.starts_with("AWS_") + || upper.starts_with("GCP_") + || upper.starts_with("GOOGLE_") + || upper.starts_with("AZURE_") + || upper.starts_with("DOCKER_") + || upper.starts_with("CARGO_REGISTRIES_") + || upper.starts_with("XDG_") + || upper.starts_with("BUNKERBOX_") + || upper.ends_with("_PROXY") + || upper.ends_with("_TOKEN") + || upper.ends_with("_PASSWORD") + || upper.ends_with("_PASS") + || upper.ends_with("_SECRET") + || upper.ends_with("_KEY") + || upper.contains("CREDENTIAL") + || upper.contains("PRIVATE_KEY") +} + fn validate_count(count: usize, max: usize, field: &str) -> Result<(), String> { if count > max { return Err(format!("{field} exceeds maximum count {max}")); diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index 50c1b07..0e1bced 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -131,8 +131,12 @@ impl RemoteTool { pub fn new(value: impl Into) -> Result { let value = value.into(); validate_remote_string("remote tool", &value, MAX_REMOTE_TOOL_BYTES)?; - if value.is_empty() { - return Err("remote tool is empty".to_string()); + if value.is_empty() + || value == "." + || value == ".." + || !value.bytes().all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b'+')) + { + return Err("remote tool must be a single executable identity".to_string()); } Ok(Self(value)) } @@ -279,7 +283,8 @@ fn validate_remote_build_fields(cwd: &WorkspaceRelativePath, tool: &RemoteTool, env.iter().try_for_each(|(key, value)| { validate_remote_string("remote environment key", key, MAX_REMOTE_ENV_KEY_BYTES)?; validate_env_key("remote environment key", key)?; - validate_remote_string("remote environment value", value, MAX_REMOTE_ENV_VALUE_BYTES) + validate_remote_string("remote environment value", value, MAX_REMOTE_ENV_VALUE_BYTES)?; + validate_env_value("remote environment value", value) }) } @@ -298,6 +303,13 @@ fn validate_remote_count(count: usize, max: usize, field: &str) -> Result<(), St Ok(()) } +fn validate_env_value(field: &str, value: &str) -> Result<(), String> { + if value.bytes().any(|byte| byte < 0x20 || byte == 0x7f) { + return Err(format!("{field} contains control data")); + } + Ok(()) +} + struct WireWriter { bytes: Vec, } From 07f5477122d41ffdefcd3442c7cb7328a0f49154 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 10:01:28 +0200 Subject: [PATCH 18/52] Add remote build policy --- src/loopback.rs | 101 +++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 91 insertions(+), 10 deletions(-) diff --git a/src/loopback.rs b/src/loopback.rs index 39362ae..d13f62f 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -18,6 +18,8 @@ use tokio::time::sleep; pub const LOOPBACK_BUILD_TIMEOUT: Duration = crate::remote::DEFAULT_REMOTE_BUILD_TIMEOUT; const LOOPBACK_PATH: &str = "/usr/local/bin:/usr/bin:/bin"; +const POST_CHILD_EXIT_DRAIN_TIMEOUT: Duration = Duration::from_millis(100); +const POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT: Duration = Duration::from_millis(100); static NEXT_JOB_ID: AtomicU64 = AtomicU64::new(1); pub struct RunRemoteSession { @@ -235,37 +237,92 @@ async fn execute_build( let stdout = child.stdout.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stdout".to_string()))?; let stderr = child.stderr.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stderr".to_string()))?; let output_bytes = Arc::new(AtomicU64::new(0)); - let mut stdout_task = - Box::pin(tokio::spawn(pump(stdout, RemoteStream::Stdout, events.clone(), output_bytes.clone(), resources.max_output_bytes))); - let mut stderr_task = Box::pin(tokio::spawn(pump(stderr, RemoteStream::Stderr, events.clone(), output_bytes, resources.max_output_bytes))); + let mut stdout_task = tokio::spawn(pump(stdout, RemoteStream::Stdout, events.clone(), output_bytes.clone(), resources.max_output_bytes)); + let mut stderr_task = tokio::spawn(pump(stderr, RemoteStream::Stderr, events.clone(), output_bytes, resources.max_output_bytes)); let mut child_wait = Box::pin(child.wait()); let mut timeout_sleep = Box::pin(sleep(resources.build_timeout)); + let mut post_exit_drain_sleep = Box::pin(sleep(POST_CHILD_EXIT_DRAIN_TIMEOUT)); + let mut final_drain_sleep = Box::pin(sleep(POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT)); let mut child_status = None; let mut stdout_done = false; let mut stderr_done = false; let mut failure = None; + let mut post_exit_drain_active = false; + let mut final_drain_active = false; + let mut group_killed = false; while child_status.is_none() || !stdout_done || !stderr_done { tokio::select! { status = &mut child_wait, if child_status.is_none() => { - child_status = Some(status.map_err(|error| RemoteBackendError::Failed(format!("wait for loopback tool: {error}")))); + let status = status.map_err(|error| RemoteBackendError::Failed(format!("wait for loopback tool: {error}"))); + child_status = Some(status); + if failure.is_some() || child_status.as_ref().is_some_and(Result::is_err) { + kill_process_group(process_group.as_ref()); + group_killed = true; + final_drain_sleep.as_mut().reset(tokio::time::Instant::now() + POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT); + final_drain_active = true; + } else { + // The child is the process-group leader. Keep draining useful + // pipe data briefly, while terminating descendants that + // inherited the build descriptors. + terminate_process_group(process_group.as_ref()); + post_exit_drain_sleep.as_mut().reset(tokio::time::Instant::now() + POST_CHILD_EXIT_DRAIN_TIMEOUT); + post_exit_drain_active = true; + } } result = &mut stdout_task, if !stdout_done => { stdout_done = true; - if let Err(error) = join_pump(result).await { failure.get_or_insert(error); kill_process_group(process_group.as_ref()); } + if let Err(error) = join_pump(result).await { + failure.get_or_insert(error); + kill_process_group(process_group.as_ref()); + group_killed = true; + post_exit_drain_active = false; + if child_status.is_some() { + final_drain_sleep.as_mut().reset(tokio::time::Instant::now() + POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT); + final_drain_active = true; + } + } } result = &mut stderr_task, if !stderr_done => { stderr_done = true; - if let Err(error) = join_pump(result).await { failure.get_or_insert(error); kill_process_group(process_group.as_ref()); } + if let Err(error) = join_pump(result).await { + failure.get_or_insert(error); + kill_process_group(process_group.as_ref()); + group_killed = true; + post_exit_drain_active = false; + if child_status.is_some() { + final_drain_sleep.as_mut().reset(tokio::time::Instant::now() + POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT); + final_drain_active = true; + } + } } _ = &mut timeout_sleep, if child_status.is_none() => { failure.get_or_insert(RemoteBackendError::Timeout); kill_process_group(process_group.as_ref()); + group_killed = true; + } + _ = &mut post_exit_drain_sleep, if post_exit_drain_active => { + post_exit_drain_active = false; + kill_process_group(process_group.as_ref()); + group_killed = true; + final_drain_sleep.as_mut().reset(tokio::time::Instant::now() + POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT); + final_drain_active = true; + } + _ = &mut final_drain_sleep, if final_drain_active => { + final_drain_active = false; + if !stdout_done { + if let Some(error) = abort_pump(&mut stdout_task).await { failure.get_or_insert(error); } + stdout_done = true; + } + if !stderr_done { + if let Some(error) = abort_pump(&mut stderr_task).await { failure.get_or_insert(error); } + stderr_done = true; + } } } } - if failure.is_none() { + if failure.is_none() || !group_killed { kill_process_group(process_group.as_ref()); } if let Some(mut process_group) = process_group { @@ -313,6 +370,18 @@ async fn join_pump(result: Result, tokio::task::J result.map_err(|error| RemoteBackendError::Failed(format!("loopback output task failed: {error}")))? } +async fn abort_pump(task: &mut tokio::task::JoinHandle>) -> Option { + if !task.is_finished() { + task.abort(); + } + match task.await { + Ok(Ok(())) => None, + Ok(Err(error)) => Some(error), + Err(error) if error.is_cancelled() => None, + Err(error) => Some(RemoteBackendError::Failed(format!("loopback output task failed: {error}"))), + } +} + async fn send_event(events: &mpsc::Sender, event: RemoteBackendEvent) -> Result<(), RemoteBackendError> { events.send(event).await.map_err(|_| RemoteBackendError::Cancelled) } @@ -341,10 +410,22 @@ impl Drop for ProcessGroupGuard { } fn kill_process_group(group: Option<&ProcessGroupGuard>) { + signal_process_group(group, libc::SIGTERM); + signal_process_group(group, libc::SIGKILL); +} + +fn terminate_process_group(group: Option<&ProcessGroupGuard>) { + signal_process_group(group, libc::SIGTERM); +} + +fn signal_process_group(group: Option<&ProcessGroupGuard>, signal: libc::c_int) { if let Some(group) = group { - unsafe { - libc::kill(-group.pgid, libc::SIGTERM); - libc::kill(-group.pgid, libc::SIGKILL); + if group.pgid > 0 { + // A negative PID targets the Unix process group, so this remains + // effective after the original group leader has exited. + unsafe { + libc::kill(-group.pgid, signal); + } } } } From 04b1c3c577f0414fa80e7659a667d10b69457fa7 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 10:01:44 +0200 Subject: [PATCH 19/52] Add remote build policy unit tests --- src/loopback_ut.rs | 45 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 45 insertions(+) diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 21cb637..456bd84 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -98,6 +98,51 @@ async fn build_preserves_argument_bytes_and_reports_nonzero_stderr() { assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code } if *exit_code != 0))); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn child_exit_does_not_wait_for_inherited_output_pipes() { + let (_temp, session, target, session_id) = fixture(); + fs::write(session.workspace_root().join("src/Makefile"), ".PHONY: leak\nleak:\n\t@sleep 30 & echo $$!\n\t@printf 'direct-output\\n'\n").unwrap(); + session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["make".to_string()]); + if tools.is_empty() { + return; + } + + let backend = LoopbackBackend::new(session.clone(), tools); + let (events, receiver) = mpsc::channel(16); + let request = authorized_build(target, session_id, "make", vec!["leak".into()], Vec::new()); + let started = tokio::time::Instant::now(); + let result = tokio::time::timeout(Duration::from_secs(1), backend.execute(request, events)).await.unwrap(); + assert_eq!(result, Ok(())); + assert!(started.elapsed() < Duration::from_millis(500)); + + let events = collect_events(receiver).await; + let stdout = events + .iter() + .filter_map(|event| match event { + RemoteBackendEvent::Stdout(bytes) => Some(bytes.as_slice()), + _ => None, + }) + .flatten() + .copied() + .collect::>(); + assert!(stdout.windows(b"direct-output\n".len()).any(|window| window == b"direct-output\n")); + let pid = std::str::from_utf8(&stdout).unwrap().split_whitespace().find_map(|value| value.parse::().ok()).unwrap(); + + let terminal_events = events + .iter() + .filter(|event| matches!(event, RemoteBackendEvent::Error { .. } | RemoteBackendEvent::Cancelled | RemoteBackendEvent::Completed { .. })) + .collect::>(); + assert_eq!(terminal_events, vec![&RemoteBackendEvent::Completed { exit_code: 0 }]); + + let deadline = tokio::time::Instant::now() + Duration::from_secs(1); + while super::process_is_alive(pid) && tokio::time::Instant::now() < deadline { + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert!(!super::process_is_alive(pid)); + assert!(fs::read_dir(&session.jobs_root).unwrap().next().is_none()); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn child_receives_guest_environment_but_trusted_target_wins() { let (_temp, session, target, session_id) = fixture(); From 932dbdaa785289b2c38c5cd5d49e22c4bb8d6418 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 10:40:52 +0200 Subject: [PATCH 20/52] Implement transparency wrapper on one build command --- src/bin/bunkerbox-image.rs | 29 +++++++- src/bin/bunkerbox-remote.rs | 133 ++++++++++++++++++++++++++++++------ src/daemon.rs | 3 +- src/kata.rs | 8 +++ src/loopback.rs | 123 +++++++++++++++++++++++++++------ src/main.rs | 13 +++- src/remote.rs | 57 ++++++++++++++-- src/remote_client.rs | 87 +++++++++++++++++++++-- src/snapshot.rs | 2 +- src/vscomm/mod.rs | 49 +++++++++++-- 10 files changed, 438 insertions(+), 66 deletions(-) diff --git a/src/bin/bunkerbox-image.rs b/src/bin/bunkerbox-image.rs index f1a7289..0094e95 100644 --- a/src/bin/bunkerbox-image.rs +++ b/src/bin/bunkerbox-image.rs @@ -226,10 +226,13 @@ run_app() {{ }} VSCOMM_BIN="/usr/local/bunkerbox/bin" +if [ -x "$VSCOMM_BIN/bunkerbox-remote" ]; then + "$VSCOMM_BIN/bunkerbox-remote" install +fi if [ -x "$VSCOMM_BIN/bunkerbox-vscomm" ]; then "$VSCOMM_BIN/bunkerbox-vscomm" install - export PATH="$VSCOMM_BIN:$PATH" fi +export PATH="$VSCOMM_BIN:$PATH" if ! command -v bunkerbox-status >/dev/null 2>&1; then bunkerbox-status() {{ :; }} @@ -286,6 +289,11 @@ fn write_build_context(config: &ImageConfig, build_dir: &Path) -> Result<(), Str fs::copy(&vscomm_path, &dest).map_err(|err| format!("failed to copy vscomm binary {}: {err}", dest.display()))?; fs::set_permissions(&dest, fs::Permissions::from_mode(0o755)).map_err(|err| format!("failed to chmod {}: {err}", dest.display()))?; + let remote_path = find_remote_binary()?; + let dest = build_dir.join("bunkerbox-remote"); + fs::copy(&remote_path, &dest).map_err(|err| format!("failed to copy remote binary {}: {err}", dest.display()))?; + fs::set_permissions(&dest, fs::Permissions::from_mode(0o755)).map_err(|err| format!("failed to chmod {}: {err}", dest.display()))?; + let status_path = find_status_binary()?; let dest = build_dir.join("bunkerbox-status"); fs::copy(&status_path, &dest).map_err(|err| format!("failed to copy status binary {}: {err}", dest.display()))?; @@ -295,7 +303,11 @@ fn write_build_context(config: &ImageConfig, build_dir: &Path) -> Result<(), Str if file.path.is_absolute() || file.path.components().any(|part| matches!(part, std::path::Component::ParentDir)) { return Err(format!("unsafe build file path: {}", file.path.display())); } - if file.path == Path::new("bunker-entrypoint") || file.path == Path::new("bunkerbox-vscomm") || file.path == Path::new("bunkerbox-status") { + if file.path == Path::new("bunker-entrypoint") + || file.path == Path::new("bunkerbox-vscomm") + || file.path == Path::new("bunkerbox-status") + || file.path == Path::new("bunkerbox-remote") + { return Err(format!("image config files must not override reserved file: {}", file.path.display())); } @@ -336,6 +348,15 @@ fn find_status_binary() -> Result { } } +fn find_remote_binary() -> Result { + let path = dist_dir()?.join("bunkerbox-remote"); + if path.is_file() { + Ok(path) + } else { + Err("bunkerbox-remote not found in target/dist/. Run: make dev".into()) + } +} + fn podman_build(config: &ImageConfig, build_dir: &Path) -> Result<(), String> { let mut args = vec!["build".to_string(), "--no-cache".to_string()]; @@ -451,3 +472,7 @@ where } } } + +#[cfg(test)] +#[path = "../bunkerbox-image_ut.rs"] +mod tests; diff --git a/src/bin/bunkerbox-remote.rs b/src/bin/bunkerbox-remote.rs index d8439e0..d3cd19e 100644 --- a/src/bin/bunkerbox-remote.rs +++ b/src/bin/bunkerbox-remote.rs @@ -1,9 +1,16 @@ -use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request}; -use bunkerbox::vscomm::{RemoteRequest, RequestId, WorkspaceSessionId, TOOLCHAIN_PORT}; -use rand::RngCore; +use bunkerbox::remote::RemoteSnapshotId; +use bunkerbox::remote_client::{ + execute_remote_request_to, logical_workspace_cwd, new_request_id, remote_build_request, remote_environment_names, remote_session_from_env, + remote_sync_request, remote_tool_enabled, selected_remote_environment, RemoteCompletion, +}; +#[cfg(test)] +use bunkerbox::vscomm::RequestId; +use bunkerbox::vscomm::{RemoteRequest, WorkspaceSessionId, TOOLCHAIN_PORT, VSCOMM_BIN_DIR}; use std::env; +use std::fs; use std::io::{self, Read, Write}; use std::mem; +use std::os::unix::fs::symlink; use std::path::Path; const HOST_CID: u32 = 2; @@ -25,21 +32,80 @@ fn main() { } fn run() -> Result { + let invoked_as = + env::args_os().next().and_then(|value| Path::new(&value).file_name().and_then(|name| name.to_str()).map(str::to_owned)).unwrap_or_default(); let args = env::args().skip(1).collect::>(); - let command = parse_command(&args)?; + + if invoked_as == "make" { + return run_transparent_make(&args); + } + if invoked_as != "bunkerbox-remote" { + return Err("bunkerbox-remote must be invoked directly or through the managed make symlink".to_string()); + } + if args.len() == 1 && args[0] == "install" { + install_remote_links()?; + return Ok(0); + } + + run_explicit(&args) +} + +fn run_explicit(args: &[String]) -> Result { + let command = parse_command(args)?; match &command { RemoteCommand::Sync => eprintln!("bunkerbox-remote: syncing"), RemoteCommand::Build { tool, .. } => eprintln!("bunkerbox-remote: building {tool}"), } + + let cwd = logical_workspace_cwd(&env::current_dir().map_err(|error| format!("current directory: {error}"))?)?; + let session = remote_session_from_env()?; + match command { + RemoteCommand::Sync => { + sync_snapshot(session)?; + Ok(0) + } + RemoteCommand::Build { tool, args } => run_build_with_sync(cwd, tool, args, session), + } +} + +fn run_transparent_make(args: &[String]) -> Result { let cwd = logical_workspace_cwd(&env::current_dir().map_err(|error| format!("current directory: {error}"))?)?; - let session = env::var("BUNKERBOX_REMOTE_SESSION") - .map_err(|_| "BUNKERBOX_REMOTE_SESSION is missing".to_string()) - .and_then(|value| WorkspaceSessionId::from_hex(&value))?; - let request = build_request(command, cwd, new_request_id(), session)?; + let session = remote_session_from_env()?; + run_build_with_sync(cwd, "make".to_string(), args.to_vec(), session) +} + +fn run_build_with_sync(cwd: String, tool: String, args: Vec, session: WorkspaceSessionId) -> Result { + let environment = selected_remote_environment(remote_environment_names()); + run_build_with_sync_using(cwd, tool, args, environment, session, execute_request_over_vsock) +} + +fn run_build_with_sync_using( + cwd: String, tool: String, args: Vec, environment: Vec<(String, String)>, session: WorkspaceSessionId, mut execute: F, +) -> Result +where + F: FnMut(RemoteRequest) -> Result, +{ + let snapshot_id = match execute(remote_sync_request(new_request_id(), session))? { + RemoteCompletion::Synced(snapshot_id) => snapshot_id, + RemoteCompletion::Completed(_) => return Err("remote sync returned a build completion".to_string()), + }; + let request = remote_build_request(new_request_id(), session, cwd, tool, args, environment, snapshot_id)?; + match execute(request)? { + RemoteCompletion::Completed(code) => Ok(code), + RemoteCompletion::Synced(_) => Err("remote build returned a sync completion".to_string()), + } +} + +fn execute_request_over_vsock(request: RemoteRequest) -> Result { let mut stream = connect_toolchain()?; - let mut stdout = io::stdout(); - let mut stderr = io::stderr(); - execute_remote_request_to(&mut stream, request, &mut stdout, &mut stderr) + execute_remote_request_to(&mut stream, request, &mut io::stdout(), &mut io::stderr()) +} + +fn sync_snapshot(session: WorkspaceSessionId) -> Result { + match execute_request_over_vsock(remote_sync_request(new_request_id(), session))? { + RemoteCompletion::Synced(snapshot_id) => Ok(snapshot_id), + RemoteCompletion::Completed(_) => Err("remote sync returned a build completion".to_string()), + } } fn parse_command(args: &[String]) -> Result { @@ -52,24 +118,47 @@ fn parse_command(args: &[String]) -> Result { } } -fn build_request(command: RemoteCommand, cwd: String, request_id: RequestId, session_id: WorkspaceSessionId) -> Result { +#[cfg(test)] +fn build_request( + command: RemoteCommand, cwd: String, request_id: RequestId, session_id: WorkspaceSessionId, snapshot_id: RemoteSnapshotId, +) -> Result { match command { RemoteCommand::Sync => Ok(remote_sync_request(request_id, session_id)), - RemoteCommand::Build { tool, args } => remote_build_request(request_id, session_id, cwd, tool, args, Vec::new()), + RemoteCommand::Build { tool, args } => remote_build_request(request_id, session_id, cwd, tool, args, Vec::new(), snapshot_id), } } -fn logical_workspace_cwd(path: &Path) -> Result { - let relative = path.strip_prefix("/workspace").map_err(|_| "current directory must be under /workspace".to_string())?; - let value = relative.to_str().ok_or_else(|| "current directory is not valid UTF-8".to_string())?.to_string(); - bunkerbox::remote::WorkspaceRelativePath::new(&value)?; - Ok(value) +fn install_remote_links() -> Result<(), String> { + fs::create_dir_all(VSCOMM_BIN_DIR).map_err(|error| format!("mkdir {VSCOMM_BIN_DIR}: {error}"))?; + let executable = env::current_exe().map_err(|error| format!("failed to locate remote binary: {error}"))?; + install_remote_make_link(Path::new(VSCOMM_BIN_DIR), &executable, remote_tool_enabled("make")) +} + +fn install_remote_make_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { + let target = bin_dir.join("make"); + let managed = is_managed_link(&target, executable); + + if enabled { + if target.exists() || fs::symlink_metadata(&target).is_ok() { + if !managed { + return Err(format!("cannot install remote make wrapper over existing {}", target.display())); + } + fs::remove_file(&target).map_err(|error| format!("remove existing remote make wrapper: {error}"))?; + } + symlink(executable, &target).map_err(|error| format!("symlink remote make wrapper: {error}"))?; + } else if managed { + fs::remove_file(&target).map_err(|error| format!("remove disabled remote make wrapper: {error}"))?; + } + Ok(()) } -fn new_request_id() -> RequestId { - let mut bytes = [0; 16]; - rand::thread_rng().fill_bytes(&mut bytes); - RequestId(bytes) +fn is_managed_link(target: &Path, executable: &Path) -> bool { + let Ok(metadata) = fs::symlink_metadata(target) else { return false }; + if !metadata.file_type().is_symlink() { + return false; + } + let Ok(link) = fs::read_link(target) else { return false }; + link == executable } fn connect_toolchain() -> Result { diff --git a/src/daemon.rs b/src/daemon.rs index 34a8449..322df6f 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -151,7 +151,8 @@ impl VsockDaemon { } (None, None) => RemoteAuthorizationPolicy::new(session.target(), session.session_id(), allowed_tools), _ => return Err("remote tool and environment policies must be configured together".to_string()), - }; + } + .with_snapshot_authority(session.clone()); let remote_context = RemoteExecutionContext { target: session.target(), workspace_session_id: session.session_id() }; let backend = LoopbackBackend::new(session, tools) .with_target_environment(target_environment) diff --git a/src/kata.rs b/src/kata.rs index 2ac635d..560675b 100644 --- a/src/kata.rs +++ b/src/kata.rs @@ -20,6 +20,8 @@ use std::thread; pub struct WorkspaceBinding<'a> { pub path: &'a Path, pub remote_session: WorkspaceSessionId, + pub remote_tools: &'a [String], + pub remote_environment: &'a [String], } const BRIDGE_SUBNET: &str = "10.247.0.0/24"; @@ -197,6 +199,12 @@ pub fn run( if vsock_enabled { container_env.push(format!("BUNKERBOX_TOOLCHAIN_PORT={TOOLCHAIN_PORT}")); container_env.push(format!("BUNKERBOX_REMOTE_SESSION={}", workspace.remote_session.to_hex())); + if !workspace.remote_tools.is_empty() { + container_env.push(format!("BUNKERBOX_REMOTE_TOOLS={}", workspace.remote_tools.join(","))); + } + if !workspace.remote_environment.is_empty() { + container_env.push(format!("BUNKERBOX_REMOTE_ENV_NAMES={}", workspace.remote_environment.join(","))); + } } if let Some(ref cmds) = config.command { diff --git a/src/loopback.rs b/src/loopback.rs index d13f62f..9c3fb62 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -1,9 +1,10 @@ use crate::remote::{ AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteResourcePolicy, - RemoteTargetId, WorkspaceSessionId, + RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, WorkspaceSessionId, }; use crate::snapshot::{SnapshotBuilder, SnapshotHandle, SnapshotStore}; -use std::collections::BTreeMap; +use rand::RngCore; +use std::collections::{BTreeMap, HashMap}; use std::fs; use std::io; use std::os::unix::fs::{MetadataExt, PermissionsExt}; @@ -28,7 +29,7 @@ pub struct RunRemoteSession { workspace_root: PathBuf, snapshot_store: SnapshotStore, snapshot_builder: SnapshotBuilder, - current_snapshot: Mutex>, + snapshot_capabilities: Mutex, snapshot_operation: Mutex<()>, jobs_root: PathBuf, } @@ -56,7 +57,7 @@ impl RunRemoteSession { workspace_root, snapshot_store, snapshot_builder, - current_snapshot: Mutex::new(None), + snapshot_capabilities: Mutex::new(SnapshotCapabilityRegistry::default()), snapshot_operation: Mutex::new(()), jobs_root, }) @@ -78,26 +79,59 @@ impl RunRemoteSession { self.snapshot_store.clone() } - pub fn current_snapshot(&self) -> Result, String> { - self.current_snapshot.lock().map(|current| current.clone()).map_err(|_| "remote session state lock poisoned".to_string()) - } - - pub fn sync_snapshot(&self) -> Result { + pub fn sync_snapshot(&self) -> Result { let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; let snapshot = self.snapshot_builder.build_root(&self.workspace_root, self.session_id)?; - let new_handle = snapshot.handle().clone(); - let old_handle = self.current_snapshot()?.clone(); - if let Some(old_handle) = old_handle.filter(|old| old != &new_handle) { - self.snapshot_store.remove(&old_handle)?; + self.register_snapshot(snapshot.handle().clone()) + } + + fn register_snapshot(&self, handle: SnapshotHandle) -> Result { + let mut registry = self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; + let snapshot_id = loop { + let mut bytes = [0u8; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + let candidate = RemoteSnapshotId::from_bytes(bytes); + if !candidate.is_zero() && !registry.capabilities.contains_key(&candidate) { + break candidate; + } + }; + registry.capabilities.insert(snapshot_id, handle.clone()); + *registry.references.entry(handle).or_insert(0) += 1; + Ok(snapshot_id) + } + + fn claim_snapshot(self: &Arc, snapshot_id: RemoteSnapshotId) -> Result { + let mut registry = self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; + let handle = registry.capabilities.remove(&snapshot_id).ok_or_else(|| "remote snapshot capability is unavailable".to_string())?; + Ok(SnapshotClaim { session: self.clone(), handle }) + } + + fn release_snapshot(&self, handle: &SnapshotHandle) { + let should_remove = self + .snapshot_capabilities + .lock() + .ok() + .map(|mut registry| { + let Some(references) = registry.references.get_mut(handle) else { + return false; + }; + *references = references.saturating_sub(1); + if *references == 0 { + registry.references.remove(handle); + true + } else { + false + } + }) + .unwrap_or(false); + if should_remove { + let _ = self.snapshot_store.remove(handle); } - self.current_snapshot.lock().map_err(|_| "remote session state lock poisoned".to_string())?.replace(new_handle.clone()); - Ok(new_handle) } - fn materialize_current_snapshot(&self, destination: &Path) -> Result<(), String> { + fn materialize_snapshot(&self, handle: &SnapshotHandle, destination: &Path) -> Result<(), String> { let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; - let handle = self.current_snapshot()?.ok_or_else(|| "remote build requires a successful sync".to_string())?; - self.snapshot_store.materialize(&handle, destination).map(|_| ()) + self.snapshot_store.materialize(handle, destination).map(|_| ()) } fn new_job_path(&self) -> Result { @@ -109,6 +143,46 @@ impl RunRemoteSession { } } +#[derive(Default)] +struct SnapshotCapabilityRegistry { + capabilities: BTreeMap, + references: HashMap, +} + +struct SnapshotClaim { + session: Arc, + handle: SnapshotHandle, +} + +impl SnapshotClaim { + fn handle(&self) -> &SnapshotHandle { + &self.handle + } + + fn clone_for_worker(&self) -> Result { + let mut registry = self.session.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; + let references = + registry.references.get_mut(&self.handle).ok_or_else(|| "remote snapshot capability reference is unavailable".to_string())?; + *references = references.checked_add(1).ok_or_else(|| "remote snapshot capability reference count overflow".to_string())?; + Ok(Self { session: self.session.clone(), handle: self.handle.clone() }) + } +} + +impl Drop for SnapshotClaim { + fn drop(&mut self) { + self.session.release_snapshot(&self.handle); + } +} + +impl RemoteSnapshotAuthority for RunRemoteSession { + fn snapshot_available(&self, session: WorkspaceSessionId, snapshot_id: RemoteSnapshotId) -> bool { + if session != self.session_id { + return false; + } + self.snapshot_capabilities.lock().map(|registry| registry.capabilities.contains_key(&snapshot_id)).unwrap_or(false) + } +} + impl Drop for RunRemoteSession { fn drop(&mut self) { let _ = fs::remove_dir_all(&self.jobs_root); @@ -181,23 +255,30 @@ async fn execute_sync(session: Arc, events: mpsc::Sender, tools: Arc>, target_environment: Arc>, resources: RemoteResourcePolicy, build: &crate::remote::RemoteBuild, events: mpsc::Sender, ) -> Result<(), RemoteBackendError> { + let snapshot = session.claim_snapshot(build.snapshot_id()).map_err(RemoteBackendError::Failed)?; let executable = tools .get(build.tool().as_str()) .ok_or_else(|| RemoteBackendError::Spawn(format!("loopback tool is not configured: {}", build.tool().as_str())))?; let job_path = session.new_job_path().map_err(RemoteBackendError::Failed)?; let _job = JobGuard { path: job_path.clone() }; + let materialization_claim = snapshot.clone_for_worker().map_err(RemoteBackendError::Failed)?; let destination = job_path.clone(); + let snapshot_handle = snapshot.handle().clone(); tokio::task::spawn_blocking({ let session = session.clone(); - move || session.materialize_current_snapshot(&destination) + move || { + let result = session.materialize_snapshot(&snapshot_handle, &destination); + drop(materialization_claim); + result + } }) .await .map_err(|error| RemoteBackendError::Failed(format!("materialization worker failed: {error}")))? diff --git a/src/main.rs b/src/main.rs index 137a494..d32f37d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -197,8 +197,10 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option>(); let remote_tool_policies = env.project.remote.tools.iter().map(|tool| (tool.name.clone(), RemoteToolPolicy::new(tool.allow_args))).collect::>(); + let configured_remote_tool_names = env.project.remote.tools.iter().map(|tool| tool.name.clone()).collect::>(); let remote_tool_names = remote_tool_names(&env.project.remote.tools); let share_dir_owned = share_dir.to_path_buf(); @@ -242,7 +244,12 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option bunkerbox::remote::RemoteTargetId { } fn remote_tool_names(entries: &[RemoteToolSpec]) -> Vec { - entries.iter().map(|tool| tool.name.clone()).collect() + entries.iter().filter(|tool| tool.name == "make").map(|tool| tool.name.clone()).collect() } fn write_run_handoff(file: &mut File, path: &Path, session_id: vscomm::WorkspaceSessionId) -> Result<(), String> { diff --git a/src/remote.rs b/src/remote.rs index 1c5c869..0dbc7da 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -3,6 +3,7 @@ use std::collections::{BTreeMap, BTreeSet}; use std::future::Future; use std::pin::Pin; +use std::sync::Arc; use std::time::Duration; use tokio::sync::mpsc; @@ -20,7 +21,24 @@ pub const DEFAULT_REMOTE_ENVIRONMENT: &[&str] = &["CC", "CXX", "AR", "RUSTFLAGS" #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RequestId(pub [u8; 16]); -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct RemoteSnapshotId([u8; 16]); + +impl RemoteSnapshotId { + pub fn from_bytes(bytes: [u8; 16]) -> Self { + Self(bytes) + } + + pub fn as_bytes(&self) -> &[u8; 16] { + &self.0 + } + + pub fn is_zero(self) -> bool { + self.0 == [0; 16] + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub struct WorkspaceSessionId(pub [u8; 16]); #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -77,10 +95,13 @@ pub struct RemoteBuild { tool: RemoteTool, argv: Vec, env: Vec<(String, String)>, + snapshot_id: RemoteSnapshotId, } impl RemoteBuild { - pub fn new(cwd: WorkspaceRelativePath, tool: RemoteTool, argv: Vec, env: Vec<(String, String)>) -> Result { + pub fn new( + cwd: WorkspaceRelativePath, tool: RemoteTool, argv: Vec, env: Vec<(String, String)>, snapshot_id: RemoteSnapshotId, + ) -> Result { validate_count(argv.len(), MAX_REMOTE_ARG_COUNT, "remote argv")?; argv.iter().try_for_each(|arg| validate_string("remote argument", arg, MAX_REMOTE_ARG_BYTES))?; validate_count(env.len(), MAX_REMOTE_ENV_COUNT, "remote environment")?; @@ -88,7 +109,10 @@ impl RemoteBuild { validate_environment_key(key)?; validate_environment_value(value) })?; - Ok(Self { cwd, tool, argv, env }) + if snapshot_id.is_zero() { + return Err("remote snapshot ID must be nonzero".to_string()); + } + Ok(Self { cwd, tool, argv, env, snapshot_id }) } pub fn cwd(&self) -> &WorkspaceRelativePath { @@ -106,6 +130,10 @@ impl RemoteBuild { pub fn env(&self) -> &[(String, String)] { &self.env } + + pub fn snapshot_id(&self) -> RemoteSnapshotId { + self.snapshot_id + } } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -263,6 +291,7 @@ pub enum RemoteAuthorizationError { ForbiddenEnvironment(String), EnvironmentNotAllowed(String), DuplicateEnvironment(String), + SnapshotNotAllowed, } impl RemoteAuthorizationError { @@ -271,18 +300,23 @@ impl RemoteAuthorizationError { } } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Clone)] pub struct RemoteAuthorizationPolicy { allowed_target: RemoteTargetId, allowed_session: WorkspaceSessionId, allowed_tools: BTreeMap, environment: RemoteEnvironmentPolicy, + snapshot_authority: Option>, +} + +pub trait RemoteSnapshotAuthority: Send + Sync { + fn snapshot_available(&self, session: WorkspaceSessionId, snapshot_id: RemoteSnapshotId) -> bool; } impl RemoteAuthorizationPolicy { pub fn new(allowed_target: RemoteTargetId, allowed_session: WorkspaceSessionId, allowed_tools: Vec) -> Self { let allowed_tools = allowed_tools.into_iter().map(|tool| (tool, RemoteToolPolicy::new(true))).collect(); - Self { allowed_target, allowed_session, allowed_tools, environment: RemoteEnvironmentPolicy::default() } + Self { allowed_target, allowed_session, allowed_tools, environment: RemoteEnvironmentPolicy::default(), snapshot_authority: None } } pub fn from_policies( @@ -296,7 +330,12 @@ impl RemoteAuthorizationPolicy { return Err(format!("duplicate remote tool policy: {tool}")); } } - Ok(Self { allowed_target, allowed_session, allowed_tools, environment }) + Ok(Self { allowed_target, allowed_session, allowed_tools, environment, snapshot_authority: None }) + } + + pub fn with_snapshot_authority(mut self, authority: Arc) -> Self { + self.snapshot_authority = Some(authority); + self } pub fn allowed_tools(&self) -> impl Iterator { @@ -314,6 +353,9 @@ impl RemoteAuthorizationPolicy { let request = match request.operation() { RemoteOperation::Sync => request, RemoteOperation::Build(build) => { + if self.snapshot_authority.as_ref().is_none_or(|authority| !authority.snapshot_available(self.allowed_session, build.snapshot_id())) { + return Err(RemoteAuthorizationError::SnapshotNotAllowed); + } let Some(tool_policy) = self.allowed_tools.get(build.tool().as_str()).copied() else { return Err(RemoteAuthorizationError::ToolNotAllowed(build.tool().as_str().to_string())); }; @@ -321,7 +363,7 @@ impl RemoteAuthorizationPolicy { return Err(RemoteAuthorizationError::ToolArgumentsNotAllowed(build.tool().as_str().to_string())); } let environment = self.environment.filter(build.env())?; - let filtered = RemoteBuild::new(build.cwd.clone(), build.tool.clone(), build.argv.clone(), environment) + let filtered = RemoteBuild::new(build.cwd.clone(), build.tool.clone(), build.argv.clone(), environment, build.snapshot_id()) .map_err(RemoteAuthorizationError::InvalidEnvironment)?; RemoteRequest::build(request.request_id, request.workspace_session_id, filtered) } @@ -334,6 +376,7 @@ impl RemoteAuthorizationPolicy { #[derive(Debug, Clone, PartialEq, Eq)] pub enum RemoteBackendEvent { SyncProgress { completed_bytes: u64, total_bytes: Option }, + SyncCompleted { snapshot_id: RemoteSnapshotId }, Stdout(Vec), Stderr(Vec), Error { message: String }, diff --git a/src/remote_client.rs b/src/remote_client.rs index 16f226c..77912cc 100644 --- a/src/remote_client.rs +++ b/src/remote_client.rs @@ -1,21 +1,93 @@ +use crate::remote::RemoteSnapshotId; use crate::vscomm::{RemoteBuild, RemoteEvent, RemoteEventKind, RemoteRequest, RemoteTool, RequestId, WorkspaceRelativePath, WorkspaceSessionId}; +use rand::RngCore; +use std::env; use std::io::{Read, Write}; +use std::path::Path; pub fn remote_sync_request(request_id: RequestId, session_id: WorkspaceSessionId) -> RemoteRequest { RemoteRequest::sync(request_id, session_id) } +pub fn remote_session_from_env() -> Result { + env::var("BUNKERBOX_REMOTE_SESSION") + .map_err(|_| "BUNKERBOX_REMOTE_SESSION is missing".to_string()) + .and_then(|value| WorkspaceSessionId::from_hex(&value)) +} + +pub fn logical_workspace_cwd(path: &Path) -> Result { + let relative = path.strip_prefix("/workspace").map_err(|_| "current directory must be under /workspace".to_string())?; + let value = relative.to_str().ok_or_else(|| "current directory is not valid UTF-8".to_string())?.to_string(); + crate::remote::WorkspaceRelativePath::new(&value)?; + Ok(value) +} + +pub fn new_request_id() -> RequestId { + let mut bytes = [0; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + RequestId(bytes) +} + +pub fn remote_tool_enabled(tool: &str) -> bool { + env::var("BUNKERBOX_REMOTE_TOOLS").ok().is_some_and(|tools| tools.split(',').any(|candidate| candidate == tool)) +} + +pub fn remote_environment_names() -> Vec { + env::var("BUNKERBOX_REMOTE_ENV_NAMES") + .ok() + .map(|names| names.split(',').filter(|name| !name.is_empty()).map(str::to_string).collect()) + .unwrap_or_default() +} + +pub fn selected_remote_environment(names: impl IntoIterator) -> Vec<(String, String)> { + names + .into_iter() + .filter(|name| !never_forward_environment(name)) + .filter_map(|name| env::var_os(&name).and_then(|value| value.into_string().ok().map(|value| (name, value)))) + .collect() +} + +fn never_forward_environment(name: &str) -> bool { + let upper = name.to_ascii_uppercase(); + matches!(upper.as_str(), "PATH" | "HOME" | "SSH_AUTH_SOCK" | "SSH_AGENT_PID" | "GITHUB_TOKEN" | "GITLAB_TOKEN" | "NPM_TOKEN" | "KUBECONFIG") + || upper.starts_with("BUNKERBOX_") + || upper.starts_with("AWS_") + || upper.starts_with("GCP_") + || upper.starts_with("GOOGLE_") + || upper.starts_with("AZURE_") + || upper.starts_with("DOCKER_") + || upper.starts_with("CARGO_REGISTRIES_") + || upper.starts_with("XDG_") + || upper.ends_with("_PROXY") + || upper.ends_with("_TOKEN") + || upper.ends_with("_PASSWORD") + || upper.ends_with("_PASS") + || upper.ends_with("_SECRET") + || upper.ends_with("_KEY") +} + +#[cfg(test)] +#[path = "remote_client_ut.rs"] +mod tests; + pub fn remote_build_request( request_id: RequestId, session_id: WorkspaceSessionId, cwd: impl Into, tool: impl Into, argv: Vec, - env: Vec<(String, String)>, + env: Vec<(String, String)>, snapshot_id: RemoteSnapshotId, ) -> Result { - let build = RemoteBuild::new(WorkspaceRelativePath::new(cwd)?, RemoteTool::new(tool)?, argv, env)?; + let snapshot_id = crate::vscomm::RemoteSnapshotId(*snapshot_id.as_bytes()); + let build = RemoteBuild::new(WorkspaceRelativePath::new(cwd)?, RemoteTool::new(tool)?, argv, env, snapshot_id)?; Ok(RemoteRequest::build(request_id, session_id, build)) } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RemoteCompletion { + Synced(RemoteSnapshotId), + Completed(i32), +} + pub fn execute_remote_request_to( stream: &mut S, request: RemoteRequest, stdout: &mut WOut, stderr: &mut WErr, -) -> Result { +) -> Result { let request_id = request.request_id; request.to_frame()?.write(stream).map_err(|e| format!("send remote request: {e}"))?; @@ -27,6 +99,13 @@ pub fn execute_remote_request_to( } match event.kind { RemoteEventKind::SyncProgress { .. } => {} + RemoteEventKind::SyncCompleted { snapshot_id } => { + let snapshot_id = RemoteSnapshotId::from_bytes(snapshot_id.0); + if snapshot_id.is_zero() { + return Err("remote snapshot ID must be nonzero".to_string()); + } + return Ok(RemoteCompletion::Synced(snapshot_id)); + } RemoteEventKind::Stdout(data) => { stdout.write_all(&data).map_err(|e| format!("stdout: {e}"))?; stdout.flush().map_err(|e| format!("flush stdout: {e}"))?; @@ -37,7 +116,7 @@ pub fn execute_remote_request_to( } RemoteEventKind::Error { message, .. } => return Err(message), RemoteEventKind::Cancelled => return Err("remote operation cancelled".to_string()), - RemoteEventKind::Completed { exit_code } => return Ok(exit_code), + RemoteEventKind::Completed { exit_code } => return Ok(RemoteCompletion::Completed(exit_code)), } } } diff --git a/src/snapshot.rs b/src/snapshot.rs index bda717a..1df733b 100644 --- a/src/snapshot.rs +++ b/src/snapshot.rs @@ -170,7 +170,7 @@ impl SnapshotExclusionPolicy { } } -#[derive(Debug, Clone, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct SnapshotHandle { session_id: WorkspaceSessionId, snapshot_id: SnapshotId, diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index 0e1bced..8dc7041 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -17,7 +17,7 @@ pub const TUI_STATUS_PORT: u32 = 10000; pub const VSCOMM_BIN_DIR: &str = "/usr/local/bunkerbox/bin"; /// Maximum payload accepted in one vsock frame. pub const MAX_FRAME_PAYLOAD: usize = 1024 * 1024; -pub const REMOTE_PROTOCOL_VERSION: u16 = 1; +pub const REMOTE_PROTOCOL_VERSION: u16 = 2; pub const MAX_REMOTE_STRING_BYTES: usize = 4 * 1024; pub const MAX_REMOTE_TOOL_BYTES: usize = 256; pub const MAX_REMOTE_ARG_COUNT: usize = 256; @@ -66,6 +66,15 @@ pub struct ExecRequest { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RequestId(pub [u8; 16]); +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct RemoteSnapshotId(pub [u8; 16]); + +impl RemoteSnapshotId { + pub fn is_zero(self) -> bool { + self.0 == [0; 16] + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct WorkspaceSessionId(pub [u8; 16]); @@ -152,12 +161,18 @@ pub struct RemoteBuild { pub tool: RemoteTool, pub argv: Vec, pub env: Vec<(String, String)>, + pub snapshot_id: RemoteSnapshotId, } impl RemoteBuild { - pub fn new(cwd: WorkspaceRelativePath, tool: RemoteTool, argv: Vec, env: Vec<(String, String)>) -> Result { + pub fn new( + cwd: WorkspaceRelativePath, tool: RemoteTool, argv: Vec, env: Vec<(String, String)>, snapshot_id: RemoteSnapshotId, + ) -> Result { validate_remote_build_fields(&cwd, &tool, &argv, &env)?; - Ok(Self { cwd, tool, argv, env }) + if snapshot_id.is_zero() { + return Err("remote snapshot ID must be nonzero".to_string()); + } + Ok(Self { cwd, tool, argv, env, snapshot_id }) } } @@ -234,7 +249,8 @@ impl RemoteRequest { RemoteOperation::Build(build) => { let cwd = remote_domain::WorkspaceRelativePath::new(build.cwd.as_str())?; let tool = remote_domain::RemoteTool::new(build.tool.as_str())?; - let build = remote_domain::RemoteBuild::new(cwd, tool, build.argv, build.env)?; + let snapshot_id = remote_domain::RemoteSnapshotId::from_bytes(build.snapshot_id.0); + let build = remote_domain::RemoteBuild::new(cwd, tool, build.argv, build.env, snapshot_id)?; Ok(remote_domain::RemoteRequest::build(request_id, session_id, build)) } } @@ -243,8 +259,12 @@ impl RemoteRequest { fn encode_remote_build(writer: &mut WireWriter, build: &RemoteBuild) -> Result<(), String> { validate_remote_build_fields(&build.cwd, &build.tool, &build.argv, &build.env)?; + if build.snapshot_id.is_zero() { + return Err("remote snapshot ID must be nonzero".to_string()); + } writer.string(build.cwd.as_str(), MAX_REMOTE_STRING_BYTES, "remote cwd")?; writer.string(build.tool.as_str(), MAX_REMOTE_TOOL_BYTES, "remote tool")?; + writer.bytes(&build.snapshot_id.0); writer.count(build.argv.len(), MAX_REMOTE_ARG_COUNT, "remote argv")?; for arg in &build.argv { writer.string(arg, MAX_REMOTE_ARG_BYTES, "remote argument")?; @@ -260,6 +280,7 @@ fn encode_remote_build(writer: &mut WireWriter, build: &RemoteBuild) -> Result<( fn decode_remote_build(reader: &mut WireReader<'_>) -> Result { let cwd = WorkspaceRelativePath::new(reader.string(MAX_REMOTE_STRING_BYTES, "remote cwd")?)?; let tool = RemoteTool::new(reader.string(MAX_REMOTE_TOOL_BYTES, "remote tool")?)?; + let snapshot_id = RemoteSnapshotId(reader.array16()?); let argv = (0..reader.count(MAX_REMOTE_ARG_COUNT, "remote argv")?) .map(|_| reader.string(MAX_REMOTE_ARG_BYTES, "remote argument")) .collect::, _>>()?; @@ -271,7 +292,7 @@ fn decode_remote_build(reader: &mut WireReader<'_>) -> Result, String>>()?; - RemoteBuild::new(cwd, tool, argv, env) + RemoteBuild::new(cwd, tool, argv, env, snapshot_id) } fn validate_remote_build_fields(cwd: &WorkspaceRelativePath, tool: &RemoteTool, argv: &[String], env: &[(String, String)]) -> Result<(), String> { @@ -484,6 +505,7 @@ impl RemoteErrorCode { #[derive(Debug, Clone, PartialEq, Eq)] pub enum RemoteEventKind { SyncProgress { completed_bytes: u64, total_bytes: Option }, + SyncCompleted { snapshot_id: RemoteSnapshotId }, Stdout(Vec), Stderr(Vec), Error { code: RemoteErrorCode, message: String }, @@ -504,6 +526,9 @@ impl RemoteEvent { remote_domain::RemoteBackendEvent::SyncProgress { completed_bytes, total_bytes } => { RemoteEventKind::SyncProgress { completed_bytes, total_bytes } } + remote_domain::RemoteBackendEvent::SyncCompleted { snapshot_id } => { + RemoteEventKind::SyncCompleted { snapshot_id: RemoteSnapshotId(snapshot_id.as_bytes().to_owned()) } + } remote_domain::RemoteBackendEvent::Stdout(data) => RemoteEventKind::Stdout(data), remote_domain::RemoteBackendEvent::Stderr(data) => RemoteEventKind::Stderr(data), remote_domain::RemoteBackendEvent::Error { message } => RemoteEventKind::Error { code: RemoteErrorCode::Failed, message }, @@ -518,6 +543,7 @@ impl RemoteEvent { writer.u16(REMOTE_PROTOCOL_VERSION); writer.u8(match &self.kind { RemoteEventKind::SyncProgress { .. } => 1, + RemoteEventKind::SyncCompleted { .. } => 7, RemoteEventKind::Stdout(_) => 2, RemoteEventKind::Stderr(_) => 3, RemoteEventKind::Error { .. } => 4, @@ -535,6 +561,12 @@ impl RemoteEvent { writer.u64(*total_bytes); } } + RemoteEventKind::SyncCompleted { snapshot_id } => { + if snapshot_id.is_zero() { + return Err("remote snapshot ID must be nonzero".to_string()); + } + writer.bytes(&snapshot_id.0); + } RemoteEventKind::Stdout(data) | RemoteEventKind::Stderr(data) => writer.blob(data, MAX_FRAME_PAYLOAD, "remote output")?, RemoteEventKind::Error { code, message } => { writer.u16(*code as u16); @@ -568,6 +600,13 @@ impl RemoteEvent { }; RemoteEventKind::SyncProgress { completed_bytes, total_bytes } } + 7 => { + let snapshot_id = RemoteSnapshotId(reader.array16()?); + if snapshot_id.is_zero() { + return Err("remote snapshot ID must be nonzero".to_string()); + } + RemoteEventKind::SyncCompleted { snapshot_id } + } 2 => RemoteEventKind::Stdout(reader.blob(MAX_FRAME_PAYLOAD, "remote stdout")?), 3 => RemoteEventKind::Stderr(reader.blob(MAX_FRAME_PAYLOAD, "remote stderr")?), 4 => RemoteEventKind::Error { From e5899565b0f8c9938a72f02085f1c4fc82896e87 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 10:41:08 +0200 Subject: [PATCH 21/52] Update images --- Makefile | 10 +++++----- images/crush.conf | 1 + images/kilocode.conf | 1 + images/opencode.conf | 1 + 4 files changed, 8 insertions(+), 5 deletions(-) diff --git a/Makefile b/Makefile index a02fd41..2cdebb6 100644 --- a/Makefile +++ b/Makefile @@ -43,8 +43,7 @@ ensure-toolchain: dev: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image - cargo build --bin bunkerbox-netrelay --target $(VSCOMM_TARGET) - cargo build --bin bunkerbox-vscomm --target $(VSCOMM_TARGET) + cargo build --bin bunkerbox-netrelay --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) cargo build --bin bunkerbox-status --target $(VSCOMM_TARGET) rm -rf target/dist mkdir -p target/dist @@ -52,13 +51,13 @@ dev: ensure-toolchain cp target/debug/bunkerbox-image target/dist/ cp target/$(VSCOMM_TARGET)/debug/bunkerbox-netrelay target/dist/ cp target/$(VSCOMM_TARGET)/debug/bunkerbox-vscomm target/dist/ + cp target/$(VSCOMM_TARGET)/debug/bunkerbox-remote target/dist/ cp target/$(VSCOMM_TARGET)/debug/bunkerbox-status target/dist/ cp target/$(VSCOMM_TARGET)/debug/bunkerbox-netrelay target/debug/bunkerbox-netrelay release: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image --release - cargo build --bin bunkerbox-netrelay --target $(VSCOMM_TARGET) --release - cargo build --bin bunkerbox-vscomm --target $(VSCOMM_TARGET) --release + cargo build --bin bunkerbox-netrelay --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) --release cargo build --bin bunkerbox-status --target $(VSCOMM_TARGET) --release rm -rf target/dist mkdir -p target/dist @@ -66,6 +65,7 @@ release: ensure-toolchain cp target/release/bunkerbox-image target/dist/ cp target/$(VSCOMM_TARGET)/release/bunkerbox-netrelay target/dist/ cp target/$(VSCOMM_TARGET)/release/bunkerbox-vscomm target/dist/ + cp target/$(VSCOMM_TARGET)/release/bunkerbox-remote target/dist/ cp target/$(VSCOMM_TARGET)/release/bunkerbox-status target/dist/ cp target/$(VSCOMM_TARGET)/release/bunkerbox-netrelay target/release/bunkerbox-netrelay @@ -83,7 +83,7 @@ setup: dev target/debug/bunkerbox setup musl-vscomm: ensure-toolchain - cargo build --bin bunkerbox-vscomm --target $(VSCOMM_TARGET) + cargo build --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) cargo build --bin bunkerbox-status --target $(VSCOMM_TARGET) image: dev diff --git a/images/crush.conf b/images/crush.conf index b785c24..550676c 100644 --- a/images/crush.conf +++ b/images/crush.conf @@ -62,6 +62,7 @@ containerfile: | RUN chmod 0755 /usr/local/bin/bunker-entrypoint COPY bunkerbox-vscomm /usr/local/bunkerbox/bin/bunkerbox-vscomm + COPY bunkerbox-remote /usr/local/bunkerbox/bin/bunkerbox-remote COPY bunkerbox-status /usr/local/bunkerbox/bin/bunkerbox-status ENV HOME=/home/bunkerbox \ diff --git a/images/kilocode.conf b/images/kilocode.conf index c1d7356..cdf0325 100644 --- a/images/kilocode.conf +++ b/images/kilocode.conf @@ -62,6 +62,7 @@ containerfile: | RUN chmod 0755 /usr/local/bin/bunker-entrypoint COPY bunkerbox-vscomm /usr/local/bunkerbox/bin/bunkerbox-vscomm + COPY bunkerbox-remote /usr/local/bunkerbox/bin/bunkerbox-remote COPY bunkerbox-status /usr/local/bunkerbox/bin/bunkerbox-status diff --git a/images/opencode.conf b/images/opencode.conf index 7aa9f20..f6c1e0d 100644 --- a/images/opencode.conf +++ b/images/opencode.conf @@ -62,6 +62,7 @@ containerfile: | RUN chmod 0755 /usr/local/bin/bunker-entrypoint COPY bunkerbox-vscomm /usr/local/bunkerbox/bin/bunkerbox-vscomm + COPY bunkerbox-remote /usr/local/bunkerbox/bin/bunkerbox-remote COPY bunkerbox-status /usr/local/bunkerbox/bin/bunkerbox-status ENV HOME=/home/bunkerbox \ From 77e050bd5a59a24b1d755077e4f0ae3a9d69a7b9 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 10:41:20 +0200 Subject: [PATCH 22/52] Add unit tests for transparency layer --- src/bunkerbox-image_ut.rs | 24 +++++ src/bunkerbox-remote_ut.rs | 99 ++++++++++++++++++-- src/bunkerbox-vscomm_ut.rs | 15 ++- src/loopback_ut.rs | 184 ++++++++++++++++++++++++++++++++----- src/main_ut.rs | 4 +- src/remote_client_ut.rs | 18 ++++ src/remote_ut.rs | 36 ++++++-- src/vscomm/mod_ut.rs | 24 ++++- 8 files changed, 362 insertions(+), 42 deletions(-) create mode 100644 src/bunkerbox-image_ut.rs create mode 100644 src/remote_client_ut.rs diff --git a/src/bunkerbox-image_ut.rs b/src/bunkerbox-image_ut.rs new file mode 100644 index 0000000..7ac099b --- /dev/null +++ b/src/bunkerbox-image_ut.rs @@ -0,0 +1,24 @@ +use super::*; + +fn config() -> ImageConfig { + ImageConfig { + name: "test".into(), + image: "test:latest".into(), + output: "test.oci".into(), + command: Vec::new(), + overwrite: false, + build_args: BTreeMap::new(), + hooks: ImageHooks::default(), + files: Vec::new(), + runtime: None, + containerfile: "FROM scratch".into(), + } +} + +#[test] +fn image_entrypoint_installs_remote_make_before_local_vscomm_links() { + let script = build_entrypoint(&config()).unwrap(); + let remote = script.find("bunkerbox-remote\" install").unwrap(); + let vscomm = script.find("bunkerbox-vscomm\" install").unwrap(); + assert!(remote < vscomm); +} diff --git a/src/bunkerbox-remote_ut.rs b/src/bunkerbox-remote_ut.rs index 6b1604f..7409e69 100644 --- a/src/bunkerbox-remote_ut.rs +++ b/src/bunkerbox-remote_ut.rs @@ -1,6 +1,10 @@ use super::*; use bunkerbox::vscomm::{Frame, RemoteEvent, RemoteEventKind}; +fn snapshot_id() -> bunkerbox::remote::RemoteSnapshotId { + bunkerbox::remote::RemoteSnapshotId::from_bytes([9; 16]) +} + struct MemoryStream { input: io::Cursor>, output: Vec, @@ -53,6 +57,7 @@ fn build_request_preserves_logical_cwd_and_arguments() { "src".into(), RequestId([1; 16]), WorkspaceSessionId([2; 16]), + snapshot_id(), ) .unwrap(); let frame = request.to_frame().unwrap(); @@ -66,15 +71,18 @@ fn build_request_preserves_logical_cwd_and_arguments() { #[test] fn sync_success_uses_existing_remote_helper_and_returns_status() { let request_id = RequestId([3; 16]); - let mut stream = MemoryStream::new(vec![RemoteEvent { request_id, kind: RemoteEventKind::Completed { exit_code: 0 } }]); + let mut stream = MemoryStream::new(vec![RemoteEvent { + request_id, + kind: RemoteEventKind::SyncCompleted { snapshot_id: bunkerbox::vscomm::RemoteSnapshotId([9; 16]) }, + }]); let status = execute_remote_request_to( &mut stream, - build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16])).unwrap(), + build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16]), snapshot_id()).unwrap(), &mut Vec::new(), &mut Vec::new(), ) .unwrap(); - assert_eq!(status, 0); + assert_eq!(status, RemoteCompletion::Synced(snapshot_id())); assert!(Frame::read(&mut io::Cursor::new(stream.output)).is_ok()); } @@ -96,6 +104,7 @@ fn build_success_preserves_output_bytes_and_nonzero_exit_code() { "src".into(), request_id, WorkspaceSessionId([2; 16]), + snapshot_id(), ) .unwrap(), &mut stdout, @@ -103,7 +112,7 @@ fn build_success_preserves_output_bytes_and_nonzero_exit_code() { ) .unwrap(); - assert_eq!(status, 17); + assert_eq!(status, RemoteCompletion::Completed(17)); assert_eq!(stdout, vec![b'o', b'\n', 0xff]); assert_eq!(stderr, b"err\n"); } @@ -117,7 +126,7 @@ fn remote_failures_return_errors_without_local_fallback() { let mut stream = MemoryStream::new(vec![RemoteEvent { request_id, kind }]); let result = execute_remote_request_to( &mut stream, - build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16])).unwrap(), + build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16]), snapshot_id()).unwrap(), &mut Vec::new(), &mut Vec::new(), ); @@ -138,12 +147,13 @@ fn rejected_tool_and_mismatched_response_fail_closed() { String::new(), request_id, WorkspaceSessionId([2; 16]), + snapshot_id(), ) .unwrap(); assert_eq!(execute_remote_request_to(&mut stream, request, &mut Vec::new(), &mut Vec::new()).unwrap_err(), "authorization rejected"); let mut stream = MemoryStream::new(vec![RemoteEvent { request_id: RequestId([8; 16]), kind: RemoteEventKind::Completed { exit_code: 0 } }]); - let request = build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16])).unwrap(); + let request = build_request(RemoteCommand::Sync, String::new(), request_id, WorkspaceSessionId([2; 16]), snapshot_id()).unwrap(); assert_eq!(execute_remote_request_to(&mut stream, request, &mut Vec::new(), &mut Vec::new()).unwrap_err(), "remote event request ID mismatch"); } @@ -152,3 +162,80 @@ fn logical_cwd_is_workspace_relative_only() { assert_eq!(logical_workspace_cwd(Path::new("/workspace/project/src")).unwrap(), "project/src"); assert!(logical_workspace_cwd(Path::new("/tmp/project")).is_err()); } + +#[test] +fn configured_remote_make_installation_precedes_native_path_resolution() { + let root = tempfile::tempdir().unwrap(); + let executable = root.path().join("bunkerbox-remote"); + std::fs::write(&executable, b"remote").unwrap(); + install_remote_make_link(root.path(), &executable, true).unwrap(); + assert_eq!(std::fs::read_link(root.path().join("make")).unwrap(), executable); +} + +#[test] +fn disabled_remote_make_removes_only_its_managed_link() { + let root = tempfile::tempdir().unwrap(); + let executable = root.path().join("bunkerbox-remote"); + std::fs::write(&executable, b"remote").unwrap(); + install_remote_make_link(root.path(), &executable, true).unwrap(); + install_remote_make_link(root.path(), &executable, false).unwrap(); + assert!(!root.path().join("make").exists()); + + std::fs::write(root.path().join("make"), b"native").unwrap(); + install_remote_make_link(root.path(), &executable, false).unwrap(); + assert_eq!(std::fs::read(root.path().join("make")).unwrap(), b"native"); +} + +#[test] +fn configured_remote_make_does_not_overwrite_unmanaged_entry() { + let root = tempfile::tempdir().unwrap(); + let executable = root.path().join("bunkerbox-remote"); + std::fs::write(&executable, b"remote").unwrap(); + std::fs::write(root.path().join("make"), b"native").unwrap(); + assert!(install_remote_make_link(root.path(), &executable, true).is_err()); + assert_eq!(std::fs::read(root.path().join("make")).unwrap(), b"native"); +} + +#[test] +fn transparent_build_syncs_first_and_reuses_that_capability() { + let session = WorkspaceSessionId([2; 16]); + let capability = snapshot_id(); + let mut requests = Vec::new(); + let result = run_build_with_sync_using( + "src".into(), + "make".into(), + vec!["release".into(), "space arg".into()], + vec![("CC".into(), "cc".into())], + session, + |request| { + requests.push(request.clone()); + if requests.len() == 1 { + Ok(RemoteCompletion::Synced(capability)) + } else { + Ok(RemoteCompletion::Completed(17)) + } + }, + ) + .unwrap(); + + assert_eq!(result, 17); + assert_eq!(requests.len(), 2); + assert!(matches!(&requests[0].operation, bunkerbox::vscomm::RemoteOperation::Sync(_))); + let bunkerbox::vscomm::RemoteOperation::Build(ref build) = requests[1].operation else { panic!("expected build") }; + assert_eq!(build.tool.as_str(), "make"); + assert_eq!(build.cwd.as_str(), "src"); + assert_eq!(build.argv, ["release", "space arg"]); + assert_eq!(build.snapshot_id, bunkerbox::vscomm::RemoteSnapshotId([9; 16])); + assert_eq!(build.env, [("CC".into(), "cc".into())]); +} + +#[test] +fn transparent_build_does_not_build_after_sync_failure() { + let mut calls = 0; + let result = run_build_with_sync_using("src".into(), "make".into(), vec!["release".into()], Vec::new(), WorkspaceSessionId([2; 16]), |_| { + calls += 1; + Err("sync failed".to_string()) + }); + assert_eq!(result, Err("sync failed".to_string())); + assert_eq!(calls, 1); +} diff --git a/src/bunkerbox-vscomm_ut.rs b/src/bunkerbox-vscomm_ut.rs index b07d66c..a196740 100644 --- a/src/bunkerbox-vscomm_ut.rs +++ b/src/bunkerbox-vscomm_ut.rs @@ -1,5 +1,5 @@ use super::{handle_response, handle_response_to, Frame, FrameType}; -use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request}; +use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request, RemoteCompletion}; use bunkerbox::vscomm::{Frame as RemoteFrame, RemoteEvent, RemoteEventKind, RequestId, WorkspaceSessionId}; use std::io::{self, Read, Write}; @@ -71,7 +71,16 @@ fn preserve_exit_status() { #[test] fn explicit_remote_client_preserves_streams_status_and_request_id() { let request_id = RequestId([9; 16]); - let request = remote_build_request(request_id, WorkspaceSessionId([8; 16]), "src", "make", vec!["release mode".into()], vec![]).unwrap(); + let request = remote_build_request( + request_id, + WorkspaceSessionId([8; 16]), + "src", + "make", + vec!["release mode".into()], + vec![], + bunkerbox::remote::RemoteSnapshotId::from_bytes([9; 16]), + ) + .unwrap(); let responses = vec![ RemoteEvent { request_id, kind: RemoteEventKind::Stdout(b"out".to_vec()) }.to_frame().unwrap(), RemoteEvent { request_id, kind: RemoteEventKind::Stderr(b"err".to_vec()) }.to_frame().unwrap(), @@ -81,7 +90,7 @@ fn explicit_remote_client_preserves_streams_status_and_request_id() { let mut stdout = FlushWriter { bytes: Vec::new(), flushes: 0 }; let mut stderr = FlushWriter { bytes: Vec::new(), flushes: 0 }; - assert_eq!(execute_remote_request_to(&mut stream, request, &mut stdout, &mut stderr).unwrap(), 23); + assert_eq!(execute_remote_request_to(&mut stream, request, &mut stdout, &mut stderr).unwrap(), RemoteCompletion::Completed(23)); assert_eq!(stdout.bytes, b"out"); assert_eq!(stderr.bytes, b"err"); assert_eq!(stdout.flushes, 1); diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 456bd84..89d6390 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -1,5 +1,8 @@ use super::*; -use crate::remote::{RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteTool, RequestId, WorkspaceRelativePath}; +use crate::remote::{ + RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTool, RequestId, + WorkspaceRelativePath, +}; use tempfile::TempDir; fn fixture() -> (TempDir, Arc, RemoteTargetId, WorkspaceSessionId) { @@ -17,12 +20,20 @@ fn fixture() -> (TempDir, Arc, RemoteTargetId, WorkspaceSessio (temp, session, target, session_id) } +struct TestSnapshotAuthority; + +impl RemoteSnapshotAuthority for TestSnapshotAuthority { + fn snapshot_available(&self, _session: WorkspaceSessionId, _snapshot_id: RemoteSnapshotId) -> bool { + true + } +} + fn authorized_build( - target: RemoteTargetId, session: WorkspaceSessionId, tool: &str, args: Vec, env: Vec<(String, String)>, + target: RemoteTargetId, session: WorkspaceSessionId, tool: &str, args: Vec, env: Vec<(String, String)>, snapshot_id: RemoteSnapshotId, ) -> crate::remote::AuthorizedRemoteRequest { - let build = RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new(tool).unwrap(), args, env).unwrap(); + let build = RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new(tool).unwrap(), args, env, snapshot_id).unwrap(); let request = RemoteRequest::build(RequestId([3; 16]), session, build); - let policy = RemoteAuthorizationPolicy::new(target, session, vec![tool.to_string()]); + let policy = RemoteAuthorizationPolicy::new(target, session, vec![tool.to_string()]).with_snapshot_authority(Arc::new(TestSnapshotAuthority)); policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session }, request).unwrap() } @@ -34,6 +45,22 @@ async fn collect_events(mut receiver: mpsc::Receiver) -> Vec events } +async fn sync_capability(backend: &LoopbackBackend, target: RemoteTargetId, session: WorkspaceSessionId) -> RemoteSnapshotId { + let request = RemoteRequest::sync(RequestId([4; 16]), session); + let policy = RemoteAuthorizationPolicy::new(target, session, Vec::new()); + let authorized = policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session }, request).unwrap(); + let (events, receiver) = mpsc::channel(8); + assert_eq!(backend.execute(authorized, events).await, Ok(())); + collect_events(receiver) + .await + .into_iter() + .find_map(|event| match event { + RemoteBackendEvent::SyncCompleted { snapshot_id } => Some(snapshot_id), + _ => None, + }) + .unwrap() +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn sync_and_build_materialize_a_bound_snapshot() { let (_temp, session, target, session_id) = fixture(); @@ -48,13 +75,18 @@ async fn sync_and_build_materialize_a_bound_snapshot() { let policy = RemoteAuthorizationPolicy::new(target, session_id, Vec::new()); let authorized = policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, sync).unwrap(); assert_eq!(backend.execute(authorized, sync_tx).await, Ok(())); - assert_eq!( - collect_events(sync_rx).await, - vec![RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }, RemoteBackendEvent::Completed { exit_code: 0 },] - ); + let sync_events = collect_events(sync_rx).await; + assert!(matches!(sync_events.first(), Some(RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }))); + let snapshot_id = sync_events + .iter() + .find_map(|event| match event { + RemoteBackendEvent::SyncCompleted { snapshot_id } => Some(*snapshot_id), + _ => None, + }) + .unwrap(); let (build_tx, build_rx) = mpsc::channel(8); - let authorized = authorized_build(target, session_id, "printf", vec!["value with spaces:$(literal)".to_string()], Vec::new()); + let authorized = authorized_build(target, session_id, "printf", vec!["value with spaces:$(literal)".to_string()], Vec::new(), snapshot_id); assert_eq!(backend.execute(authorized, build_tx).await, Ok(())); assert_eq!( collect_events(build_rx).await, @@ -63,6 +95,100 @@ async fn sync_and_build_materialize_a_bound_snapshot() { assert!(fs::read_dir(&session_state.jobs_root).unwrap().next().is_none()); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn interleaved_syncs_build_their_own_snapshot_capabilities() { + let (_temp, session, target, session_id) = fixture(); + let tools = resolve_fixed_tools(["cat".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session.clone(), tools); + fs::write(session.workspace_root().join("src/input.txt"), b"A\n").unwrap(); + let snapshot_a = sync_capability(&backend, target, session_id).await; + fs::write(session.workspace_root().join("src/input.txt"), b"B\n").unwrap(); + let snapshot_b = sync_capability(&backend, target, session_id).await; + assert_ne!(snapshot_a, snapshot_b); + + let authority = session.clone(); + let authorize = |snapshot_id| { + let build = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("cat").unwrap(), + vec!["input.txt".into()], + Vec::new(), + snapshot_id, + ) + .unwrap(); + let request = RemoteRequest::build(RequestId([8; 16]), session_id, build); + RemoteAuthorizationPolicy::new(target, session_id, vec!["cat".into()]) + .with_snapshot_authority(authority.clone()) + .authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, request) + .unwrap() + }; + let request_a = authorize(snapshot_a); + let request_b = authorize(snapshot_b); + let (events_a, receiver_a) = mpsc::channel(8); + let (events_b, receiver_b) = mpsc::channel(8); + let (result_a, result_b) = tokio::join!(backend.execute(request_a, events_a), backend.execute(request_b, events_b)); + assert_eq!(result_a, Ok(())); + assert_eq!(result_b, Ok(())); + let output_a = collect_events(receiver_a).await; + let output_b = collect_events(receiver_b).await; + assert!(output_a.iter().any(|event| matches!(event, RemoteBackendEvent::Stdout(bytes) if bytes == b"A\n"))); + assert!(output_b.iter().any(|event| matches!(event, RemoteBackendEvent::Stdout(bytes) if bytes == b"B\n"))); + assert!(output_a.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); + assert!(output_b.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); + + let replay = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("cat").unwrap(), + vec!["input.txt".into()], + Vec::new(), + snapshot_a, + ) + .unwrap(); + let replay_request = RemoteRequest::build(RequestId([9; 16]), session_id, replay); + assert_eq!( + RemoteAuthorizationPolicy::new(target, session_id, vec!["cat".into()]) + .with_snapshot_authority(session.clone()) + .authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, replay_request), + Err(crate::remote::RemoteAuthorizationError::SnapshotNotAllowed) + ); + + let unknown = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("cat").unwrap(), + vec!["input.txt".into()], + Vec::new(), + RemoteSnapshotId::from_bytes([6; 16]), + ) + .unwrap(); + let unknown_request = RemoteRequest::build(RequestId([10; 16]), session_id, unknown); + assert_eq!( + RemoteAuthorizationPolicy::new(target, session_id, vec!["cat".into()]) + .with_snapshot_authority(session.clone()) + .authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, unknown_request), + Err(crate::remote::RemoteAuthorizationError::SnapshotNotAllowed) + ); + + let other_session = WorkspaceSessionId([7; 16]); + let cross_session = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("cat").unwrap(), + vec!["input.txt".into()], + Vec::new(), + snapshot_b, + ) + .unwrap(); + let cross_request = RemoteRequest::build(RequestId([11; 16]), other_session, cross_session); + assert_eq!( + RemoteAuthorizationPolicy::new(target, other_session, vec!["cat".into()]) + .with_snapshot_authority(session) + .authorize(&RemoteExecutionContext { target, workspace_session_id: other_session }, cross_request), + Err(crate::remote::RemoteAuthorizationError::SnapshotNotAllowed) + ); +} + #[test] fn unapproved_remote_environment_is_rejected_before_backend_execution() { let (_temp, _session, target, session_id) = fixture(); @@ -71,27 +197,38 @@ fn unapproved_remote_environment_is_rejected_before_backend_execution() { RemoteTool::new("printf").unwrap(), Vec::new(), vec![("UNTRUSTED".to_string(), "1".to_string())], + RemoteSnapshotId::from_bytes([9; 16]), ) .unwrap(); let request = RemoteRequest::build(RequestId([3; 16]), session_id, build); - let policy = RemoteAuthorizationPolicy::new(target, session_id, vec!["printf".to_string()]); + let policy = + RemoteAuthorizationPolicy::new(target, session_id, vec!["printf".to_string()]).with_snapshot_authority(Arc::new(TestSnapshotAuthority)); assert_eq!( policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, request), Err(crate::remote::RemoteAuthorizationError::EnvironmentNotAllowed("UNTRUSTED".to_string())) ); } +#[test] +fn fixed_tool_resolution_never_uses_guest_wrapper_path() { + let tools = resolve_fixed_tools(["make".to_string()]); + if let Some(path) = tools.get("make") { + assert!(matches!(path.to_str(), Some(value) if value == "/usr/local/bin/make" || value == "/usr/bin/make" || value == "/bin/make")); + assert!(!path.starts_with("/usr/local/bunkerbox/bin")); + } +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn build_preserves_argument_bytes_and_reports_nonzero_stderr() { let (_temp, session, target, session_id) = fixture(); - session.sync_snapshot().unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); let tools = resolve_fixed_tools(["ls".to_string()]); if tools.is_empty() { return; } let backend = LoopbackBackend::new(session, tools); let (events, receiver) = mpsc::channel(8); - let request = authorized_build(target, session_id, "ls", vec!["$(not-a-shell-argument)".to_string()], Vec::new()); + let request = authorized_build(target, session_id, "ls", vec!["$(not-a-shell-argument)".to_string()], Vec::new(), snapshot_id); assert_eq!(backend.execute(request, events).await, Ok(())); let events = collect_events(receiver).await; assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Stderr(bytes) if !bytes.is_empty()))); @@ -102,7 +239,7 @@ async fn build_preserves_argument_bytes_and_reports_nonzero_stderr() { async fn child_exit_does_not_wait_for_inherited_output_pipes() { let (_temp, session, target, session_id) = fixture(); fs::write(session.workspace_root().join("src/Makefile"), ".PHONY: leak\nleak:\n\t@sleep 30 & echo $$!\n\t@printf 'direct-output\\n'\n").unwrap(); - session.sync_snapshot().unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); let tools = resolve_fixed_tools(["make".to_string()]); if tools.is_empty() { return; @@ -110,11 +247,9 @@ async fn child_exit_does_not_wait_for_inherited_output_pipes() { let backend = LoopbackBackend::new(session.clone(), tools); let (events, receiver) = mpsc::channel(16); - let request = authorized_build(target, session_id, "make", vec!["leak".into()], Vec::new()); - let started = tokio::time::Instant::now(); - let result = tokio::time::timeout(Duration::from_secs(1), backend.execute(request, events)).await.unwrap(); + let request = authorized_build(target, session_id, "make", vec!["leak".into()], Vec::new(), snapshot_id); + let result = tokio::time::timeout(Duration::from_secs(2), backend.execute(request, events)).await.unwrap(); assert_eq!(result, Ok(())); - assert!(started.elapsed() < Duration::from_millis(500)); let events = collect_events(receiver).await; let stdout = events @@ -146,14 +281,14 @@ async fn child_exit_does_not_wait_for_inherited_output_pipes() { #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn child_receives_guest_environment_but_trusted_target_wins() { let (_temp, session, target, session_id) = fixture(); - session.sync_snapshot().unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); let tools = resolve_fixed_tools(["printenv".to_string()]); if tools.is_empty() { return; } let backend = LoopbackBackend::new(session, tools).with_target_environment(BTreeMap::from([("CC".into(), "trusted-target".into())])); let (events, receiver) = mpsc::channel(8); - let request = authorized_build(target, session_id, "printenv", vec!["CC".into()], vec![("CC".into(), "guest-value".into())]); + let request = authorized_build(target, session_id, "printenv", vec!["CC".into()], vec![("CC".into(), "guest-value".into())], snapshot_id); assert_eq!(backend.execute(request, events).await, Ok(())); let events = collect_events(receiver).await; assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Stdout(bytes) if bytes == b"trusted-target\n"))); @@ -163,36 +298,37 @@ async fn child_receives_guest_environment_but_trusted_target_wins() { #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn missing_tool_fails_before_execution() { let (_temp, session, target, session_id) = fixture(); + let snapshot_id = session.sync_snapshot().unwrap(); let backend = LoopbackBackend::new(session, BTreeMap::new()); let (events, _receiver) = mpsc::channel(8); - let request = authorized_build(target, session_id, "missing-tool", Vec::new(), Vec::new()); + let request = authorized_build(target, session_id, "missing-tool", Vec::new(), Vec::new(), snapshot_id); assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Spawn("loopback tool is not configured: missing-tool".to_string()))); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn timeout_kills_a_direct_child_process() { let (_temp, session, target, session_id) = fixture(); - session.sync_snapshot().unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); let tools = resolve_fixed_tools(["sleep".to_string()]); if tools.is_empty() { return; } let backend = LoopbackBackend::new(session, tools).with_timeout(Duration::from_millis(50)); let (events, _receiver) = mpsc::channel(8); - let request = authorized_build(target, session_id, "sleep", vec!["5".to_string()], Vec::new()); + let request = authorized_build(target, session_id, "sleep", vec!["5".to_string()], Vec::new(), snapshot_id); assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Timeout)); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn output_limit_kills_a_flooding_direct_child() { let (_temp, session, target, session_id) = fixture(); - session.sync_snapshot().unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); let tools = resolve_fixed_tools(["printf".to_string()]); if tools.is_empty() { return; } let backend = LoopbackBackend::new(session, tools).with_output_limit(8); let (events, _receiver) = mpsc::channel(8); - let request = authorized_build(target, session_id, "printf", vec!["0123456789".into()], Vec::new()); + let request = authorized_build(target, session_id, "printf", vec!["0123456789".into()], Vec::new(), snapshot_id); assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::OutputLimit { limit: 8 })); } diff --git a/src/main_ut.rs b/src/main_ut.rs index 2affc20..66bcb2b 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -59,10 +59,10 @@ fn run_handoff_rejects_zero_session() { } #[test] -fn remote_tool_names_preserve_configured_order() { +fn remote_tool_names_only_enables_the_make_wrapper() { assert_eq!( remote_tool_names(&[RemoteToolSpec { name: "make".into(), allow_args: true }, RemoteToolSpec { name: "cargo".into(), allow_args: false },]), - vec!["make", "cargo"] + vec!["make"] ); } diff --git a/src/remote_client_ut.rs b/src/remote_client_ut.rs new file mode 100644 index 0000000..93d15df --- /dev/null +++ b/src/remote_client_ut.rs @@ -0,0 +1,18 @@ +use super::*; +use std::path::Path; + +#[test] +fn logical_cwd_is_workspace_relative() { + assert_eq!(logical_workspace_cwd(Path::new("/workspace/foo/bar")).unwrap(), "foo/bar"); + assert!(logical_workspace_cwd(Path::new("/tmp/foo")).is_err()); +} + +#[test] +fn selected_environment_uses_only_targeted_names() { + std::env::set_var("BB_TEST_REMOTE_ALLOWED", "selected"); + std::env::set_var("BB_TEST_REMOTE_PATH", "should-not-forward"); + let values = selected_remote_environment(vec!["BB_TEST_REMOTE_ALLOWED".into(), "PATH".into(), "BB_TEST_REMOTE_PATH".into()]); + assert_eq!(values, vec![("BB_TEST_REMOTE_ALLOWED".into(), "selected".into()), ("BB_TEST_REMOTE_PATH".into(), "should-not-forward".into())]); + std::env::remove_var("BB_TEST_REMOTE_ALLOWED"); + std::env::remove_var("BB_TEST_REMOTE_PATH"); +} diff --git a/src/remote_ut.rs b/src/remote_ut.rs index 8e061f5..b977ba2 100644 --- a/src/remote_ut.rs +++ b/src/remote_ut.rs @@ -9,18 +9,32 @@ fn request(tool: &str) -> RemoteRequest { RemoteTool::new(tool).unwrap(), vec!["build".into()], vec![("CC".into(), "cc".into())], + RemoteSnapshotId::from_bytes([9; 16]), ) .unwrap(), ) } +struct TestSnapshotAuthority; + +impl RemoteSnapshotAuthority for TestSnapshotAuthority { + fn snapshot_available(&self, _session: WorkspaceSessionId, _snapshot_id: RemoteSnapshotId) -> bool { + true + } +} + +fn policy(tools: Vec) -> RemoteAuthorizationPolicy { + RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), tools) + .with_snapshot_authority(std::sync::Arc::new(TestSnapshotAuthority)) +} + fn context() -> RemoteExecutionContext { RemoteExecutionContext { target: RemoteTargetId([3; 16]), workspace_session_id: WorkspaceSessionId([2; 16]) } } #[test] fn policy_authorizes_typed_request() { - let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let policy = policy(vec!["make".into()]); let authorized = policy.authorize(&context(), request("make")).unwrap(); assert_eq!(authorized.request_id(), RequestId([1; 16])); @@ -30,11 +44,17 @@ fn policy_authorizes_typed_request() { #[test] fn policy_rejects_unapproved_tool() { - let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["cargo".into()]); + let policy = policy(vec!["cargo".into()]); assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::ToolNotAllowed("make".into()))); } +#[test] +fn policy_requires_snapshot_authority_for_builds() { + let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::SnapshotNotAllowed)); +} + #[test] fn backend_errors_have_typed_events() { assert_eq!(RemoteBackendError::Spawn("could not start".into()).event(), RemoteBackendEvent::Error { message: "could not start".into() }); @@ -44,7 +64,7 @@ fn backend_errors_have_typed_events() { #[test] fn environment_policy_preserves_allowed_entries() { - let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let policy = policy(vec!["make".into()]); let authorized = policy.authorize(&context(), request("make")).unwrap(); let RemoteOperation::Build(build) = authorized.request().operation() else { panic!("expected build") }; assert_eq!(build.env(), [("CC".into(), "cc".into())]); @@ -61,10 +81,11 @@ fn environment_policy_rejects_forbidden_and_unlisted_entries() { RemoteTool::new("make").unwrap(), Vec::new(), vec![(name.into(), "value".into())], + RemoteSnapshotId::from_bytes([9; 16]), ) .unwrap(); let request = RemoteRequest::build(RequestId([1; 16]), WorkspaceSessionId([2; 16]), build); - let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let policy = policy(vec!["make".into()]); assert_eq!(policy.authorize(&context(), request), Err(expected)); } } @@ -76,9 +97,10 @@ fn environment_policy_rejects_duplicates_and_control_data() { RemoteTool::new("make").unwrap(), Vec::new(), vec![("CC".into(), "one".into()), ("CC".into(), "two".into())], + RemoteSnapshotId::from_bytes([9; 16]), ) .unwrap(); - let policy = RemoteAuthorizationPolicy::new(RemoteTargetId([3; 16]), WorkspaceSessionId([2; 16]), vec!["make".into()]); + let policy = policy(vec!["make".into()]); let request = RemoteRequest::build(RequestId([1; 16]), WorkspaceSessionId([2; 16]), duplicate); assert_eq!(policy.authorize(&context(), request), Err(RemoteAuthorizationError::DuplicateEnvironment("CC".into()))); @@ -87,6 +109,7 @@ fn environment_policy_rejects_duplicates_and_control_data() { RemoteTool::new("make").unwrap(), Vec::new(), vec![("CC".into(), "bad\nvalue".into())], + RemoteSnapshotId::from_bytes([9; 16]), ) .is_err()); } @@ -99,6 +122,7 @@ fn command_policy_rejects_unapproved_arguments() { [("make".into(), RemoteToolPolicy::new(false))], RemoteEnvironmentPolicy::default(), ) - .unwrap(); + .unwrap() + .with_snapshot_authority(std::sync::Arc::new(TestSnapshotAuthority)); assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::ToolArgumentsNotAllowed("make".into()))); } diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index e7e021e..a38d760 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -19,7 +19,7 @@ fn workspace_session_hex_rejects_non_ascii_without_panicking() { } fn build(argv: Vec, env: Vec<(String, String)>) -> Result { - RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new("make").unwrap(), argv, env) + RemoteBuild::new(WorkspaceRelativePath::new("src").unwrap(), RemoteTool::new("make").unwrap(), argv, env, RemoteSnapshotId([9; 16])) } fn raw_build_frame(argv_count: u16, arg: Option<&str>, env_count: u16, env: Option<(&str, &str)>) -> Frame { @@ -99,6 +99,8 @@ fn remote_build_round_trips_structured_arguments_and_environment() { ); let decoded = RemoteRequest::from_frame(request.to_frame().unwrap()).unwrap(); assert_eq!(decoded, request); + let RemoteOperation::Build(build) = decoded.operation else { panic!("expected build") }; + assert_eq!(build.snapshot_id, RemoteSnapshotId([9; 16])); } #[test] @@ -106,6 +108,7 @@ fn every_remote_event_round_trips() { let request_id = ids().0; let events = vec![ RemoteEventKind::SyncProgress { completed_bytes: 4, total_bytes: Some(9) }, + RemoteEventKind::SyncCompleted { snapshot_id: RemoteSnapshotId([9; 16]) }, RemoteEventKind::Stdout(b"out".to_vec()), RemoteEventKind::Stderr(b"err".to_vec()), RemoteEventKind::Error { code: RemoteErrorCode::Failed, message: "failed".into() }, @@ -193,6 +196,24 @@ fn oversized_environment_key_and_value_are_rejected() { assert!(RemoteRequest::from_frame(raw_build_frame(0, None, 1, Some(("KEY", &"V".repeat(MAX_REMOTE_ENV_VALUE_BYTES + 1))))).is_err()); } +#[test] +fn zero_remote_snapshot_id_is_rejected() { + assert!(RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("make").unwrap(), + Vec::new(), + Vec::new(), + RemoteSnapshotId([0; 16]), + ) + .is_err()); +} + +#[test] +fn zero_sync_completion_snapshot_id_is_rejected() { + let event = RemoteEvent { request_id: ids().0, kind: RemoteEventKind::SyncCompleted { snapshot_id: RemoteSnapshotId([0; 16]) } }; + assert!(event.to_frame().is_err()); +} + #[test] fn control_data_in_remote_environment_value_is_rejected() { assert!(build(Vec::new(), vec![("CC".into(), "bad\nvalue".into())]).is_err()); @@ -231,6 +252,7 @@ fn protocol_request_converts_to_transport_independent_domain_request() { assert_eq!(build.tool().as_str(), "make"); assert_eq!(build.argv(), ["--release"]); assert_eq!(build.env(), [("MODE".into(), "debug".into())]); + assert_eq!(build.snapshot_id(), crate::remote::RemoteSnapshotId::from_bytes([9; 16])); } #[test] From 7321c3c7a6a8e85a45d87b96ad646e06143dcb45 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 11:04:15 +0200 Subject: [PATCH 23/52] Implement shared remote-make ownership and diagnostic sync cleanup with bounded capabilities --- src/bin/bunkerbox-remote.rs | 50 +++++++--------------- src/bin/bunkerbox-vscomm.rs | 46 ++------------------ src/guest_install.rs | 85 +++++++++++++++++++++++++++++++++++++ src/lib.rs | 1 + src/loopback.rs | 49 ++++++++++++++++++--- src/loopback_ut.rs | 68 +++++++++++++++++++++++++++++ src/remote.rs | 25 +++++++++-- src/remote_client.rs | 4 ++ src/vscomm/mod.rs | 37 +++++++++++++--- 9 files changed, 274 insertions(+), 91 deletions(-) create mode 100644 src/guest_install.rs diff --git a/src/bin/bunkerbox-remote.rs b/src/bin/bunkerbox-remote.rs index d3cd19e..da8ff45 100644 --- a/src/bin/bunkerbox-remote.rs +++ b/src/bin/bunkerbox-remote.rs @@ -1,7 +1,9 @@ +use bunkerbox::guest_install::install_remote_make_link; +#[cfg(test)] use bunkerbox::remote::RemoteSnapshotId; use bunkerbox::remote_client::{ - execute_remote_request_to, logical_workspace_cwd, new_request_id, remote_build_request, remote_environment_names, remote_session_from_env, - remote_sync_request, remote_tool_enabled, selected_remote_environment, RemoteCompletion, + execute_remote_request_to, logical_workspace_cwd, new_request_id, remote_build_request, remote_diagnostic_sync_request, remote_environment_names, + remote_session_from_env, remote_sync_request, remote_tool_enabled, selected_remote_environment, RemoteCompletion, }; #[cfg(test)] use bunkerbox::vscomm::RequestId; @@ -10,7 +12,6 @@ use std::env; use std::fs; use std::io::{self, Read, Write}; use std::mem; -use std::os::unix::fs::symlink; use std::path::Path; const HOST_CID: u32 = 2; @@ -101,10 +102,18 @@ fn execute_request_over_vsock(request: RemoteRequest) -> Result Result { - match execute_request_over_vsock(remote_sync_request(new_request_id(), session))? { - RemoteCompletion::Synced(snapshot_id) => Ok(snapshot_id), - RemoteCompletion::Completed(_) => Err("remote sync returned a build completion".to_string()), +fn sync_snapshot(session: WorkspaceSessionId) -> Result<(), String> { + sync_snapshot_using(session, execute_request_over_vsock) +} + +fn sync_snapshot_using(session: WorkspaceSessionId, mut execute: F) -> Result<(), String> +where + F: FnMut(RemoteRequest) -> Result, +{ + match execute(remote_diagnostic_sync_request(new_request_id(), session))? { + RemoteCompletion::Synced(_) => Err("diagnostic sync returned a retained capability".to_string()), + RemoteCompletion::Completed(0) => Ok(()), + RemoteCompletion::Completed(code) => Err(format!("remote sync returned exit code {code}")), } } @@ -134,33 +143,6 @@ fn install_remote_links() -> Result<(), String> { install_remote_make_link(Path::new(VSCOMM_BIN_DIR), &executable, remote_tool_enabled("make")) } -fn install_remote_make_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { - let target = bin_dir.join("make"); - let managed = is_managed_link(&target, executable); - - if enabled { - if target.exists() || fs::symlink_metadata(&target).is_ok() { - if !managed { - return Err(format!("cannot install remote make wrapper over existing {}", target.display())); - } - fs::remove_file(&target).map_err(|error| format!("remove existing remote make wrapper: {error}"))?; - } - symlink(executable, &target).map_err(|error| format!("symlink remote make wrapper: {error}"))?; - } else if managed { - fs::remove_file(&target).map_err(|error| format!("remove disabled remote make wrapper: {error}"))?; - } - Ok(()) -} - -fn is_managed_link(target: &Path, executable: &Path) -> bool { - let Ok(metadata) = fs::symlink_metadata(target) else { return false }; - if !metadata.file_type().is_symlink() { - return false; - } - let Ok(link) = fs::read_link(target) else { return false }; - link == executable -} - fn connect_toolchain() -> Result { vsock_connect(HOST_CID, TOOLCHAIN_PORT).map_err(|error| format!("toolchain vsock connect: {error}")) } diff --git a/src/bin/bunkerbox-vscomm.rs b/src/bin/bunkerbox-vscomm.rs index aea40b5..ea68edd 100644 --- a/src/bin/bunkerbox-vscomm.rs +++ b/src/bin/bunkerbox-vscomm.rs @@ -7,9 +7,9 @@ use std::env; use std::fs; use std::io::{self, Read, Write}; use std::mem; -use std::os::unix::fs::PermissionsExt; use std::path::{Path, PathBuf}; +use bunkerbox::guest_install::install_vscomm_links; pub use bunkerbox::remote_client::{execute_remote_request_to, remote_build_request, remote_sync_request}; use vscomm::{encode_ui_payload, validate_exec_request, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT, TUI_STATUS_PORT, VSCOMM_BIN_DIR}; @@ -101,27 +101,10 @@ fn install_symlinks() -> Result<(), String> { let config_path = find_config().ok_or_else(|| "no whitelist config found".to_string())?; let entries = read_whitelist_entries(&config_path)?; - - fs::create_dir_all(VSCOMM_BIN_DIR).map_err(|e| format!("mkdir {VSCOMM_BIN_DIR}: {e}"))?; - let vscomm_path = env::current_exe().map_err(|e| format!("failed to locate vscomm binary: {e}"))?; - - for entry in &entries { - let cmd = extract_command_name(entry); - if cmd.is_empty() { - continue; - } - if command_exists_in_path_except(&cmd, &vscomm_path) { - continue; - } - let target = PathBuf::from(VSCOMM_BIN_DIR).join(&cmd); - if target.exists() { - let _ = fs::remove_file(&target); - } - std::os::unix::fs::symlink(&vscomm_path, &target).map_err(|e| format!("symlink {cmd}: {e}"))?; - } - - Ok(()) + let commands = entries.iter().map(|entry| extract_command_name(entry)); + let path = env::var("PATH").unwrap_or_default(); + install_vscomm_links(commands, Path::new(VSCOMM_BIN_DIR), &vscomm_path, &path) } fn find_config() -> Option { @@ -180,27 +163,6 @@ fn extract_command_name(entry: &str) -> String { } } -fn command_exists_in_path_except(cmd: &str, except: &Path) -> bool { - if let Ok(path) = env::var("PATH") { - for dir in path.split(':') { - let candidate = Path::new(dir).join(cmd); - if candidate == except { - continue; - } - if candidate.is_file() { - let metadata = match fs::metadata(&candidate) { - Ok(m) => m, - Err(_) => continue, - }; - if metadata.permissions().mode() & 0o111 != 0 { - return true; - } - } - } - } - false -} - fn vsock_connect(cid: u32, port: u32) -> io::Result { unsafe { let fd = libc::socket(libc::AF_VSOCK, libc::SOCK_STREAM, 0); diff --git a/src/guest_install.rs b/src/guest_install.rs new file mode 100644 index 0000000..35d2807 --- /dev/null +++ b/src/guest_install.rs @@ -0,0 +1,85 @@ +use std::ffi::OsStr; +use std::fs; +use std::os::unix::fs::{symlink, PermissionsExt}; +use std::path::{Path, PathBuf}; + +const REMOTE_MAKE_OWNER: &str = "bunkerbox-remote"; + +pub fn install_remote_make_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { + let target = bin_dir.join("make"); + let managed = is_managed_remote_make_link(&target); + + if enabled { + if let Ok(link) = fs::read_link(&target) { + if link == executable { + return Ok(()); + } + } + if fs::symlink_metadata(&target).is_ok() { + if !managed { + return Err(format!("cannot install remote make wrapper over existing {}", target.display())); + } + fs::remove_file(&target).map_err(|error| format!("remove existing remote make wrapper: {error}"))?; + } + symlink(executable, &target).map_err(|error| format!("symlink remote make wrapper: {error}"))?; + } else if managed { + fs::remove_file(&target).map_err(|error| format!("remove disabled remote make wrapper: {error}"))?; + } + Ok(()) +} + +pub fn install_vscomm_links(commands: impl IntoIterator, bin_dir: &Path, vscomm_path: &Path, path: &str) -> Result<(), String> { + fs::create_dir_all(bin_dir).map_err(|error| format!("mkdir {}: {error}", bin_dir.display()))?; + + for command in commands { + if command.is_empty() { + continue; + } + let target = bin_dir.join(&command); + if command == "make" && is_managed_remote_make_link(&target) { + continue; + } + if command_exists_in_path_except(&command, vscomm_path, path) { + continue; + } + if let Ok(link) = fs::read_link(&target) { + if link == vscomm_path { + continue; + } + } + if fs::symlink_metadata(&target).is_ok() { + fs::remove_file(&target).map_err(|error| format!("remove existing {command} link: {error}"))?; + } + symlink(vscomm_path, &target).map_err(|error| format!("symlink {command}: {error}"))?; + } + + Ok(()) +} + +fn is_managed_remote_make_link(target: &Path) -> bool { + let Ok(metadata) = fs::symlink_metadata(target) else { return false }; + if !metadata.file_type().is_symlink() { + return false; + } + fs::read_link(target).ok().and_then(|link| link.file_name().map(OsStr::to_owned)).is_some_and(|name| name == REMOTE_MAKE_OWNER) +} + +fn command_exists_in_path_except(command: &str, except: &Path, path: &str) -> bool { + for directory in path.split(':') { + let candidate = PathBuf::from(directory).join(command); + if candidate == except { + continue; + } + if candidate.is_file() { + let Ok(metadata) = fs::metadata(&candidate) else { continue }; + if metadata.permissions().mode() & 0o111 != 0 { + return true; + } + } + } + false +} + +#[cfg(test)] +#[path = "guest_install_ut.rs"] +mod tests; diff --git a/src/lib.rs b/src/lib.rs index c7b0059..e4edf99 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -3,6 +3,7 @@ pub mod cfgsetup; pub mod clidef; pub mod cmdrun; pub mod daemon; +pub mod guest_install; pub mod kata; pub mod logging; pub mod netrelay; diff --git a/src/loopback.rs b/src/loopback.rs index 9c3fb62..8fd9fab 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -18,6 +18,7 @@ use tokio::sync::mpsc; use tokio::time::sleep; pub const LOOPBACK_BUILD_TIMEOUT: Duration = crate::remote::DEFAULT_REMOTE_BUILD_TIMEOUT; +pub const MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES: usize = 16; const LOOPBACK_PATH: &str = "/usr/local/bin:/usr/bin:/bin"; const POST_CHILD_EXIT_DRAIN_TIMEOUT: Duration = Duration::from_millis(100); const POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT: Duration = Duration::from_millis(100); @@ -80,13 +81,35 @@ impl RunRemoteSession { } pub fn sync_snapshot(&self) -> Result { + self.sync_snapshot_for_request(true)?.ok_or_else(|| "retained snapshot capability was not created".to_string()) + } + + fn sync_snapshot_for_request(&self, retain_capability: bool) -> Result, String> { let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; let snapshot = self.snapshot_builder.build_root(&self.workspace_root, self.session_id)?; - self.register_snapshot(snapshot.handle().clone()) + let handle = snapshot.handle().clone(); + if !retain_capability { + self.discard_unclaimed_snapshot(&handle)?; + return Ok(None); + } + + match self.register_snapshot(handle.clone()) { + Ok(snapshot_id) => Ok(Some(snapshot_id)), + Err(error) => { + let cleanup = self.discard_unclaimed_snapshot(&handle); + if let Err(cleanup_error) = cleanup { + return Err(format!("{error}; snapshot cleanup failed: {cleanup_error}")); + } + Err(error) + } + } } fn register_snapshot(&self, handle: SnapshotHandle) -> Result { let mut registry = self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; + if registry.capabilities.len() >= MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES { + return Err(format!("remote snapshot capability limit reached ({MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES})")); + } let snapshot_id = loop { let mut bytes = [0u8; 16]; rand::thread_rng().fill_bytes(&mut bytes); @@ -100,6 +123,16 @@ impl RunRemoteSession { Ok(snapshot_id) } + fn discard_unclaimed_snapshot(&self, handle: &SnapshotHandle) -> Result<(), String> { + let referenced = + self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?.references.contains_key(handle); + if referenced { + Ok(()) + } else { + self.snapshot_store.remove(handle) + } + } + fn claim_snapshot(self: &Arc, snapshot_id: RemoteSnapshotId) -> Result { let mut registry = self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; let handle = registry.capabilities.remove(&snapshot_id).ok_or_else(|| "remote snapshot capability is unavailable".to_string())?; @@ -243,20 +276,24 @@ impl RemoteBackend for LoopbackBackend { let resources = self.resources; Box::pin(async move { match request.request().operation() { - RemoteOperation::Sync => execute_sync(session, events).await, + RemoteOperation::Sync(sync) => execute_sync(session, sync.retain_capability(), events).await, RemoteOperation::Build(build) => execute_build(session, tools, target_environment, resources, build, events).await, } }) } } -async fn execute_sync(session: Arc, events: mpsc::Sender) -> Result<(), RemoteBackendError> { +async fn execute_sync( + session: Arc, retain_capability: bool, events: mpsc::Sender, +) -> Result<(), RemoteBackendError> { send_event(&events, RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }).await?; - let sync = tokio::task::spawn_blocking(move || session.sync_snapshot()) + let sync = tokio::task::spawn_blocking(move || session.sync_snapshot_for_request(retain_capability)) .await .map_err(|error| RemoteBackendError::Failed(format!("snapshot worker failed: {error}")))?; - let snapshot_id = sync.map_err(RemoteBackendError::Failed)?; - send_event(&events, RemoteBackendEvent::SyncCompleted { snapshot_id }).await + match sync.map_err(RemoteBackendError::Failed)? { + Some(snapshot_id) => send_event(&events, RemoteBackendEvent::SyncCompleted { snapshot_id }).await, + None => send_event(&events, RemoteBackendEvent::Completed { exit_code: 0 }).await, + } } async fn execute_build( diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 89d6390..cf280ab 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -61,6 +61,24 @@ async fn sync_capability(backend: &LoopbackBackend, target: RemoteTargetId, sess .unwrap() } +async fn diagnostic_sync(backend: &LoopbackBackend, target: RemoteTargetId, session: WorkspaceSessionId) -> Vec { + let request = RemoteRequest::diagnostic_sync(RequestId([5; 16]), session); + let policy = RemoteAuthorizationPolicy::new(target, session, Vec::new()); + let authorized = policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session }, request).unwrap(); + let (events, receiver) = mpsc::channel(8); + assert_eq!(backend.execute(authorized, events).await, Ok(())); + collect_events(receiver).await +} + +fn published_snapshot_count(session: &RunRemoteSession) -> usize { + fs::read_dir(session.snapshot_store_root()) + .unwrap() + .filter_map(Result::ok) + .filter(|entry| entry.file_name() != ".staging") + .map(|entry| fs::read_dir(entry.path()).map(|entries| entries.filter_map(Result::ok).count()).unwrap_or_default()) + .sum() +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn sync_and_build_materialize_a_bound_snapshot() { let (_temp, session, target, session_id) = fixture(); @@ -95,6 +113,56 @@ async fn sync_and_build_materialize_a_bound_snapshot() { assert!(fs::read_dir(&session_state.jobs_root).unwrap().next().is_none()); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn diagnostic_sync_releases_capability_and_snapshot_storage() { + let (_temp, session, target, session_id) = fixture(); + let backend = LoopbackBackend::new(session.clone(), BTreeMap::new()); + + for _ in 0..(MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES * 2) { + let events = diagnostic_sync(&backend, target, session_id).await; + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); + assert!(!events.iter().any(|event| matches!(event, RemoteBackendEvent::SyncCompleted { .. }))); + assert!(session.snapshot_capabilities.lock().unwrap().capabilities.is_empty()); + assert_eq!(published_snapshot_count(&session), 0); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn failed_diagnostic_sync_registers_nothing() { + let (_temp, session, target, session_id) = fixture(); + fs::remove_dir_all(session.workspace_root()).unwrap(); + let backend = LoopbackBackend::new(session.clone(), BTreeMap::new()); + let request = RemoteRequest::diagnostic_sync(RequestId([5; 16]), session_id); + let policy = RemoteAuthorizationPolicy::new(target, session_id, Vec::new()); + let authorized = policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session_id }, request).unwrap(); + let (events, _receiver) = mpsc::channel(8); + + assert!(matches!(backend.execute(authorized, events).await, Err(RemoteBackendError::Failed(_)))); + assert!(session.snapshot_capabilities.lock().unwrap().capabilities.is_empty()); + assert_eq!(published_snapshot_count(&session), 0); +} + +#[test] +fn outstanding_snapshot_capabilities_are_bounded_and_consumption_frees_a_slot() { + let (temp, session, _target, _session_id) = fixture(); + let mut capabilities = Vec::new(); + for _ in 0..MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES { + capabilities.push(session.sync_snapshot().unwrap()); + } + assert_eq!(session.snapshot_capabilities.lock().unwrap().capabilities.len(), MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES); + assert!(session.sync_snapshot().unwrap_err().contains("capability limit reached")); + + let claim = session.claim_snapshot(capabilities.pop().unwrap()).unwrap(); + drop(claim); + assert_eq!(session.snapshot_capabilities.lock().unwrap().capabilities.len(), MAX_OUTSTANDING_SNAPSHOT_CAPABILITIES - 1); + session.sync_snapshot().unwrap(); + + let snapshot_root = session.snapshot_store_root(); + drop(session); + assert!(!snapshot_root.exists()); + drop(temp); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn interleaved_syncs_build_their_own_snapshot_capabilities() { let (_temp, session, target, session_id) = fixture(); diff --git a/src/remote.rs b/src/remote.rs index 0dbc7da..5f7e3f9 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -220,9 +220,20 @@ impl RemoteToolPolicy { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteSync { + retain_capability: bool, +} + +impl RemoteSync { + pub fn retain_capability(self) -> bool { + self.retain_capability + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum RemoteOperation { - Sync, + Sync(RemoteSync), Build(RemoteBuild), } @@ -235,7 +246,15 @@ pub struct RemoteRequest { impl RemoteRequest { pub fn sync(request_id: RequestId, workspace_session_id: WorkspaceSessionId) -> Self { - Self { request_id, workspace_session_id, operation: RemoteOperation::Sync } + Self::sync_with_capability(request_id, workspace_session_id, true) + } + + pub fn diagnostic_sync(request_id: RequestId, workspace_session_id: WorkspaceSessionId) -> Self { + Self::sync_with_capability(request_id, workspace_session_id, false) + } + + fn sync_with_capability(request_id: RequestId, workspace_session_id: WorkspaceSessionId, retain_capability: bool) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Sync(RemoteSync { retain_capability }) } } pub fn build(request_id: RequestId, workspace_session_id: WorkspaceSessionId, build: RemoteBuild) -> Self { @@ -351,7 +370,7 @@ impl RemoteAuthorizationPolicy { } let request = match request.operation() { - RemoteOperation::Sync => request, + RemoteOperation::Sync(_) => request, RemoteOperation::Build(build) => { if self.snapshot_authority.as_ref().is_none_or(|authority| !authority.snapshot_available(self.allowed_session, build.snapshot_id())) { return Err(RemoteAuthorizationError::SnapshotNotAllowed); diff --git a/src/remote_client.rs b/src/remote_client.rs index 77912cc..c67a281 100644 --- a/src/remote_client.rs +++ b/src/remote_client.rs @@ -9,6 +9,10 @@ pub fn remote_sync_request(request_id: RequestId, session_id: WorkspaceSessionId RemoteRequest::sync(request_id, session_id) } +pub fn remote_diagnostic_sync_request(request_id: RequestId, session_id: WorkspaceSessionId) -> RemoteRequest { + RemoteRequest::diagnostic_sync(request_id, session_id) +} + pub fn remote_session_from_env() -> Result { env::var("BUNKERBOX_REMOTE_SESSION") .map_err(|_| "BUNKERBOX_REMOTE_SESSION is missing".to_string()) diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index 8dc7041..d444dfd 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -75,6 +75,11 @@ impl RemoteSnapshotId { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteSync { + pub retain_capability: bool, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct WorkspaceSessionId(pub [u8; 16]); @@ -182,9 +187,6 @@ pub enum RemoteOperation { Build(RemoteBuild), } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct RemoteSync; - #[derive(Debug, Clone, PartialEq, Eq)] pub struct RemoteRequest { pub request_id: RequestId, @@ -194,7 +196,15 @@ pub struct RemoteRequest { impl RemoteRequest { pub fn sync(request_id: RequestId, workspace_session_id: WorkspaceSessionId) -> Self { - Self { request_id, workspace_session_id, operation: RemoteOperation::Sync(RemoteSync) } + Self::sync_with_capability(request_id, workspace_session_id, true) + } + + pub fn diagnostic_sync(request_id: RequestId, workspace_session_id: WorkspaceSessionId) -> Self { + Self::sync_with_capability(request_id, workspace_session_id, false) + } + + fn sync_with_capability(request_id: RequestId, workspace_session_id: WorkspaceSessionId, retain_capability: bool) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Sync(RemoteSync { retain_capability }) } } pub fn build(request_id: RequestId, workspace_session_id: WorkspaceSessionId, build: RemoteBuild) -> Self { @@ -214,6 +224,8 @@ impl RemoteRequest { if let RemoteOperation::Build(build) = &self.operation { encode_remote_build(&mut writer, build)?; + } else if let RemoteOperation::Sync(sync) = &self.operation { + writer.u8(u8::from(sync.retain_capability)); } writer.into_frame(FrameType::RemoteRequest) @@ -232,7 +244,14 @@ impl RemoteRequest { let request_id = RequestId(reader.array16()?); let workspace_session_id = WorkspaceSessionId(reader.array16()?); let operation = match operation_kind { - 1 => RemoteOperation::Sync(RemoteSync), + 1 => { + let retain_capability = match reader.u8()? { + 0 => false, + 1 => true, + value => return Err(format!("invalid remote sync capability flag: {value}")), + }; + RemoteOperation::Sync(RemoteSync { retain_capability }) + } 2 => RemoteOperation::Build(decode_remote_build(&mut reader)?), value => return Err(format!("unknown remote operation: {value}")), }; @@ -245,7 +264,13 @@ impl RemoteRequest { let request_id = remote_domain::RequestId(self.request_id.0); let session_id = remote_domain::WorkspaceSessionId(self.workspace_session_id.0); match self.operation { - RemoteOperation::Sync(_) => Ok(remote_domain::RemoteRequest::sync(request_id, session_id)), + RemoteOperation::Sync(sync) => { + if sync.retain_capability { + Ok(remote_domain::RemoteRequest::sync(request_id, session_id)) + } else { + Ok(remote_domain::RemoteRequest::diagnostic_sync(request_id, session_id)) + } + } RemoteOperation::Build(build) => { let cwd = remote_domain::WorkspaceRelativePath::new(build.cwd.as_str())?; let tool = remote_domain::RemoteTool::new(build.tool.as_str())?; From 4e2d7ae1cbffa76d97259b64884755fedc78a76b Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 3 Aug 2026 11:04:31 +0200 Subject: [PATCH 24/52] Add regression tests for installers, cleanup, bounds, and preserved binding --- src/bunkerbox-remote_ut.rs | 11 + src/daemon_ut.rs | 483 +++++++++++++++++++++++++------------ src/guest_install_ut.rs | 82 +++++++ src/vscomm/mod_ut.rs | 19 ++ 4 files changed, 437 insertions(+), 158 deletions(-) create mode 100644 src/guest_install_ut.rs diff --git a/src/bunkerbox-remote_ut.rs b/src/bunkerbox-remote_ut.rs index 7409e69..b5fbacf 100644 --- a/src/bunkerbox-remote_ut.rs +++ b/src/bunkerbox-remote_ut.rs @@ -86,6 +86,17 @@ fn sync_success_uses_existing_remote_helper_and_returns_status() { assert!(Frame::read(&mut io::Cursor::new(stream.output)).is_ok()); } +#[test] +fn standalone_sync_uses_diagnostic_non_retaining_request() { + let session = WorkspaceSessionId([2; 16]); + sync_snapshot_using(session, |request| { + let bunkerbox::vscomm::RemoteOperation::Sync(sync) = request.operation else { panic!("expected sync") }; + assert!(!sync.retain_capability); + Ok(RemoteCompletion::Completed(0)) + }) + .unwrap(); +} + #[test] fn build_success_preserves_output_bytes_and_nonzero_exit_code() { let request_id = RequestId([6; 16]); diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index 83ea583..367332f 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -1,45 +1,20 @@ -use super::{build_command, find_netrelay_binary, make_proxy_runtime_dir, monitor_bwrap_status, ChildEvent, SandboxProxyConfig, VsockSession}; +use super::{dispatch_remote_frame, is_allowed, RemoteBroker, RemoteDispatchError}; +use super::{monitor_bwrap_status, ChildEvent}; use crate::cfg::EnvMode; -use crate::sandbox::{MergedProfile, NetworkMode}; -use crate::vscomm::{validate_exec_request, ExecRequest}; -use std::ffi::OsStr; +use crate::remote::{ + AuthorizedRemoteRequest, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionContext, RemoteFuture, + RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, RequestId, WorkspaceRelativePath, WorkspaceSessionId, +}; +use crate::vscomm::{ + Frame, RemoteBuild as WireRemoteBuild, RemoteRequest as WireRemoteRequest, RemoteTool as WireRemoteTool, RequestId as WireRequestId, + WorkspaceRelativePath as WireWorkspaceRelativePath, WorkspaceSessionId as WireWorkspaceSessionId, +}; use std::io::Write; -use std::os::unix::fs::PermissionsExt; -use std::path::PathBuf; -use std::sync::Arc; - -fn session_no_proxy() -> VsockSession { - VsockSession { - passthrough: Arc::new(vec!["cargo *".into()]), - env_mode: EnvMode::Relaxed, - workspace: PathBuf::from("/tmp/ws"), - merged_profile: Some(Arc::new(MergedProfile { name: "test".into(), network: NetworkMode::None, ..Default::default() })), - proxy_config: None, - } -} - -fn session_with_proxy() -> VsockSession { - VsockSession { - passthrough: Arc::new(vec!["cargo *".into()]), - env_mode: EnvMode::Relaxed, - workspace: PathBuf::from("/tmp/ws"), - merged_profile: Some(Arc::new(MergedProfile { name: "test".into(), network: NetworkMode::None, ..Default::default() })), - proxy_config: Some(Arc::new(SandboxProxyConfig { - socket_path: PathBuf::from("/tmp/proxy.sock"), - netrelay_path: PathBuf::from("/tmp/bunkerbox-netrelay"), - })), - } -} - -fn session_no_profile() -> VsockSession { - VsockSession { - passthrough: Arc::new(vec!["cargo *".into()]), - env_mode: EnvMode::Relaxed, - workspace: PathBuf::from("/tmp/ws"), - merged_profile: None, - proxy_config: None, - } -} +use std::path::Path; +use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll}; +use tokio::io::AsyncWrite; +use tokio::sync::mpsc; #[test] fn bwrap_status_reports_command_start() { @@ -65,147 +40,339 @@ fn bwrap_status_reports_setup_failure_without_child() { assert!(matches!(rx.try_recv().unwrap(), ChildEvent::LauncherFailed(_))); } -#[test] -fn a_profile_no_allowlist_has_unshare_net_no_proxy() { - let req = ExecRequest { cwd: "/workspace".into(), command: "cargo".into(), args: vec!["build".into()], env: vec![] }; - validate_exec_request(&req).unwrap(); - let session = session_no_proxy(); - let cmd = build_command(&session, &req, &PathBuf::from("/tmp/ws"), "/workspace").unwrap(); - let cmd = cmd.as_std(); - let args: Vec<_> = cmd.get_args().collect(); - let args_str: Vec = args.iter().map(|a| a.to_string_lossy().to_string()).collect(); - assert!(args_str.contains(&"--unshare-net".to_string())); - assert!(!args_str.contains(&"/run/bunkerbox/netrelay".to_string())); - assert!(!args_str.contains(&"/run/bunkerbox/proxy.sock".to_string())); - assert!(!args_str.contains(&"--setenv".to_string()) || !args_str.iter().any(|a| a.contains("HTTP_PROXY"))); +struct RecordingBackend { + calls: Mutex>, + emit: Vec, + result: Option, } -#[test] -fn b_profile_allowlist_has_unshare_net_and_relay() { - let req = ExecRequest { cwd: "/workspace".into(), command: "cargo".into(), args: vec!["build".into()], env: vec![] }; - validate_exec_request(&req).unwrap(); - let session = session_with_proxy(); - let cmd = build_command(&session, &req, &PathBuf::from("/tmp/ws"), "/workspace").unwrap(); - let cmd = cmd.as_std(); - let args: Vec<_> = cmd.get_args().collect(); - let args_str: Vec = args.iter().map(|a| a.to_string_lossy().to_string()).collect(); - assert!(args_str.contains(&"--unshare-net".to_string())); - assert!(args_str.contains(&"/run/bunkerbox/netrelay".to_string())); - assert!(args_str.contains(&"/run/bunkerbox/proxy.sock".to_string())); - assert!(args_str.contains(&"--socket".to_string())); - assert!(args_str.iter().any(|a| a.contains("HTTP_PROXY"))); +struct TestSnapshotAuthority; + +impl RemoteSnapshotAuthority for TestSnapshotAuthority { + fn snapshot_available(&self, _session: WorkspaceSessionId, _snapshot_id: RemoteSnapshotId) -> bool { + true + } } -#[test] -fn c_no_profile_allowlist_direct_host_unchanged() { - let req = ExecRequest { cwd: "/workspace".into(), command: "cargo".into(), args: vec!["build".into()], env: vec![] }; - validate_exec_request(&req).unwrap(); - let mut session = session_no_profile(); - session.proxy_config = - Some(Arc::new(SandboxProxyConfig { socket_path: PathBuf::from("/tmp/proxy.sock"), netrelay_path: PathBuf::from("/tmp/bunkerbox-netrelay") })); - let cmd = build_command(&session, &req, &PathBuf::from("/tmp/ws"), "/workspace").unwrap(); - let cmd = cmd.as_std(); - let args: Vec<_> = cmd.get_args().collect(); - let args_str: Vec = args.iter().map(|a| a.to_string_lossy().to_string()).collect(); - assert_eq!(cmd.get_program(), "cargo"); - assert!(args_str.contains(&"build".to_string())); - assert!(!args_str.contains(&"--unshare-net".to_string())); - assert!(!args_str.iter().any(|a| a.contains("HTTP_PROXY"))); +struct RejectSnapshotAuthority; + +impl RemoteSnapshotAuthority for RejectSnapshotAuthority { + fn snapshot_available(&self, _session: WorkspaceSessionId, _snapshot_id: RemoteSnapshotId) -> bool { + false + } } -#[test] -fn d_critical_regression_no_proxy_with_unshare_net() { - let req = ExecRequest { cwd: "/workspace".into(), command: "cargo".into(), args: vec!["build".into()], env: vec![] }; - validate_exec_request(&req).unwrap(); +impl RemoteBackend for RecordingBackend { + fn execute<'a>( + &'a self, request: AuthorizedRemoteRequest, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + self.calls.lock().unwrap().push(request); + let emit = self.emit.clone(); + let result = self.result.clone(); + Box::pin(async move { + for event in emit { + events.send(event).await.map_err(|_| RemoteBackendError::Cancelled)?; + } + result.map_or(Ok(()), Err) + }) + } +} - // No proxy -> --unshare-net present - let session = session_no_proxy(); - let cmd = build_command(&session, &req, &PathBuf::from("/tmp/ws"), "/workspace").unwrap(); - let args: Vec<_> = cmd.as_std().get_args().collect(); - let args_str: Vec = args.iter().map(|a| a.to_string_lossy().to_string()).collect(); - assert!(args_str.contains(&"--unshare-net".to_string())); +struct StreamingBackend { + event_count: usize, + release: Arc, +} - // Proxy -> --unshare-net STILL present - let session = session_with_proxy(); - let cmd = build_command(&session, &req, &PathBuf::from("/tmp/ws"), "/workspace").unwrap(); - let args: Vec<_> = cmd.as_std().get_args().collect(); - let args_str: Vec = args.iter().map(|a| a.to_string_lossy().to_string()).collect(); - assert!(args_str.contains(&"--unshare-net".to_string())); +impl RemoteBackend for StreamingBackend { + fn execute<'a>( + &'a self, _request: AuthorizedRemoteRequest, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + let event_count = self.event_count; + let release = self.release.clone(); + Box::pin(async move { + for index in 0..event_count { + events.send(RemoteBackendEvent::Stdout(vec![index as u8])).await.map_err(|_| RemoteBackendError::Cancelled)?; + } + release.notified().await; + events.send(RemoteBackendEvent::Completed { exit_code: 0 }).await.map_err(|_| RemoteBackendError::Cancelled) + }) + } } -#[test] -fn e_runtime_dir_exclusive_and_private() { - let dir = make_proxy_runtime_dir().unwrap(); - assert!(dir.exists()); - let meta = std::fs::symlink_metadata(&dir).unwrap(); - assert!(meta.is_dir()); - let mode = meta.permissions().mode(); - assert_eq!(mode & 0o777, 0o700); - std::fs::remove_dir(&dir).unwrap(); +struct HangingBackend; + +impl RemoteBackend for HangingBackend { + fn execute<'a>( + &'a self, _request: AuthorizedRemoteRequest, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + Box::pin(async move { + events.send(RemoteBackendEvent::Stdout(b"first".to_vec())).await.map_err(|_| RemoteBackendError::Cancelled)?; + std::future::pending::>().await + }) + } } -#[test] -fn e_runtime_dir_rejects_existing() { - let dir = make_proxy_runtime_dir().unwrap(); - let result = make_proxy_runtime_dir(); - // dir still exists from first call -> create fails (not the same name but - // proves the function works when path is available) - std::fs::remove_dir(&dir).unwrap(); - assert!(result.is_ok()); +struct FailingWriter; + +impl AsyncWrite for FailingWriter { + fn poll_write(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>, _buf: &[u8]) -> Poll> { + Poll::Ready(Err(std::io::Error::new(std::io::ErrorKind::BrokenPipe, "writer closed"))) + } + + fn poll_flush(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown(self: std::pin::Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } } -#[test] -fn e_runtime_dir_rejects_existing_file() { - let tmp = std::env::temp_dir().join(format!("bunkerbox-daemon-test-file-{}", std::process::id())); - std::fs::write(&tmp, "data").unwrap(); - let meta = std::fs::symlink_metadata(&tmp).unwrap(); - assert!(meta.is_file()); - let _ = std::fs::remove_file(&tmp); +fn remote_request(tool: &str) -> RemoteRequest { + RemoteRequest::build( + RequestId([1; 16]), + WorkspaceSessionId([2; 16]), + crate::remote::RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + crate::remote::RemoteTool::new(tool).unwrap(), + vec!["build".into(), "--release".into()], + vec![("CC".into(), "cc".into())], + RemoteSnapshotId::from_bytes([9; 16]), + ) + .unwrap(), + ) } -#[test] -fn f_missing_netrelay_fails_closed() { - let exe = std::env::current_exe().unwrap(); - let dir = exe.parent().unwrap().join("nonexistent-dir-for-test"); - let path = dir.join("bunkerbox-netrelay"); - assert!(!path.is_file()); - // find_netrelay_binary looks for sibling -> succeeds if sibling exists, - // fails if not. This test proves a missing sibling returns Err. - // We can't test missing_from_nonexistent_dir without modifying the - // function, but the code path is: sibling doesn't exist -> Err. - // This is a structural test: assert the function returns Err when sibling absent. - // Since the sibling may actually exist (if built), we just verify the function - // name and error message pattern. - assert!(!path.exists()); +fn remote_broker(backend: Arc) -> RemoteBroker { + let target = RemoteTargetId([3; 16]); + let session = WorkspaceSessionId([2; 16]); + RemoteBroker::new( + RemoteAuthorizationPolicy::new(target, session, vec!["make".into()]).with_snapshot_authority(Arc::new(TestSnapshotAuthority)), + RemoteExecutionContext { target, workspace_session_id: session }, + backend, + ) } -#[test] -fn g_make_proxy_runtime_dir_rejects_existing_path() { - let existing = std::env::temp_dir().join(format!("bunkerbox-daemon-test-{}", std::process::id())); - std::fs::create_dir(&existing).unwrap(); - let exists = existing.exists(); - assert!(exists); +#[tokio::test] +async fn authorized_remote_request_reaches_typed_backend() { + let backend = Arc::new(RecordingBackend { + calls: Mutex::new(Vec::new()), + emit: vec![ + RemoteBackendEvent::Stdout(b"out".to_vec()), + RemoteBackendEvent::Stderr(b"err".to_vec()), + RemoteBackendEvent::Completed { exit_code: 7 }, + ], + result: None, + }); + let broker = remote_broker(backend.clone()); + let (tx, mut rx) = mpsc::channel(8); + + broker.dispatch(remote_request("make"), tx).await.unwrap(); + + assert_eq!(backend.calls.lock().unwrap().len(), 1); + let request = backend.calls.lock().unwrap()[0].request().clone(); + let crate::remote::RemoteOperation::Build(build) = request.operation() else { panic!("expected build") }; + assert_eq!(build.tool().as_str(), "make"); + assert_eq!(build.cwd().as_str(), "src"); + assert_eq!(build.argv(), ["build", "--release"]); + assert_eq!(build.env(), [("CC".into(), "cc".into())]); + assert_eq!(rx.recv().await, Some(RemoteBackendEvent::Stdout(b"out".to_vec()))); + assert_eq!(rx.recv().await, Some(RemoteBackendEvent::Stderr(b"err".to_vec()))); + assert_eq!(rx.recv().await, Some(RemoteBackendEvent::Completed { exit_code: 7 })); +} + +#[tokio::test] +async fn rejected_remote_request_never_calls_backend() { + let backend = Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }); + let broker = remote_broker(backend.clone()); + let (tx, mut rx) = mpsc::channel(8); + + let error = broker.dispatch(remote_request("cargo"), tx).await.unwrap_err(); + + assert!(matches!(error, RemoteDispatchError::Unauthorized(_))); + assert!(backend.calls.lock().unwrap().is_empty()); + assert!(matches!(rx.recv().await, Some(RemoteBackendEvent::Error { message }) if message.contains("authorization rejected"))); +} + +#[tokio::test] +async fn rejected_snapshot_capability_never_calls_backend() { + let backend = Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }); + let target = RemoteTargetId([3; 16]); + let session = WorkspaceSessionId([2; 16]); + let policy = RemoteAuthorizationPolicy::new(target, session, vec!["make".into()]).with_snapshot_authority(Arc::new(RejectSnapshotAuthority)); + let broker = RemoteBroker::new(policy, RemoteExecutionContext { target, workspace_session_id: session }, backend.clone()); + let (tx, mut rx) = mpsc::channel(8); + + let error = broker.dispatch(remote_request("make"), tx).await.unwrap_err(); + assert_eq!(error, RemoteDispatchError::Unauthorized(crate::remote::RemoteAuthorizationError::SnapshotNotAllowed)); + assert!(backend.calls.lock().unwrap().is_empty()); + assert!(matches!(rx.recv().await, Some(RemoteBackendEvent::Error { message }) if message.contains("SnapshotNotAllowed"))); +} + +#[tokio::test] +async fn backend_failure_is_reported_as_typed_error() { + let backend = + Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: Some(RemoteBackendError::Spawn("not found".into())) }); + let broker = remote_broker(backend); + let (tx, mut rx) = mpsc::channel(8); + + assert!(matches!(broker.dispatch(remote_request("make"), tx).await, Err(RemoteDispatchError::Backend(RemoteBackendError::Spawn(_))))); + assert_eq!(rx.recv().await, Some(RemoteBackendEvent::Error { message: "not found".into() })); +} + +#[tokio::test] +async fn fake_backend_can_report_timeout_and_cancellation_states() { + for (failure, expected) in [ + (RemoteBackendError::Timeout, RemoteBackendEvent::Error { message: "remote backend timed out".into() }), + (RemoteBackendError::Cancelled, RemoteBackendEvent::Cancelled), + ] { + let backend = Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: Some(failure.clone()) }); + let broker = remote_broker(backend); + let (tx, mut rx) = mpsc::channel(8); + + assert!(matches!(broker.dispatch(remote_request("make"), tx).await, Err(RemoteDispatchError::Backend(error)) if error == failure)); + assert_eq!(rx.recv().await, Some(expected)); + } +} + +#[tokio::test] +async fn framed_remote_sync_runs_full_dispatch_and_event_conversion_chain() { + let backend = Arc::new(RecordingBackend { + calls: Mutex::new(Vec::new()), + emit: vec![RemoteBackendEvent::SyncCompleted { snapshot_id: RemoteSnapshotId::from_bytes([9; 16]) }], + result: None, + }); + let broker = remote_broker(backend.clone()); + let request = WireRemoteRequest::sync(WireRequestId([6; 16]), WireWorkspaceSessionId([2; 16])); + let (mut guest, mut host) = tokio::io::duplex(4096); + + dispatch_remote_frame(request.to_frame().unwrap(), &broker, &mut host).await.unwrap(); + + let event = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + assert_eq!(event.request_id, crate::vscomm::RequestId([6; 16])); + assert_eq!(event.kind, crate::vscomm::RemoteEventKind::SyncCompleted { snapshot_id: crate::vscomm::RemoteSnapshotId([9; 16]) }); + let calls = backend.calls.lock().unwrap(); + assert_eq!(calls.len(), 1); + assert!(matches!(calls[0].request().operation(), crate::remote::RemoteOperation::Sync(_))); +} + +#[tokio::test] +async fn framed_remote_build_preserves_typed_fields_and_output_order() { + let backend = Arc::new(RecordingBackend { + calls: Mutex::new(Vec::new()), + emit: vec![ + RemoteBackendEvent::Stdout(b"out".to_vec()), + RemoteBackendEvent::Stderr(b"err".to_vec()), + RemoteBackendEvent::Completed { exit_code: 23 }, + ], + result: None, + }); + let broker = remote_broker(backend.clone()); + let request = WireRemoteRequest::build( + WireRequestId([7; 16]), + WireWorkspaceSessionId([2; 16]), + WireRemoteBuild::new( + WireWorkspaceRelativePath::new("src").unwrap(), + WireRemoteTool::new("make").unwrap(), + vec!["release mode".into(), "$(literal)".into()], + vec![("CC".into(), "cc".into())], + crate::vscomm::RemoteSnapshotId([9; 16]), + ) + .unwrap(), + ); + let (mut guest, mut host) = tokio::io::duplex(4096); + + dispatch_remote_frame(request.to_frame().unwrap(), &broker, &mut host).await.unwrap(); + + let stdout = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + let stderr = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + let completed = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + assert_eq!(stdout.kind, crate::vscomm::RemoteEventKind::Stdout(b"out".to_vec())); + assert_eq!(stderr.kind, crate::vscomm::RemoteEventKind::Stderr(b"err".to_vec())); + assert_eq!(completed.kind, crate::vscomm::RemoteEventKind::Completed { exit_code: 23 }); + + let calls = backend.calls.lock().unwrap(); + let crate::remote::RemoteOperation::Build(build) = calls[0].request().operation() else { panic!("expected build") }; + assert_eq!(calls[0].request_id(), crate::remote::RequestId([7; 16])); + assert_eq!(build.cwd().as_str(), "src"); + assert_eq!(build.tool().as_str(), "make"); + assert_eq!(build.argv(), ["release mode", "$(literal)"]); + assert_eq!(build.env(), [("CC".into(), "cc".into())]); +} + +#[tokio::test] +async fn malformed_remote_frame_fails_before_backend_dispatch() { + let backend = Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }); + let broker = remote_broker(backend.clone()); + let mut frame = WireRemoteRequest::sync(WireRequestId([6; 16]), WireWorkspaceSessionId([2; 16])).to_frame().unwrap(); + frame.payload[6] = 99; + let (_, mut host) = tokio::io::duplex(128); + + assert!(dispatch_remote_frame(frame, &broker, &mut host).await.is_err()); + assert!(backend.calls.lock().unwrap().is_empty()); +} - std::fs::remove_dir(&existing).unwrap(); +#[tokio::test] +async fn remote_events_stream_past_bounded_channel_capacity_before_completion() { + let release = Arc::new(tokio::sync::Notify::new()); + let backend = Arc::new(StreamingBackend { event_count: 65, release: release.clone() }); + let broker = remote_broker(backend); + let request = WireRemoteRequest::sync(WireRequestId([6; 16]), WireWorkspaceSessionId([2; 16])); + let (mut guest, mut host) = tokio::io::duplex(8192); + let dispatch = tokio::spawn(async move { dispatch_remote_frame(request.to_frame().unwrap(), &broker, &mut host).await }); + + let first = tokio::time::timeout(std::time::Duration::from_secs(1), Frame::read_async(&mut guest)).await.unwrap().unwrap(); + let first = crate::vscomm::RemoteEvent::from_frame(first).unwrap(); + assert_eq!(first.kind, crate::vscomm::RemoteEventKind::Stdout(vec![0])); + release.notify_one(); + + let result = tokio::time::timeout(std::time::Duration::from_secs(1), dispatch).await.unwrap().unwrap(); + result.unwrap(); + for index in 1..65 { + let event = crate::vscomm::RemoteEvent::from_frame( + tokio::time::timeout(std::time::Duration::from_secs(1), Frame::read_async(&mut guest)).await.unwrap().unwrap(), + ) + .unwrap(); + assert_eq!(event.kind, crate::vscomm::RemoteEventKind::Stdout(vec![index as u8])); + } + let completed = crate::vscomm::RemoteEvent::from_frame( + tokio::time::timeout(std::time::Duration::from_secs(1), Frame::read_async(&mut guest)).await.unwrap().unwrap(), + ) + .unwrap(); + assert_eq!(completed.kind, crate::vscomm::RemoteEventKind::Completed { exit_code: 0 }); +} + +#[tokio::test] +async fn writer_failure_cancels_hanging_backend_without_waiting_forever() { + let broker = remote_broker(Arc::new(HangingBackend)); + let request = WireRemoteRequest::sync(WireRequestId([6; 16]), WireWorkspaceSessionId([2; 16])); + let result = + tokio::time::timeout(std::time::Duration::from_secs(1), dispatch_remote_frame(request.to_frame().unwrap(), &broker, &mut FailingWriter)) + .await + .unwrap(); + + assert!(result.is_err()); } #[test] -fn h_literal_argv_preserved() { - let req = ExecRequest { cwd: "/workspace".into(), command: "make".into(), args: vec!["A=a b".into(), "$HOME".into(), "x;y".into()], env: vec![] }; - validate_exec_request(&req).unwrap(); - let session = session_with_proxy(); - let cmd = build_command(&session, &req, &PathBuf::from("/tmp/ws"), "/workspace").unwrap(); - let args: Vec<_> = cmd.as_std().get_args().collect(); - assert!(args.iter().any(|a| *a == OsStr::new("A=a b"))); - assert!(args.iter().any(|a| *a == OsStr::new("$HOME"))); - assert!(args.iter().any(|a| *a == OsStr::new("x;y"))); +fn local_passthrough_authorization_remains_separate() { + assert!(is_allowed(&["make *".into()], "make", &["--release".into()])); + assert!(!is_allowed(&["make *".into()], "cargo", &["build".into()])); } #[test] -fn i_static_netrelay_smoke() { - // find_netrelay_binary returns Ok if sibling exists - let result = find_netrelay_binary(); - if let Ok(path) = &result { - assert!(path.is_file()); - } +fn local_exec_request_still_builds_on_the_local_path() { + let workspace = tempfile::tempdir().unwrap(); + let cwd = crate::workspace::WorkspaceCwd::resolve(workspace.path(), Path::new("/workspace")).unwrap(); + let session = super::VsockSession { + passthrough: Arc::new(vec!["true *".into()]), + env_mode: EnvMode::Paranoid, + workspace: workspace.path().to_path_buf(), + merged_profile: None, + has_proxy: false, + remote_broker: Arc::new(remote_broker(Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }))), + }; + let request = crate::vscomm::ExecRequest { cwd: "/workspace".into(), command: "true".into(), args: Vec::new(), env: Vec::new() }; + + assert!(super::build_command(&session, &request, &cwd).is_ok()); } diff --git a/src/guest_install_ut.rs b/src/guest_install_ut.rs new file mode 100644 index 0000000..a3ab65e --- /dev/null +++ b/src/guest_install_ut.rs @@ -0,0 +1,82 @@ +use super::*; +use std::os::unix::fs::PermissionsExt; + +fn executable(path: &Path) { + fs::write(path, b"binary").unwrap(); + fs::set_permissions(path, fs::Permissions::from_mode(0o755)).unwrap(); +} + +#[test] +fn remote_make_ownership_survives_local_passthrough_install() { + let root = tempfile::tempdir().unwrap(); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_make_link(root.path(), &remote, true).unwrap(); + install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, "").unwrap(); + + assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); + assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); +} + +#[test] +fn disabled_remote_make_preserves_native_and_vscomm_behavior() { + let root = tempfile::tempdir().unwrap(); + let native = root.path().join("native"); + fs::create_dir(&native).unwrap(); + let native_make = native.join("make"); + executable(&native_make); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_make_link(root.path(), &remote, false).unwrap(); + install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + + assert_eq!(fs::read_link(root.path().join("make")).unwrap_err().kind(), std::io::ErrorKind::NotFound); + assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); +} + +#[test] +fn remote_make_wins_when_native_make_is_present() { + let root = tempfile::tempdir().unwrap(); + let native = root.path().join("native"); + fs::create_dir(&native).unwrap(); + executable(&native.join("make")); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_make_link(root.path(), &remote, true).unwrap(); + install_vscomm_links(vec!["make".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + + assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); +} + +#[test] +fn repeated_install_is_idempotent_and_stale_managed_links_are_replaced() { + let root = tempfile::tempdir().unwrap(); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_make_link(root.path(), &remote, true).unwrap(); + install_remote_make_link(root.path(), &remote, true).unwrap(); + install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, "").unwrap(); + install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, "").unwrap(); + assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); + assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); + + fs::remove_file(root.path().join("make")).unwrap(); + symlink(root.path().join("old/bunkerbox-remote"), root.path().join("make")).unwrap(); + install_remote_make_link(root.path(), &remote, true).unwrap(); + assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); + + install_remote_make_link(root.path(), &remote, false).unwrap(); + assert!(fs::symlink_metadata(root.path().join("make")).is_err()); +} diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index a38d760..2712a11 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -87,6 +87,25 @@ fn remote_sync_round_trips() { let request = RemoteRequest::sync(request_id, session_id); let decoded = RemoteRequest::from_frame(request.to_frame().unwrap()).unwrap(); assert_eq!(decoded, request); + let RemoteOperation::Sync(sync) = decoded.operation else { panic!("expected sync") }; + assert!(sync.retain_capability); +} + +#[test] +fn diagnostic_remote_sync_does_not_retain_capability() { + let (request_id, session_id) = ids(); + let request = RemoteRequest::diagnostic_sync(request_id, session_id); + let decoded = RemoteRequest::from_frame(request.to_frame().unwrap()).unwrap(); + let RemoteOperation::Sync(sync) = decoded.operation else { panic!("expected sync") }; + assert!(!sync.retain_capability); +} + +#[test] +fn invalid_remote_sync_capability_flag_is_rejected() { + let (request_id, session_id) = ids(); + let mut frame = RemoteRequest::diagnostic_sync(request_id, session_id).to_frame().unwrap(); + frame.payload[40] = 2; + assert!(RemoteRequest::from_frame(frame).unwrap_err().contains("invalid remote sync capability flag")); } #[test] From 842f4202e33b47a88581961927c71ac5b70ad4a8 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 4 Aug 2026 23:07:00 +0200 Subject: [PATCH 25/52] Linters --- src/daemon_ut.rs | 2 +- src/lib.rs | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index 367332f..07b0e0f 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -369,7 +369,7 @@ fn local_exec_request_still_builds_on_the_local_path() { env_mode: EnvMode::Paranoid, workspace: workspace.path().to_path_buf(), merged_profile: None, - has_proxy: false, + proxy_config: None, remote_broker: Arc::new(remote_broker(Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }))), }; let request = crate::vscomm::ExecRequest { cwd: "/workspace".into(), command: "true".into(), args: Vec::new(), env: Vec::new() }; diff --git a/src/lib.rs b/src/lib.rs index e4edf99..8706a2d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -6,8 +6,8 @@ pub mod daemon; pub mod guest_install; pub mod kata; pub mod logging; -pub mod netrelay; pub mod loopback; +pub mod netrelay; pub mod overlay; pub mod proxy; pub mod remote; From a453d81fa2223da18016b6ddb4f979c3f79b11d0 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 17:02:24 +0200 Subject: [PATCH 26/52] Add project data exclusion paths --- src/cfg.rs | 14 +++++++++++++- src/main.rs | 2 +- src/snapshot.rs | 4 ++++ 3 files changed, 18 insertions(+), 2 deletions(-) diff --git a/src/cfg.rs b/src/cfg.rs index b308273..907ed51 100644 --- a/src/cfg.rs +++ b/src/cfg.rs @@ -5,6 +5,7 @@ use std::path::{Path, PathBuf}; use serde::Deserialize; use crate::remote::{RemoteEnvironmentPolicy, RemoteTool}; +use crate::snapshot::SnapshotExclusionPolicy; use crate::vscomm::buildsys::{self, PassthroughMode}; pub const DEFAULT_SHARE_DIR: &str = "/usr/share/bunkerbox"; @@ -203,6 +204,8 @@ pub struct ProjectSection { #[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)] pub struct RemoteSection { + #[serde(default)] + pub exclude: Vec, #[serde(default)] pub environment: Vec, #[serde(default)] @@ -278,6 +281,7 @@ impl ProjectConfig { } } RemoteEnvironmentPolicy::from_names(self.project.remote.environment.clone())?; + SnapshotExclusionPolicy::from_patterns(self.project.remote.exclude.clone())?; let mut tools = std::collections::BTreeSet::new(); for tool in &self.project.remote.tools { RemoteTool::new(tool.name.clone())?; @@ -381,8 +385,16 @@ impl ProjectConfig { } } - if !self.project.remote.environment.is_empty() || !self.project.remote.tools.is_empty() { + if !self.project.remote.exclude.is_empty() || !self.project.remote.environment.is_empty() || !self.project.remote.tools.is_empty() { y.push_str(" remote:\n"); + y.push_str(" exclude:\n"); + if self.project.remote.exclude.is_empty() { + y.push_str(" []\n"); + } else { + for pattern in &self.project.remote.exclude { + y.push_str(&format!(" - {pattern}\n")); + } + } y.push_str(" environment:\n"); if self.project.remote.environment.is_empty() { y.push_str(" []\n"); diff --git a/src/main.rs b/src/main.rs index d32f37d..c602de8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -330,7 +330,7 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option) -> Result { + Self::from_patterns(config.effective_exclude(runtime_exclude).into_iter().chain(config.project.remote.exclude.iter().cloned())) + } + pub fn from_patterns(patterns: impl IntoIterator) -> Result { let mut policy = Self { basename_prunes: BTreeSet::new(), anchored_prunes: BTreeSet::new() }; for name in [".git", ".bunker", ".bunkerbox", ".env", ".envrc", ".ssh"] { From 4a56e235fa783bb9aaccb8fa282e6c0a7fbb89ab Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 17:02:31 +0200 Subject: [PATCH 27/52] Add data exclusion UT --- src/cfg_ut.rs | 42 ++++++++++++++++++++++++++ src/snapshot_ut.rs | 73 +++++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 114 insertions(+), 1 deletion(-) diff --git a/src/cfg_ut.rs b/src/cfg_ut.rs index d932fbf..6d4ddde 100644 --- a/src/cfg_ut.rs +++ b/src/cfg_ut.rs @@ -45,6 +45,48 @@ fn load_or_create_loads_existing_config() { assert_eq!(cfg.project.exclude, vec!["build/", "logs/"]); } +#[test] +fn load_or_create_parses_remote_exclusions() { + let root = TempDir::new().unwrap(); + write_project_conf(root.path(), "project:\n remote:\n exclude:\n - .tmp\n - docs/generated\n"); + + let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); + + assert_eq!(cfg.project.remote.exclude, vec![".tmp", "docs/generated"]); +} + +#[test] +fn remote_exclusions_default_to_empty() { + let root = TempDir::new().unwrap(); + write_project_conf(root.path(), "project:\n remote:\n exclude: []\n"); + + let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); + + assert!(cfg.project.remote.exclude.is_empty()); + assert!(ProjectConfig::default().project.remote.exclude.is_empty()); +} + +#[test] +fn load_or_create_rejects_invalid_remote_exclusions() { + for exclusion in ["/absolute", "foo/../bar"] { + let root = TempDir::new().unwrap(); + write_project_conf(root.path(), &format!("project:\n remote:\n exclude:\n - \"{exclusion}\"\n")); + + assert!(ProjectConfig::load_or_create(root.path()).is_err(), "{exclusion}"); + } +} + +#[test] +fn remote_exclusions_do_not_change_general_workspace_exclusions() { + let cfg = ProjectConfig { + project: ProjectSection { remote: RemoteSection { exclude: vec![".tmp".into()], ..Default::default() }, ..Default::default() }, + ..Default::default() + }; + + assert!(cfg.project.exclude.is_empty()); + assert!(!cfg.effective_exclude(None).iter().any(|pattern| pattern.trim_end_matches('/') == ".tmp")); +} + /// Invalid YAML in project.conf produces an error. #[test] fn load_or_create_invalid_yaml_is_error() { diff --git a/src/snapshot_ut.rs b/src/snapshot_ut.rs index 173774e..5c520ce 100644 --- a/src/snapshot_ut.rs +++ b/src/snapshot_ut.rs @@ -1,5 +1,5 @@ use super::*; -use crate::cfg::{ProjectConfig, ProjectSection}; +use crate::cfg::{ProjectConfig, ProjectSection, RemoteSection}; use crate::remote::WorkspaceSessionId; use std::fs; use std::os::unix::fs::{symlink, PermissionsExt}; @@ -24,6 +24,24 @@ fn build_at(source: &TempDir, store: &TempDir, limits: SnapshotLimits, patterns: builder(store, limits, patterns).build_root(source.path(), session(1)) } +fn remote_config(patterns: &[&str]) -> ProjectConfig { + ProjectConfig { + project: ProjectSection { + remote: RemoteSection { exclude: patterns.iter().map(|pattern| (*pattern).to_string()).collect(), ..Default::default() }, + ..Default::default() + }, + ..Default::default() + } +} + +fn remote_builder(store: &TempDir, limits: SnapshotLimits, config: &ProjectConfig) -> SnapshotBuilder { + SnapshotBuilder::new(SnapshotStore::new(store.path()), limits, SnapshotExclusionPolicy::from_remote_config(config, None).unwrap()) +} + +fn build_remote_at(source: &TempDir, store: &TempDir, limits: SnapshotLimits, config: &ProjectConfig) -> Result { + remote_builder(store, limits, config).build_root(source.path(), session(1)) +} + fn write_file(root: &Path, path: &str, contents: &[u8]) { let path = root.join(path); fs::create_dir_all(path.parent().unwrap()).unwrap(); @@ -104,6 +122,28 @@ fn config_and_runtime_exclusions_use_explicit_snapshot_semantics() { assert!(!policy.excludes("vendorized/file")); } +#[test] +fn remote_exclusions_use_basename_and_root_anchored_semantics() { + let config = remote_config(&[".tmp", "docs/generated"]); + let policy = SnapshotExclusionPolicy::from_remote_config(&config, None).unwrap(); + + assert!(policy.excludes(".tmp/electron")); + assert!(policy.excludes("nested/.tmp/electron")); + assert!(policy.excludes("docs/generated/file")); + assert!(!policy.excludes("nested/docs/generated/file")); + assert!(!policy.excludes("docs/generated-other/file")); +} + +#[test] +fn mandatory_snapshot_exclusions_remain_enforced_for_remote_config() { + let config = remote_config(&[".git", ".bunker", ".bunkerbox", ".ssh", ".env", ".envrc"]); + let policy = SnapshotExclusionPolicy::from_remote_config(&config, None).unwrap(); + + for path in [".git/config", "nested/.bunker/state", ".bunkerbox/control", ".ssh/key", ".env", "nested/.envrc"] { + assert!(policy.excludes(path), "{path}"); + } +} + #[test] fn malformed_exclusions_are_rejected() { for pattern in ["/absolute", "foo/../bar", "foo//bar", "foo/./bar", ""] { @@ -136,6 +176,37 @@ fn every_symlink_is_rejected_without_following_it() { } } +#[test] +fn remote_exclusions_skip_oversized_files_and_symlinks_before_validation() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), ".tmp/electron", b"oversized"); + let symlink_path = source.path().join(".venv-docs/bin/python"); + fs::create_dir_all(symlink_path.parent().unwrap()).unwrap(); + symlink("missing-python", &symlink_path).unwrap(); + + let config = remote_config(&[".tmp", ".venv-docs"]); + let snapshot = build_remote_at(&source, &store, SnapshotLimits { max_file_bytes: 4, ..SnapshotLimits::default() }, &config).unwrap(); + + assert!(snapshot.entries().is_empty()); +} + +#[test] +fn remote_exclusions_do_not_bypass_snapshot_validation_outside_excluded_trees() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "other/electron", b"oversized"); + let config = remote_config(&[".tmp", ".venv-docs"]); + assert!(build_remote_at(&source, &store, SnapshotLimits { max_file_bytes: 4, ..SnapshotLimits::default() }, &config).is_err()); + + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + let symlink_path = source.path().join("other/python"); + fs::create_dir_all(symlink_path.parent().unwrap()).unwrap(); + symlink("missing-python", &symlink_path).unwrap(); + assert!(build_remote_at(&source, &store, SnapshotLimits::default(), &config).is_err()); +} + #[test] fn special_files_and_hard_links_are_rejected() { let source = TempDir::new().unwrap(); From f3d7cb65e4a7a9b3f2dbb046921724d98a674def Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 19:48:34 +0200 Subject: [PATCH 28/52] Implement SSH worker transport --- src/daemon.rs | 48 +- src/lib.rs | 3 + src/main.rs | 20 +- src/remote.rs | 33 + src/remote_target.rs | 831 +++++++++++++++++++++++ src/snapshot.rs | 143 +++- src/ssh.rs | 704 +++++++++++++++++++ src/worker_protocol.rs | 1448 ++++++++++++++++++++++++++++++++++++++++ 8 files changed, 3208 insertions(+), 22 deletions(-) create mode 100644 src/remote_target.rs create mode 100644 src/ssh.rs create mode 100644 src/worker_protocol.rs diff --git a/src/daemon.rs b/src/daemon.rs index 322df6f..5645990 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -6,7 +6,9 @@ use crate::remote::{ RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteEnvironmentPolicy, RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteToolPolicy, }; +use crate::remote_target::SshTarget; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; +use crate::ssh::SshBackend; use crate::vscomm::{validate_exec_request, validate_process_path, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; use crate::workspace::WorkspaceCwd; use rand::Rng; @@ -98,24 +100,39 @@ pub struct RemoteDaemonConfig { allowed_tools: Vec, tool_policies: Option>, environment: Option, - tools: std::collections::BTreeMap, - target_environment: std::collections::BTreeMap, + backend: RemoteBackendSelection, resources: RemoteResourcePolicy, } +enum RemoteBackendSelection { + Loopback { tools: std::collections::BTreeMap, target_environment: std::collections::BTreeMap }, + Ssh { target: Box }, +} + impl RemoteDaemonConfig { - pub fn new(session: Arc, allowed_tools: Vec, tools: std::collections::BTreeMap) -> Self { + pub fn loopback(session: Arc, allowed_tools: Vec, tools: std::collections::BTreeMap) -> Self { Self { session, allowed_tools, tool_policies: None, environment: None, - tools, - target_environment: std::collections::BTreeMap::new(), + backend: RemoteBackendSelection::Loopback { tools, target_environment: std::collections::BTreeMap::new() }, resources: RemoteResourcePolicy::default(), } } + pub fn ssh(session: Arc, target: SshTarget) -> Result { + crate::ssh::SshLaunchSpec::from_target(&target)?; + Ok(Self { + session, + allowed_tools: Vec::new(), + tool_policies: None, + environment: None, + backend: RemoteBackendSelection::Ssh { target: Box::new(target) }, + resources: RemoteResourcePolicy::default(), + }) + } + pub fn with_policy(mut self, tools: Vec<(String, RemoteToolPolicy)>, environment: RemoteEnvironmentPolicy) -> Self { self.tool_policies = Some(tools); self.environment = Some(environment); @@ -123,7 +140,9 @@ impl RemoteDaemonConfig { } pub fn with_target_environment(mut self, environment: std::collections::BTreeMap) -> Self { - self.target_environment = environment; + if let RemoteBackendSelection::Loopback { target_environment, .. } = &mut self.backend { + *target_environment = environment; + } self } @@ -144,7 +163,7 @@ impl VsockDaemon { passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, remote: RemoteDaemonConfig, ) -> Result { - let RemoteDaemonConfig { session, allowed_tools, tool_policies, environment, tools, target_environment, resources } = remote; + let RemoteDaemonConfig { session, allowed_tools, tool_policies, environment, backend, resources } = remote; let remote_policy = match (tool_policies, environment) { (Some(tool_policies), Some(environment)) => { RemoteAuthorizationPolicy::from_policies(session.target(), session.session_id(), tool_policies, environment)? @@ -154,11 +173,16 @@ impl VsockDaemon { } .with_snapshot_authority(session.clone()); let remote_context = RemoteExecutionContext { target: session.target(), workspace_session_id: session.session_id() }; - let backend = LoopbackBackend::new(session, tools) - .with_target_environment(target_environment) - .with_timeout(resources.build_timeout) - .with_output_limit(resources.max_output_bytes); - let remote_components = RemoteComponents { context: remote_context, policy: remote_policy, backend: Arc::new(backend) }; + let backend: Arc = match backend { + RemoteBackendSelection::Loopback { tools, target_environment } => Arc::new( + LoopbackBackend::new(session, tools) + .with_target_environment(target_environment) + .with_timeout(resources.build_timeout) + .with_output_limit(resources.max_output_bytes), + ), + RemoteBackendSelection::Ssh { target } => Arc::new(SshBackend::new(session, *target)?), + }; + let remote_components = RemoteComponents { context: remote_context, policy: remote_policy, backend }; Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, remote_components) } diff --git a/src/lib.rs b/src/lib.rs index 8706a2d..bb36354 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -12,10 +12,13 @@ pub mod overlay; pub mod proxy; pub mod remote; pub mod remote_client; +pub mod remote_target; pub mod sandbox; pub mod snapshot; +pub mod ssh; pub mod tui; pub mod vscomm; +pub mod worker_protocol; pub mod workspace; pub mod wrap; diff --git a/src/main.rs b/src/main.rs index c602de8..6bcabce 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,6 @@ use bunkerbox::cfg::{ProjectConfig, RemoteToolSpec, WorkspaceMode}; use bunkerbox::remote::{RemoteEnvironmentPolicy, RemoteToolPolicy}; -use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, loopback, overlay, snapshot, tui, vscomm, workspace}; +use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, loopback, overlay, remote_target, snapshot, tui, vscomm, workspace}; use rand::RngCore; use std::ffi::OsString; use std::fs::File; @@ -190,6 +190,15 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option = config.allow.clone().unwrap_or_default().into_iter().chain(env.image.allow.clone().unwrap_or_default()).collect(); @@ -342,6 +351,13 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option daemon::RemoteDaemonConfig::loopback(session.clone(), Vec::new(), tools), + remote_target::BackendMode::Ssh => { + let target = remote_backend.target().cloned().ok_or_else(|| "SSH backend selection has no target".to_string())?; + daemon::RemoteDaemonConfig::ssh(session.clone(), target)? + } + }; let daemon = daemon::VsockDaemon::start_with_remote( passthrough, env_mode, @@ -349,7 +365,7 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option &'static str { + match self { + Self::Dns => "dns", + Self::Connect => "connect", + Self::Authentication => "authentication", + Self::HostIdentity => "host identity", + Self::WorkerUnavailable => "worker unavailable", + Self::WorkerVersion => "worker version", + Self::WorkerProtocol => "worker protocol", + Self::SnapshotTransfer => "snapshot transfer", + Self::Disconnect => "disconnect", + Self::Cleanup => "cleanup", + } + } +} + impl RemoteBackendError { pub fn event(&self) -> RemoteBackendEvent { match self { Self::Failed(message) | Self::Spawn(message) => RemoteBackendEvent::Error { message: message.clone() }, + Self::Transport { class, message } => RemoteBackendEvent::Error { message: format!("remote {} failure: {message}", class.as_str()) }, Self::Timeout => RemoteBackendEvent::Error { message: "remote backend timed out".to_string() }, Self::OutputLimit { limit } => RemoteBackendEvent::Error { message: format!("remote output exceeded limit of {limit} bytes") }, Self::Cancelled => RemoteBackendEvent::Cancelled, diff --git a/src/remote_target.rs b/src/remote_target.rs new file mode 100644 index 0000000..316a620 --- /dev/null +++ b/src/remote_target.rs @@ -0,0 +1,831 @@ +use serde::de::{self, MapAccess, Visitor}; +use serde::Deserialize; +use std::collections::BTreeMap; +use std::env; +use std::fmt; +use std::fs::{self, File}; +use std::marker::PhantomData; +use std::net::IpAddr; +use std::path::{Component, Path, PathBuf}; +use std::str::FromStr; +use std::time::Duration; + +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; + +pub const CONFIG_VERSION: u64 = 1; +pub const CONFIG_FILE_NAME: &str = "remote-targets.yaml"; +pub const CONFIG_DIRECTORY_NAME: &str = "bunkerbox"; + +const MAX_CONFIG_PATH_BYTES: usize = 4096; +const MAX_REMOTE_PATH_BYTES: usize = 4096; +const MAX_HOST_BYTES: usize = 253; +const MAX_USER_BYTES: usize = 64; +const MAX_NAME_BYTES: usize = 64; +const MAX_TOOL_PATH_BYTES: usize = 4096; +const MAX_ENV_NAME_BYTES: usize = 256; +const MAX_ENV_VALUE_BYTES: usize = 16 * 1024; + +/// Selects the host-side facility used for a project. +#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "lowercase")] +pub enum BackendMode { + Loopback, + Ssh, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ResourceLimits { + pub connect_timeout: Duration, + pub sync_timeout: Duration, + pub build_timeout: Duration, + pub max_output_bytes: u64, +} + +impl ResourceLimits { + pub fn connect_timeout(&self) -> Duration { + self.connect_timeout + } + + pub fn sync_timeout(&self) -> Duration { + self.sync_timeout + } + + pub fn build_timeout(&self) -> Duration { + self.build_timeout + } + + pub fn max_output(&self) -> u64 { + self.max_output_bytes + } + + pub fn max_output_bytes(&self) -> u64 { + self.max_output_bytes + } +} + +/// An SSH target after all configuration and local-file checks have passed. +/// +/// The identity and known-hosts files are intentionally represented only by +/// their paths. The files are never read by this module. +#[derive(Clone, PartialEq, Eq)] +pub struct SshTarget { + name: String, + host: String, + port: u16, + user: String, + identity_file: PathBuf, + known_hosts_file: PathBuf, + worker_path: String, + workspace_root: String, + tools: BTreeMap, + environment: BTreeMap, + resources: ResourceLimits, +} + +/// Alias emphasizing that an `SshTarget` can only be obtained after validation. +pub type ValidatedSshTarget = SshTarget; + +impl fmt::Debug for SshTarget { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SshTarget") + .field("name", &self.name) + .field("host", &self.host) + .field("port", &self.port) + .field("user", &self.user) + .field("identity_file", &self.identity_file) + .field("known_hosts_file", &self.known_hosts_file) + .field("worker_path", &self.worker_path) + .field("workspace_root", &self.workspace_root) + .field("tools", &self.tools) + .field("environment", &RedactedEnvironment(self.environment.len())) + .field("resources", &self.resources) + .finish() + } +} + +impl SshTarget { + pub fn name(&self) -> &str { + &self.name + } + + pub fn host(&self) -> &str { + &self.host + } + + pub fn port(&self) -> u16 { + self.port + } + + pub fn user(&self) -> &str { + &self.user + } + + pub fn identity_file(&self) -> &Path { + &self.identity_file + } + + pub fn known_hosts_file(&self) -> &Path { + &self.known_hosts_file + } + + pub fn worker_path(&self) -> &str { + &self.worker_path + } + + pub fn workspace_root(&self) -> &str { + &self.workspace_root + } + + pub fn tools(&self) -> &BTreeMap { + &self.tools + } + + pub fn environment(&self) -> &BTreeMap { + &self.environment + } + + pub fn resources(&self) -> ResourceLimits { + self.resources + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ProjectBinding { + pub backend: BackendMode, + pub target: Option, +} + +impl ProjectBinding { + pub fn backend(&self) -> BackendMode { + self.backend + } + + pub fn target(&self) -> Option<&str> { + self.target.as_deref() + } +} + +/// The result of resolving a canonical project path. +/// +/// Loopback resolutions always contain `None` for `target`, even if a +/// loopback project entry happens to contain an unused target name. +#[derive(Clone, PartialEq, Eq)] +pub struct ResolvedBackend { + pub project_root: PathBuf, + pub backend: BackendMode, + pub target: Option, +} + +impl fmt::Debug for ResolvedBackend { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("ResolvedBackend") + .field("project_root", &self.project_root) + .field("backend", &self.backend) + .field("target", &self.target) + .finish() + } +} + +impl ResolvedBackend { + pub fn backend(&self) -> BackendMode { + self.backend + } + + pub fn mode(&self) -> BackendMode { + self.backend + } + + pub fn project_root(&self) -> &Path { + &self.project_root + } + + pub fn target(&self) -> Option<&SshTarget> { + self.target.as_ref() + } +} + +/// Configuration loaded from the host's remote-targets file. +pub struct RemoteTargetConfig { + source_path: PathBuf, + targets: BTreeMap, + projects: BTreeMap, +} + +pub type RemoteConfig = RemoteTargetConfig; + +impl fmt::Debug for RemoteTargetConfig { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("RemoteTargetConfig") + .field("source_path", &self.source_path) + .field("targets", &self.targets) + .field("projects", &self.projects) + .finish() + } +} + +impl RemoteTargetConfig { + pub fn load_default() -> Result { + let helper = ConfigPathHelper::from_environment()?; + Self::load_default_with(&helper) + } + + pub fn load_default_with(helper: &ConfigPathHelper) -> Result { + Self::load_from(helper.config_path()?) + } + + pub fn load_default_with_path_helper(helper: &ConfigPathHelper) -> Result { + Self::load_default_with(helper) + } + + pub fn load_default_with_paths(xdg_config_home: Option, home: Option) -> Result { + Self::load_default_with(&ConfigPathHelper::new(xdg_config_home, home)) + } + + pub fn load_from(path: impl AsRef) -> Result { + let path = path.as_ref(); + let contents = fs::read_to_string(path).map_err(|error| format!("failed to read remote target config {}: {error}", path.display()))?; + let raw: RawConfig = + serde_yaml::from_str(&contents).map_err(|error| format!("failed to parse remote target config {}: {error}", path.display()))?; + Self::from_raw(raw, path.to_path_buf()) + } + + pub fn source_path(&self) -> &Path { + &self.source_path + } + + pub fn targets(&self) -> &BTreeMap { + &self.targets + } + + pub fn target(&self, name: &str) -> Option<&SshTarget> { + self.targets.get(name) + } + + pub fn ssh_target(&self, name: &str) -> Result<&SshTarget, String> { + self.targets.get(name).ok_or_else(|| format!("unknown SSH target '{name}'")) + } + + pub fn projects(&self) -> &BTreeMap { + &self.projects + } + + pub fn binding_for_project(&self, project: impl AsRef) -> Result<&ProjectBinding, String> { + let canonical = canonical_project_for_resolution(project.as_ref())?; + self.projects.get(&canonical).ok_or_else(|| format!("project has no remote backend binding: {}", canonical.display())) + } + + pub fn resolve_for_project(&self, project: impl AsRef) -> Result { + let canonical = canonical_project_for_resolution(project.as_ref())?; + let binding = self.projects.get(&canonical).ok_or_else(|| format!("project has no remote backend binding: {}", canonical.display()))?; + + match binding.backend { + BackendMode::Loopback => Ok(ResolvedBackend { project_root: canonical, backend: BackendMode::Loopback, target: None }), + BackendMode::Ssh => { + let target_name = binding.target.as_deref().ok_or_else(|| "SSH project binding is missing a target".to_string())?; + let target = self.targets.get(target_name).ok_or_else(|| format!("unknown SSH target '{target_name}'"))?; + Ok(ResolvedBackend { project_root: canonical, backend: BackendMode::Ssh, target: Some(target.clone()) }) + } + } + } + + fn from_raw(raw: RawConfig, source_path: PathBuf) -> Result { + if raw.version != CONFIG_VERSION { + return Err(format!("unsupported remote target config version {}; expected {}", raw.version, CONFIG_VERSION)); + } + + let mut targets = BTreeMap::new(); + for (name, target) in raw.targets.0 { + validate_name("target name", &name)?; + let validated = validate_target(name.clone(), target)?; + if targets.insert(name.clone(), validated).is_some() { + return Err(format!("duplicate target name '{name}'")); + } + } + + let mut projects = BTreeMap::new(); + for (project, binding) in raw.projects.0 { + let canonical = validate_project_binding_path(&project)?; + validate_project_binding(&binding)?; + if binding.backend == BackendMode::Ssh { + let Some(target_name) = binding.target.as_deref() else { + return Err("SSH project binding requires a target".to_string()); + }; + if !targets.contains_key(target_name) { + return Err(format!("unknown SSH target '{target_name}'")); + } + } + if projects.insert(canonical, binding.into_public()).is_some() { + return Err("duplicate project binding after canonicalization".to_string()); + } + } + + Ok(Self { source_path, targets, projects }) + } +} + +/// Inputs used to resolve the host-level default configuration path. +/// +/// Keeping environment lookup outside `config_path` makes default-path +/// behavior deterministic in callers and unit tests. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ConfigPathHelper { + xdg_config_home: Option, + home: Option, +} + +impl ConfigPathHelper { + pub fn new(xdg_config_home: Option, home: Option) -> Self { + Self { xdg_config_home: nonempty_path(xdg_config_home), home: nonempty_path(home) } + } + + pub fn from_paths(xdg_config_home: Option<&Path>, home: Option<&Path>) -> Self { + Self::new(xdg_config_home.map(Path::to_path_buf), home.map(Path::to_path_buf)) + } + + pub fn from_environment() -> Result { + Ok(Self::new(env::var_os("XDG_CONFIG_HOME").map(PathBuf::from), env::var_os("HOME").map(PathBuf::from))) + } + + pub fn xdg_config_home(&self) -> Option<&Path> { + self.xdg_config_home.as_deref() + } + + pub fn home(&self) -> Option<&Path> { + self.home.as_deref() + } + + pub fn config_path(&self) -> Result { + let base = match (&self.xdg_config_home, &self.home) { + (Some(xdg), _) => xdg.clone(), + (None, Some(home)) => home.join(".config"), + (None, None) => return Err("cannot resolve remote target config path: HOME is not set".to_string()), + }; + + validate_config_base(&base)?; + Ok(base.join(CONFIG_DIRECTORY_NAME).join(CONFIG_FILE_NAME)) + } + + pub fn default_path(&self) -> Result { + self.config_path() + } +} + +pub fn default_config_path() -> Result { + ConfigPathHelper::from_environment()?.config_path() +} + +pub fn default_config_path_with(xdg_config_home: Option<&Path>, home: Option<&Path>) -> Result { + ConfigPathHelper::from_paths(xdg_config_home, home).config_path() +} + +struct UniqueMap(BTreeMap); + +impl Default for UniqueMap { + fn default() -> Self { + Self(BTreeMap::new()) + } +} + +impl<'de, K, V> Deserialize<'de> for UniqueMap +where + K: Deserialize<'de> + Ord, + V: Deserialize<'de>, +{ + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + deserializer.deserialize_map(UniqueMapVisitor(PhantomData)) + } +} + +struct UniqueMapVisitor(PhantomData<(K, V)>); + +impl<'de, K, V> Visitor<'de> for UniqueMapVisitor +where + K: Deserialize<'de> + Ord, + V: Deserialize<'de>, +{ + type Value = UniqueMap; + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a YAML mapping with unique keys") + } + + fn visit_map(self, mut access: A) -> Result + where + A: MapAccess<'de>, + { + let mut entries = BTreeMap::new(); + while let Some(key) = access.next_key::()? { + if entries.contains_key(&key) { + return Err(de::Error::custom("duplicate YAML map key")); + } + let value = access.next_value::()?; + entries.insert(key, value); + } + Ok(UniqueMap(entries)) + } +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct RawConfig { + version: u64, + targets: UniqueMap, + projects: UniqueMap, +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct RawTarget { + transport: String, + host: String, + port: u64, + user: String, + #[serde(rename = "identity-file", alias = "identity_file")] + identity_file: String, + #[serde(rename = "known-hosts-file", alias = "known_hosts_file")] + known_hosts_file: String, + #[serde(rename = "worker-path", alias = "worker_path")] + worker_path: String, + #[serde(rename = "workspace-root", alias = "workspace_root")] + workspace_root: String, + #[serde(default)] + tools: UniqueMap, + #[serde(default)] + environment: UniqueMap, + resources: RawResources, +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct RawProjectBinding { + backend: BackendMode, + #[serde(default)] + target: Option, +} + +impl RawProjectBinding { + fn into_public(self) -> ProjectBinding { + ProjectBinding { backend: self.backend, target: self.target } + } +} + +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct RawResources { + #[serde( + rename = "connect-timeout-seconds", + alias = "connect-timeout", + alias = "connect_timeout_seconds", + alias = "connect_timeout", + alias = "connect" + )] + connect_timeout: RawQuantity, + #[serde(rename = "sync-timeout-seconds", alias = "sync-timeout", alias = "sync_timeout_seconds", alias = "sync_timeout", alias = "sync")] + sync_timeout: RawQuantity, + #[serde(rename = "build-timeout-seconds", alias = "build-timeout", alias = "build_timeout_seconds", alias = "build_timeout", alias = "build")] + build_timeout: RawQuantity, + #[serde(rename = "max-output-bytes", alias = "max-output", alias = "max_output_bytes", alias = "max_output")] + max_output: RawQuantity, +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum RawQuantity { + Integer(u64), + Text(String), +} + +fn validate_target(name: String, raw: RawTarget) -> Result { + if raw.transport != "ssh" { + return Err("remote target transport must be 'ssh'".to_string()); + } + + validate_hostname(&raw.host)?; + if raw.port == 0 || raw.port > u16::MAX as u64 { + return Err("remote target port must be between 1 and 65535".to_string()); + } + validate_username(&raw.user)?; + + let identity_file = validate_local_absolute_path("identity-file", &raw.identity_file)?; + validate_identity_file(&identity_file)?; + + let known_hosts_file = validate_local_absolute_path("known-hosts-file", &raw.known_hosts_file)?; + validate_known_hosts_file(&known_hosts_file)?; + + let worker_path = validate_remote_path("worker-path", &raw.worker_path, false)?; + let workspace_root = validate_remote_path("workspace-root", &raw.workspace_root, true)?; + + let mut tools = BTreeMap::new(); + for (identity, path) in raw.tools.0 { + validate_tool_identity(&identity)?; + let path = validate_remote_tool_path(&path)?; + if tools.insert(identity.clone(), path).is_some() { + return Err(format!("duplicate tool identity '{identity}'")); + } + } + + let mut environment = BTreeMap::new(); + for (name, value) in raw.environment.0 { + validate_environment_name(&name)?; + validate_environment_value(&value)?; + if environment.insert(name.clone(), value).is_some() { + return Err(format!("duplicate environment name '{name}'")); + } + } + + let resources = validate_resources(raw.resources)?; + + Ok(SshTarget { + name, + host: raw.host, + port: raw.port as u16, + user: raw.user, + identity_file, + known_hosts_file, + worker_path, + workspace_root, + tools, + environment, + resources, + }) +} + +fn validate_project_binding(binding: &RawProjectBinding) -> Result<(), String> { + if let Some(target) = &binding.target { + validate_name("project target name", target)?; + } + if binding.backend == BackendMode::Ssh && binding.target.is_none() { + return Err("SSH project binding requires a target".to_string()); + } + Ok(()) +} + +fn validate_resources(raw: RawResources) -> Result { + let connect_timeout = parse_duration("connect-timeout", raw.connect_timeout)?; + let sync_timeout = parse_duration("sync-timeout", raw.sync_timeout)?; + let build_timeout = parse_duration("build-timeout", raw.build_timeout)?; + let max_output_bytes = parse_size("max-output", raw.max_output)?; + Ok(ResourceLimits { connect_timeout, sync_timeout, build_timeout, max_output_bytes }) +} + +fn parse_duration(field: &str, quantity: RawQuantity) -> Result { + let (number, suffix) = quantity_parts(field, quantity)?; + let multiplier_nanos = match suffix.to_ascii_lowercase().as_str() { + "" | "s" => 1_000_000_000u64, + "ms" => 1_000_000, + "us" => 1_000, + "ns" => 1, + "m" => 60 * 1_000_000_000, + "h" => 60 * 60 * 1_000_000_000, + "d" => 24 * 60 * 60 * 1_000_000_000, + _ => return Err(format!("{field} has an unsupported duration unit")), + }; + let nanos = number.checked_mul(multiplier_nanos).ok_or_else(|| format!("{field} is too large"))?; + if nanos == 0 { + return Err(format!("{field} must be positive")); + } + let seconds = nanos / 1_000_000_000; + let subsecond_nanos = (nanos % 1_000_000_000) as u32; + Ok(Duration::new(seconds, subsecond_nanos)) +} + +fn parse_size(field: &str, quantity: RawQuantity) -> Result { + let (number, suffix) = quantity_parts(field, quantity)?; + let multiplier = match suffix.to_ascii_lowercase().as_str() { + "" | "b" => 1u64, + "k" | "kb" | "kib" => 1024, + "m" | "mb" | "mib" => 1024 * 1024, + "g" | "gb" | "gib" => 1024 * 1024 * 1024, + "t" | "tb" | "tib" => 1024 * 1024 * 1024 * 1024, + _ => return Err(format!("{field} has an unsupported size unit")), + }; + let bytes = number.checked_mul(multiplier).ok_or_else(|| format!("{field} is too large"))?; + if bytes == 0 { + return Err(format!("{field} must be positive")); + } + Ok(bytes) +} + +fn quantity_parts(field: &str, quantity: RawQuantity) -> Result<(u64, String), String> { + let text = match quantity { + RawQuantity::Integer(number) => return Ok((number, String::new())), + RawQuantity::Text(text) => text, + }; + let text = text.trim(); + let split = text.find(|character: char| !character.is_ascii_digit()).unwrap_or(text.len()); + if split == 0 { + return Err(format!("{field} must start with a positive integer")); + } + let number = text[..split].parse::().map_err(|_| format!("{field} is too large"))?; + Ok((number, text[split..].to_string())) +} + +fn validate_name(field: &str, value: &str) -> Result<(), String> { + if value.is_empty() || value.len() > MAX_NAME_BYTES || value == "." || value == ".." { + return Err(format!("{field} is invalid")); + } + let mut characters = value.bytes(); + let Some(first) = characters.next() else { + return Err(format!("{field} is invalid")); + }; + if !first.is_ascii_alphanumeric() { + return Err(format!("{field} is invalid")); + } + if !characters.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b'+')) { + return Err(format!("{field} is invalid")); + } + Ok(()) +} + +fn validate_hostname(host: &str) -> Result<(), String> { + if host.is_empty() || host.len() > MAX_HOST_BYTES || !host.is_ascii() || host.chars().any(char::is_whitespace) { + return Err("remote target host has unsafe hostname syntax".to_string()); + } + if IpAddr::from_str(host).is_ok() { + return Ok(()); + } + if host.ends_with('.') { + return Err("remote target host has unsafe hostname syntax".to_string()); + } + for label in host.split('.') { + if label.is_empty() || label.len() > 63 { + return Err("remote target host has unsafe hostname syntax".to_string()); + } + let bytes = label.as_bytes(); + if !bytes[0].is_ascii_alphanumeric() || !bytes[bytes.len() - 1].is_ascii_alphanumeric() { + return Err("remote target host has unsafe hostname syntax".to_string()); + } + if !bytes.iter().all(|byte| byte.is_ascii_alphanumeric() || *byte == b'-') { + return Err("remote target host has unsafe hostname syntax".to_string()); + } + } + Ok(()) +} + +fn validate_username(user: &str) -> Result<(), String> { + if user.is_empty() || user.len() > MAX_USER_BYTES || !user.is_ascii() { + return Err("remote target username is invalid".to_string()); + } + let bytes = user.as_bytes(); + if !bytes[0].is_ascii_alphabetic() && bytes[0] != b'_' { + return Err("remote target username is invalid".to_string()); + } + if !bytes[1..].iter().all(|byte| byte.is_ascii_alphanumeric() || matches!(*byte, b'_' | b'-' | b'.')) { + return Err("remote target username is invalid".to_string()); + } + Ok(()) +} + +fn validate_local_absolute_path(field: &str, value: &str) -> Result { + if value.is_empty() || value.len() > MAX_CONFIG_PATH_BYTES || value.chars().any(char::is_control) { + return Err(format!("{field} is invalid")); + } + let path = Path::new(value); + validate_absolute_no_parent_path(field, path)?; + Ok(path.to_path_buf()) +} + +fn validate_absolute_no_parent_path(field: &str, path: &Path) -> Result<(), String> { + if !path.is_absolute() { + return Err(format!("{field} must be absolute")); + } + if path.components().any(|component| matches!(component, Component::ParentDir | Component::CurDir)) { + return Err(format!("{field} must not contain '.' or '..' path components")); + } + Ok(()) +} + +fn validate_remote_path(field: &str, value: &str, allow_root: bool) -> Result { + if value.is_empty() || value.len() > MAX_REMOTE_PATH_BYTES || !value.is_ascii() || value.chars().any(char::is_whitespace) { + return Err(format!("{field} has invalid path syntax")); + } + if !value.starts_with('/') || value.contains("//") || (value != "/" && value.ends_with('/')) { + return Err(format!("{field} must be an absolute normalized path")); + } + let path = Path::new(value); + validate_absolute_no_parent_path(field, path)?; + if !allow_root && value == "/" { + return Err(format!("{field} must name a worker")); + } + if value.bytes().any(|byte| !byte.is_ascii_alphanumeric() && !matches!(byte, b'/' | b'.' | b'_' | b'-' | b'+' | b'@' | b'%' | b'~')) { + return Err(format!("{field} has invalid path syntax")); + } + Ok(value.to_string()) +} + +fn validate_remote_tool_path(value: &str) -> Result { + if value.len() > MAX_TOOL_PATH_BYTES { + return Err("tool path is too long".to_string()); + } + validate_remote_path("tool path", value, false) +} + +fn validate_tool_identity(identity: &str) -> Result<(), String> { + validate_name("tool identity", identity) +} + +fn validate_environment_name(name: &str) -> Result<(), String> { + if name.is_empty() || name.len() > MAX_ENV_NAME_BYTES { + return Err("environment name is invalid".to_string()); + } + let bytes = name.as_bytes(); + if !bytes[0].is_ascii_alphabetic() && bytes[0] != b'_' { + return Err("environment name is invalid".to_string()); + } + if !bytes[1..].iter().all(|byte| byte.is_ascii_alphanumeric() || *byte == b'_') { + return Err("environment name is invalid".to_string()); + } + Ok(()) +} + +fn validate_environment_value(value: &str) -> Result<(), String> { + if value.len() > MAX_ENV_VALUE_BYTES || value.chars().any(char::is_control) { + return Err("environment value is invalid".to_string()); + } + Ok(()) +} + +fn validate_project_binding_path(value: &str) -> Result { + let path = validate_local_absolute_path("project binding", value)?; + let canonical = fs::canonicalize(&path).map_err(|_| "project binding must refer to an existing canonical project path".to_string())?; + if canonical != path { + return Err("project binding must use the canonical project path".to_string()); + } + let metadata = fs::metadata(&canonical).map_err(|_| "project binding cannot be inspected".to_string())?; + if !metadata.is_dir() { + return Err("project binding must refer to a directory".to_string()); + } + Ok(canonical) +} + +fn canonical_project_for_resolution(path: &Path) -> Result { + validate_absolute_no_parent_path("project path", path)?; + let canonical = fs::canonicalize(path).map_err(|_| "project path cannot be canonicalized".to_string())?; + let metadata = fs::metadata(&canonical).map_err(|_| "project path cannot be inspected".to_string())?; + if !metadata.is_dir() { + return Err("project path must refer to a directory".to_string()); + } + Ok(canonical) +} + +fn validate_identity_file(path: &Path) -> Result<(), String> { + let metadata = regular_file_metadata(path, "identity-file")?; + #[cfg(unix)] + { + let mode = metadata.permissions().mode(); + if mode & 0o7777 != 0o400 && mode & 0o7777 != 0o600 { + return Err("identity-file has insecure permissions".to_string()); + } + } + File::open(path).map_err(|_| "identity-file is not readable".to_string())?; + Ok(()) +} + +fn validate_known_hosts_file(path: &Path) -> Result<(), String> { + let metadata = regular_file_metadata(path, "known-hosts-file")?; + #[cfg(unix)] + if metadata.permissions().mode() & 0o444 == 0 { + return Err("known-hosts-file is not readable".to_string()); + } + File::open(path).map_err(|_| "known-hosts-file is not readable".to_string())?; + Ok(()) +} + +fn regular_file_metadata(path: &Path, field: &str) -> Result { + let metadata = fs::symlink_metadata(path).map_err(|_| format!("{field} does not exist"))?; + if !metadata.file_type().is_file() { + return Err(format!("{field} must be a regular file")); + } + Ok(metadata) +} + +fn validate_config_base(path: &Path) -> Result<(), String> { + validate_absolute_no_parent_path("configuration directory", path)?; + if path.as_os_str().len() > MAX_CONFIG_PATH_BYTES { + return Err("configuration directory path is too long".to_string()); + } + Ok(()) +} + +fn nonempty_path(path: Option) -> Option { + path.filter(|path| !path.as_os_str().is_empty()) +} + +struct RedactedEnvironment(usize); + +impl fmt::Debug for RedactedEnvironment { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.debug_struct("RedactedEnvironment").field("entries", &self.0).finish() + } +} + +#[cfg(test)] +#[path = "remote_target_ut.rs"] +mod tests; diff --git a/src/snapshot.rs b/src/snapshot.rs index bc52e8a..ac4be5e 100644 --- a/src/snapshot.rs +++ b/src/snapshot.rs @@ -222,16 +222,21 @@ impl SnapshotStore { } pub fn resolve(&self, handle: &SnapshotHandle) -> Result { - let manifest_path = self.manifest_path(handle); - let metadata = fs::symlink_metadata(&manifest_path).map_err(|error| format!("snapshot is unavailable: {error}"))?; - if !metadata.file_type().is_file() { + let mut manifest = open_snapshot_manifest(self, handle)?; + let metadata = stat_fd(manifest.as_raw_fd())?; + if metadata.st_mode & libc::S_IFMT != libc::S_IFREG { return Err("snapshot manifest is not a regular file".to_string()); } - if metadata.len() > MAX_SNAPSHOT_MANIFEST_BYTES as u64 { + if metadata.st_size < 0 || metadata.st_size as u64 > MAX_SNAPSHOT_MANIFEST_BYTES as u64 { return Err("stored snapshot manifest exceeds limit".to_string()); } - let stored: StoredSnapshot = serde_json::from_slice(&fs::read(&manifest_path).map_err(|error| format!("read snapshot manifest: {error}"))?) - .map_err(|error| format!("decode snapshot manifest: {error}"))?; + let mut manifest_bytes = Vec::new(); + let mut bounded_manifest = (&mut manifest).take(MAX_SNAPSHOT_MANIFEST_BYTES as u64 + 1); + bounded_manifest.read_to_end(&mut manifest_bytes).map_err(|error| format!("read snapshot manifest: {error}"))?; + if manifest_bytes.len() > MAX_SNAPSHOT_MANIFEST_BYTES { + return Err("stored snapshot manifest exceeds limit".to_string()); + } + let stored: StoredSnapshot = serde_json::from_slice(&manifest_bytes).map_err(|error| format!("decode snapshot manifest: {error}"))?; if stored.session_id != handle.session_id.0 { return Err("snapshot session mismatch".to_string()); } @@ -243,6 +248,13 @@ impl SnapshotStore { Ok(WorkspaceSnapshot { handle: handle.clone(), entries, total_file_bytes }) } + #[allow(dead_code)] + pub(crate) fn resolve_export(&self, handle: &SnapshotHandle) -> Result { + let snapshot = self.resolve(handle)?; + let _files = open_snapshot_files(self, handle)?; + Ok(SnapshotExport { store: self.clone(), snapshot }) + } + pub fn remove(&self, handle: &SnapshotHandle) -> Result<(), String> { let path = self.snapshot_path(handle); match fs::symlink_metadata(&path) { @@ -288,8 +300,66 @@ impl SnapshotStore { Ok(MaterializedWorkspace { root: destination.to_path_buf() }) } - fn manifest_path(&self, handle: &SnapshotHandle) -> PathBuf { - self.snapshot_path(handle).join("manifest.json") + #[allow(dead_code)] + fn read_staged_file(&self, handle: &SnapshotHandle, entry: &SnapshotEntry, max_bytes: u64) -> Result, String> { + if entry.kind != SnapshotEntryKind::RegularFile { + return Err(format!("snapshot export entry is not a regular file: {}", entry.path.as_str())); + } + let expected_digest = *entry.content_digest.as_ref().ok_or_else(|| format!("regular file has no digest: {}", entry.path.as_str()))?; + let limit = max_bytes.min(MAX_SNAPSHOT_FILE_BYTES); + if entry.size > limit { + return Err(format!("snapshot export file exceeds read limit: {}", entry.path.as_str())); + } + + let source_root = open_snapshot_files(self, handle)?; + let source = open_relative_file(&source_root, entry.path.as_str(), libc::O_RDONLY)?; + let initial_stat = stat_fd(source.as_raw_fd())?; + if initial_stat.st_mode & libc::S_IFMT != libc::S_IFREG { + return Err(format!("snapshot export content is not a regular file: {}", entry.path.as_str())); + } + if initial_stat.st_nlink != 1 { + return Err(format!("snapshot export rejects linked content: {}", entry.path.as_str())); + } + if initial_stat.st_size < 0 || initial_stat.st_size as u64 != entry.size { + return Err(format!("snapshot export content size mismatch: {}", entry.path.as_str())); + } + if normalized_mode(initial_stat.st_mode) != entry.mode { + return Err(format!("snapshot export content mode mismatch: {}", entry.path.as_str())); + } + + let capacity = usize::try_from(entry.size).map_err(|_| format!("snapshot export file is too large to read: {}", entry.path.as_str()))?; + let mut contents = Vec::with_capacity(capacity); + let mut buffer = vec![0u8; SNAPSHOT_COPY_BUFFER_BYTES]; + let mut hasher = Sha256::new(); + let mut read_bytes = 0u64; + loop { + let count = (&source).read(&mut buffer).map_err(|error| format!("read snapshot export file {}: {error}", entry.path.as_str()))?; + if count == 0 { + break; + } + read_bytes = + read_bytes.checked_add(count as u64).ok_or_else(|| format!("snapshot export file size overflow: {}", entry.path.as_str()))?; + if read_bytes > entry.size || read_bytes > limit { + return Err(format!("snapshot export content exceeds manifest size: {}", entry.path.as_str())); + } + hasher.update(&buffer[..count]); + contents.extend_from_slice(&buffer[..count]); + } + + let final_stat = stat_fd(source.as_raw_fd())?; + if final_stat.st_dev != initial_stat.st_dev + || final_stat.st_ino != initial_stat.st_ino + || final_stat.st_mode & libc::S_IFMT != initial_stat.st_mode & libc::S_IFMT + || final_stat.st_size != initial_stat.st_size + || final_stat.st_nlink != initial_stat.st_nlink + || final_stat.st_mode & 0o777 != initial_stat.st_mode & 0o777 + { + return Err(format!("snapshot export content changed while reading: {}", entry.path.as_str())); + } + if read_bytes != entry.size || hasher.finalize().as_slice() != expected_digest { + return Err(format!("snapshot export content digest mismatch: {}", entry.path.as_str())); + } + Ok(contents) } fn snapshot_path(&self, handle: &SnapshotHandle) -> PathBuf { @@ -305,6 +375,40 @@ impl SnapshotStore { } } +#[allow(dead_code)] +pub(crate) struct SnapshotExport { + store: SnapshotStore, + snapshot: WorkspaceSnapshot, +} + +#[allow(dead_code)] +impl SnapshotExport { + pub(crate) fn entries(&self) -> &[SnapshotEntry] { + self.snapshot.entries() + } + + pub(crate) fn total_file_bytes(&self) -> u64 { + self.snapshot.total_file_bytes() + } + + pub(crate) fn read_file(&self, entry: &SnapshotEntry) -> Result, String> { + self.read_file_bounded(entry, MAX_SNAPSHOT_FILE_BYTES) + } + + pub(crate) fn read_file_bounded(&self, entry: &SnapshotEntry, max_bytes: u64) -> Result, String> { + let manifest_entry = self + .snapshot + .entries + .iter() + .find(|candidate| candidate.path == entry.path) + .ok_or_else(|| format!("snapshot export entry is not in the manifest: {}", entry.path.as_str()))?; + if manifest_entry != entry { + return Err(format!("snapshot export entry does not match the manifest: {}", entry.path.as_str())); + } + self.store.read_staged_file(self.snapshot.handle(), manifest_entry, max_bytes) + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct MaterializedWorkspace { root: PathBuf, @@ -629,6 +733,29 @@ fn open_directory(path: &Path) -> Result { .map_err(|error| format!("open snapshot workspace: {error}")) } +fn open_snapshot_directory(store: &SnapshotStore, handle: &SnapshotHandle) -> Result { + let store_root = open_directory(&store.root).map_err(|error| format!("open snapshot store: {error}"))?; + let session_name = hex(&handle.session_id.0); + let session = open_at(store_root.as_raw_fd(), OsStr::new(&session_name), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot session: {error}"))?; + let snapshot_name = hex(&handle.snapshot_id.0); + open_at(session.as_raw_fd(), OsStr::new(&snapshot_name), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot publication: {error}")) +} + +fn open_snapshot_manifest(store: &SnapshotStore, handle: &SnapshotHandle) -> Result { + let snapshot = open_snapshot_directory(store, handle)?; + open_at(snapshot.as_raw_fd(), OsStr::new("manifest.json"), libc::O_RDONLY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot manifest: {error}")) +} + +#[allow(dead_code)] +fn open_snapshot_files(store: &SnapshotStore, handle: &SnapshotHandle) -> Result { + let snapshot = open_snapshot_directory(store, handle)?; + open_at(snapshot.as_raw_fd(), OsStr::new("files"), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open snapshot content directory: {error}")) +} + fn open_child_directory(parent: RawFd, name: &OsStr, path: &str) -> Result { open_at(parent, name, libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) .map_err(|error| format!("open snapshot directory {path}: {error}")) diff --git a/src/ssh.rs b/src/ssh.rs new file mode 100644 index 0000000..5334973 --- /dev/null +++ b/src/ssh.rs @@ -0,0 +1,704 @@ +use crate::loopback::{RunRemoteSession, SnapshotExportClaim}; +use crate::remote::{ + AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFailureClass, RemoteFuture, RemoteOperation, + RemoteSnapshotId, +}; +use crate::remote_target::{ResourceLimits, SshTarget}; +use crate::snapshot::SnapshotEntryKind; +use crate::worker_protocol::{ + self, WorkerBuild, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRelativePath, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, + WorkerUploadId, MAX_WORKER_CHUNK_BYTES, MAX_WORKER_FILE_BYTES, +}; +use rand::RngCore; +use std::collections::BTreeMap; +use std::io; +use std::path::{Path, PathBuf}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite}; +use tokio::process::Command; +use tokio::task::JoinHandle; +use tokio::time::timeout; + +const SSH_PROGRAM: &str = "/usr/bin/ssh"; +const MAX_SSH_DIAGNOSTIC_BYTES: usize = 16 * 1024; +const CLEANUP_TIMEOUT: Duration = Duration::from_secs(1); +const PROCESS_REAP_TIMEOUT: Duration = Duration::from_secs(2); + +type WorkerReader = Box; +type WorkerWriter = Box; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SshLaunchSpec { + program: PathBuf, + args: Vec, + remote_command: String, +} + +impl SshLaunchSpec { + pub fn from_target(target: &SshTarget) -> Result { + if target.host().is_empty() || target.user().is_empty() || target.port() == 0 { + return Err("SSH target is not fully validated".to_string()); + } + + let remote_command = format!("exec {} --stdio --workspace-root {}", shell_quote(target.worker_path()), shell_quote(target.workspace_root())); + let connect_timeout = target.resources().connect_timeout().as_secs().max(1).to_string(); + let args = vec![ + "-F".to_string(), + "/dev/null".to_string(), + "-o".to_string(), + "BatchMode=yes".to_string(), + "-o".to_string(), + "StrictHostKeyChecking=yes".to_string(), + "-o".to_string(), + format!("UserKnownHostsFile={}", target.known_hosts_file().display()), + "-o".to_string(), + "GlobalKnownHostsFile=/dev/null".to_string(), + "-o".to_string(), + "IdentitiesOnly=yes".to_string(), + "-o".to_string(), + "IdentityAgent=none".to_string(), + "-o".to_string(), + "ForwardAgent=no".to_string(), + "-o".to_string(), + "ClearAllForwardings=yes".to_string(), + "-o".to_string(), + "RequestTTY=no".to_string(), + "-o".to_string(), + "PasswordAuthentication=no".to_string(), + "-o".to_string(), + "KbdInteractiveAuthentication=no".to_string(), + "-o".to_string(), + "ControlMaster=no".to_string(), + "-o".to_string(), + "EscapeChar=none".to_string(), + "-o".to_string(), + format!("ConnectTimeout={connect_timeout}"), + "-p".to_string(), + target.port().to_string(), + "-i".to_string(), + target.identity_file().display().to_string(), + "-l".to_string(), + target.user().to_string(), + "--".to_string(), + target.host().to_string(), + remote_command.clone(), + ]; + + Ok(Self { program: PathBuf::from(SSH_PROGRAM), args, remote_command }) + } + + pub fn program(&self) -> &Path { + &self.program + } + + pub fn args(&self) -> &[String] { + &self.args + } + + pub fn remote_command(&self) -> &str { + &self.remote_command + } +} + +pub trait SshProcess: Send { + fn take_stdin(&mut self) -> Option; + fn take_stdout(&mut self) -> Option; + fn take_stderr(&mut self) -> Option; + fn terminate_group(&mut self); + fn kill_group(&mut self); + fn wait<'a>(&'a mut self) -> RemoteFuture<'a, Result>; +} + +pub trait SshProcessFactory: Send + Sync { + fn spawn(&self, spec: &SshLaunchSpec) -> Result, String>; +} + +#[derive(Debug, Default)] +pub struct SystemSshProcessFactory; + +impl SshProcessFactory for SystemSshProcessFactory { + fn spawn(&self, spec: &SshLaunchSpec) -> Result, String> { + let mut command = Command::new(spec.program()); + command + .args(spec.args()) + .env_clear() + .stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .kill_on_drop(true); + unsafe { + command.pre_exec(|| { + if libc::setpgid(0, 0) != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) + }); + } + + let mut child = command.spawn().map_err(|error| format!("spawn SSH transport: {error}"))?; + let pid = child.id().map(|pid| pid as i32); + let stdin = child.stdin.take().ok_or_else(|| "SSH transport has no stdin".to_string())?; + let stdout = child.stdout.take().ok_or_else(|| "SSH transport has no stdout".to_string())?; + let stderr = child.stderr.take().ok_or_else(|| "SSH transport has no stderr".to_string())?; + Ok(Box::new(SystemSshProcess { + child, + pid, + stdin: Some(Box::new(stdin)), + stdout: Some(Box::new(stdout)), + stderr: Some(Box::new(stderr)), + active: true, + })) + } +} + +struct SystemSshProcess { + child: tokio::process::Child, + pid: Option, + stdin: Option, + stdout: Option, + stderr: Option, + active: bool, +} + +impl SshProcess for SystemSshProcess { + fn take_stdin(&mut self) -> Option { + self.stdin.take() + } + + fn take_stdout(&mut self) -> Option { + self.stdout.take() + } + + fn take_stderr(&mut self) -> Option { + self.stderr.take() + } + + fn terminate_group(&mut self) { + signal_group(self.pid, libc::SIGTERM); + } + + fn kill_group(&mut self) { + signal_group(self.pid, libc::SIGTERM); + signal_group(self.pid, libc::SIGKILL); + } + + fn wait<'a>(&'a mut self) -> RemoteFuture<'a, Result> { + Box::pin(async move { + let status = self.child.wait().await.map_err(|error| format!("wait for SSH transport: {error}"))?; + self.active = false; + Ok(status.code().unwrap_or(-1)) + }) + } +} + +impl Drop for SystemSshProcess { + fn drop(&mut self) { + if self.active { + self.kill_group(); + } + } +} + +pub struct SshBackend { + session: Arc, + target: SshTarget, + factory: Arc, + uploads: Arc>>, +} + +impl SshBackend { + pub fn new(session: Arc, target: SshTarget) -> Result { + let _ = SshLaunchSpec::from_target(&target)?; + Ok(Self { session, target, factory: Arc::new(SystemSshProcessFactory), uploads: Arc::new(Mutex::new(BTreeMap::new())) }) + } + + pub fn with_process_factory(mut self, factory: Arc) -> Self { + self.factory = factory; + self + } + + pub fn target(&self) -> &SshTarget { + &self.target + } + + pub fn pending_uploads(&self) -> usize { + self.uploads.lock().map(|uploads| uploads.len()).unwrap_or(0) + } +} + +impl RemoteBackend for SshBackend { + fn execute<'a>( + &'a self, request: AuthorizedRemoteRequest, events: tokio::sync::mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + let operation = request.request().operation().clone(); + let request_id = request.request_id(); + let session = self.session.clone(); + let target = self.target.clone(); + let factory = self.factory.clone(); + let uploads = self.uploads.clone(); + Box::pin(async move { + let backend = SshExecution { session, target, factory, uploads }; + match operation { + RemoteOperation::Sync(sync) => backend.execute_sync(request_id.0, sync.retain_capability(), events).await, + RemoteOperation::Build(build) => backend.execute_build(request_id.0, &build, events).await, + } + }) + } +} + +struct SshExecution { + session: Arc, + target: SshTarget, + factory: Arc, + uploads: Arc>>, +} + +impl SshExecution { + async fn execute_sync( + &self, request_id: [u8; 16], retain_capability: bool, events: tokio::sync::mpsc::Sender, + ) -> Result<(), RemoteBackendError> { + send_event(&events, RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }).await?; + let snapshot_id = self.session.sync_snapshot().map_err(RemoteBackendError::Failed)?; + let export = match self.session.claim_snapshot_for_export(snapshot_id) { + Ok(export) => export, + Err(error) => { + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(RemoteBackendError::Failed(error)); + } + }; + let upload_id = random_upload_id(); + let mut connection = match WorkerConnection::spawn(&self.factory, &self.target) { + Ok(connection) => connection, + Err(error) => { + drop(export); + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(error); + } + }; + let session_id = WorkerSessionId(self.session.session_id().0); + let operation = upload_and_finish(&mut connection, WorkerRequestId(request_id), session_id, upload_id, &export, &events, !retain_capability); + let result = match timeout(self.target.resources().sync_timeout(), operation).await { + Ok(result) => result, + Err(_) => { + connection.kill_and_reap().await; + Err(RemoteBackendError::Timeout) + } + }; + drop(export); + + if let Err(error) = result { + connection.kill_and_reap().await; + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(error); + } + + if retain_capability { + self.uploads + .lock() + .map_err(|_| RemoteBackendError::Failed("SSH upload registry lock poisoned".to_string()))? + .insert(snapshot_id, upload_id); + if let Err(error) = send_event(&events, RemoteBackendEvent::SyncCompleted { snapshot_id }).await { + self.remove_upload(snapshot_id); + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(error); + } + } else { + self.session.abort_snapshot_capability(snapshot_id).map_err(RemoteBackendError::Failed)?; + send_event(&events, RemoteBackendEvent::Completed { exit_code: 0 }).await?; + } + Ok(()) + } + + async fn execute_build( + &self, request_id: [u8; 16], build: &crate::remote::RemoteBuild, events: tokio::sync::mpsc::Sender, + ) -> Result<(), RemoteBackendError> { + let snapshot_id = build.snapshot_id(); + let upload_id = self + .uploads + .lock() + .map_err(|_| RemoteBackendError::Failed("SSH upload registry lock poisoned".to_string()))? + .remove(&snapshot_id) + .ok_or_else(|| RemoteBackendError::Transport { + class: RemoteFailureClass::SnapshotTransfer, + message: "remote snapshot upload is unavailable".to_string(), + })?; + let claim = self.session.claim_snapshot(snapshot_id).map_err(RemoteBackendError::Failed)?; + let executable = match self.target.tools().get(build.tool().as_str()) { + Some(executable) => executable.clone(), + None => { + drop(claim); + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::WorkerUnavailable, + message: format!("remote tool is not configured: {}", build.tool().as_str()), + }); + } + }; + let guest_env = build.env().to_vec(); + let target_env = self.target.environment().iter().map(|(key, value)| (key.clone(), value.clone())).collect::>(); + let worker_build = + WorkerBuild::new(build.tool().as_str(), executable, build.argv().to_vec(), build.cwd().as_str(), guest_env, target_env, upload_id) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, message: error.to_string() })?; + + let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; + let session_id = WorkerSessionId(self.session.session_id().0); + let operation = + build_and_finish(&mut connection, WorkerRequestId(request_id), session_id, worker_build, &events, upload_id, self.target.resources()); + let result = match timeout(self.target.resources().build_timeout(), operation).await { + Ok(result) => result, + Err(_) => { + connection.kill_and_reap().await; + Err(RemoteBackendError::Timeout) + } + }; + if let Err(error) = result { + let _ = timeout(CLEANUP_TIMEOUT, connection.cleanup(WorkerRequestId(request_id), session_id, upload_id)).await; + connection.kill_and_reap().await; + return Err(error); + } + drop(claim); + Ok(()) + } + + fn remove_upload(&self, snapshot_id: RemoteSnapshotId) { + if let Ok(mut uploads) = self.uploads.lock() { + uploads.remove(&snapshot_id); + } + } +} + +struct WorkerConnection { + process: Box, + writer: Option, + reader: WorkerReader, + stderr_task: Option>>, +} + +impl WorkerConnection { + fn spawn(factory: &Arc, target: &SshTarget) -> Result { + let spec = SshLaunchSpec::from_target(target) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::WorkerUnavailable, message: error })?; + let mut process = factory.spawn(&spec).map_err(classify_spawn_error)?; + let writer = process.take_stdin().ok_or_else(|| unavailable("SSH transport has no stdin"))?; + let reader = process.take_stdout().ok_or_else(|| unavailable("SSH transport has no stdout"))?; + let stderr = process.take_stderr().ok_or_else(|| unavailable("SSH transport has no stderr"))?; + let stderr_task = tokio::spawn(read_diagnostic(stderr)); + Ok(Self { process, writer: Some(writer), reader, stderr_task: Some(stderr_task) }) + } + + async fn handshake(&mut self, request_id: WorkerRequestId, session_id: WorkerSessionId) -> Result<(), RemoteBackendError> { + self.write(&WorkerMessage::hello(request_id, session_id, false)).await?; + let message = self.read().await?; + match message { + WorkerMessage::Hello { request_id: received_request, session_id: received_session, version, response } => { + if received_request != request_id || received_session != session_id { + return Err(worker_protocol("worker hello correlation mismatch")); + } + if version != worker_protocol::WORKER_PROTOCOL_VERSION { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::WorkerVersion, + message: format!("unsupported worker version: {version}"), + }); + } + if !response { + return Err(worker_protocol("worker hello was not a response")); + } + Ok(()) + } + WorkerMessage::Error { kind, message, .. } => Err(worker_error(WorkerOperation::Protocol, kind, message)), + _ => Err(worker_protocol("worker did not respond with Hello")), + } + } + + async fn write(&mut self, message: &WorkerMessage) -> Result<(), RemoteBackendError> { + let writer = self.writer.as_mut().ok_or_else(|| disconnected("SSH worker stdin is closed"))?; + worker_protocol::write_message(writer, message).await.map_err(worker_io_error) + } + + async fn read(&mut self) -> Result { + worker_protocol::read_message(&mut self.reader).await.map_err(worker_io_error) + } + + async fn cleanup( + &mut self, request_id: WorkerRequestId, session_id: WorkerSessionId, upload_id: WorkerUploadId, + ) -> Result<(), RemoteBackendError> { + self.write(&WorkerMessage::Cleanup { request_id, session_id, upload_token: upload_id }).await?; + match self.read().await? { + WorkerMessage::Completed { + request_id: received_request, + session_id: received_session, + operation: WorkerOperation::Cleanup, + exit_code, + } => { + if received_request != request_id || received_session != session_id { + return Err(worker_protocol("worker cleanup correlation mismatch")); + } + if exit_code != 0 { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::Cleanup, + message: format!("worker cleanup exited with status {exit_code}"), + }); + } + Ok(()) + } + WorkerMessage::Error { kind, message, .. } => Err(worker_error(WorkerOperation::Cleanup, kind, message)), + _ => Err(worker_protocol("unexpected worker cleanup response")), + } + } + + async fn finish(&mut self) -> Result<(), RemoteBackendError> { + self.writer.take(); + let status = self.process.wait().await.map_err(disconnected)?; + let diagnostic = self.stderr_task.take().map(|task| async move { task.await.unwrap_or_default() }); + let diagnostic = match diagnostic { + Some(future) => future.await, + None => Vec::new(), + }; + if status != 0 { + return Err(classify_exit(status, &diagnostic)); + } + Ok(()) + } + + async fn kill_and_reap(&mut self) { + self.process.kill_group(); + let _ = timeout(PROCESS_REAP_TIMEOUT, self.process.wait()).await; + if let Some(task) = self.stderr_task.take() { + task.abort(); + } + } +} + +impl Drop for WorkerConnection { + fn drop(&mut self) { + self.process.kill_group(); + if let Some(task) = self.stderr_task.take() { + task.abort(); + } + } +} + +async fn upload_and_finish( + connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, upload_id: WorkerUploadId, + export: &SnapshotExportClaim, events: &tokio::sync::mpsc::Sender, cleanup: bool, +) -> Result<(), RemoteBackendError> { + connection.handshake(request_id, session_id).await?; + let entries = export.entries().iter().map(worker_entry).collect::, _>>()?; + let total_bytes = export.total_file_bytes(); + connection.write(&WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries }).await?; + let mut completed_bytes = 0u64; + for entry in export.entries() { + if entry.kind() != SnapshotEntryKind::RegularFile { + continue; + } + let contents = export + .read_file_bounded(entry, MAX_WORKER_FILE_BYTES) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::SnapshotTransfer, message: error })?; + for (index, chunk) in contents.chunks(MAX_WORKER_CHUNK_BYTES).enumerate() { + let offset = u64::try_from(index).unwrap_or(u64::MAX).saturating_mul(MAX_WORKER_CHUNK_BYTES as u64); + connection + .write(&WorkerMessage::UploadFileChunk { + request_id, + session_id, + upload_id, + path: WorkerRelativePath::new(entry.path().as_str()).map_err(worker_io_error)?, + offset, + data: chunk.to_vec(), + }) + .await?; + completed_bytes = completed_bytes.saturating_add(chunk.len() as u64); + send_event(events, RemoteBackendEvent::SyncProgress { completed_bytes, total_bytes: Some(total_bytes) }).await?; + } + } + connection.write(&WorkerMessage::UploadComplete { request_id, session_id, upload_id }).await?; + loop { + match connection.read().await? { + WorkerMessage::SyncProgress { + request_id: received_request, + session_id: received_session, + upload_id: received_upload, + completed_bytes, + total_bytes, + } => { + if received_request != request_id || received_session != session_id || received_upload != upload_id { + return Err(worker_protocol("worker sync progress correlation mismatch")); + } + send_event(events, RemoteBackendEvent::SyncProgress { completed_bytes, total_bytes }).await?; + } + WorkerMessage::UploadComplete { request_id: received_request, session_id: received_session, upload_id: received_upload } => { + if received_request != request_id || received_session != session_id || received_upload != upload_id { + return Err(worker_protocol("worker upload completion correlation mismatch")); + } + break; + } + WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Upload, kind, message)), + _ => return Err(worker_protocol("unexpected worker upload response")), + } + } + if cleanup { + connection.cleanup(request_id, session_id, upload_id).await?; + } + connection.finish().await +} + +async fn build_and_finish( + connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild, + events: &tokio::sync::mpsc::Sender, upload_id: WorkerUploadId, resources: ResourceLimits, +) -> Result<(), RemoteBackendError> { + connection.handshake(request_id, session_id).await?; + connection.write(&WorkerMessage::build(request_id, session_id, build)).await?; + let mut output_bytes = 0u64; + let exit_code = loop { + match connection.read().await? { + WorkerMessage::Stdout { request_id: received_request, session_id: received_session, data } => { + check_correlation(received_request, received_session, request_id, session_id, "worker stdout")?; + output_bytes = output_bytes.saturating_add(data.len() as u64); + if output_bytes > resources.max_output_bytes() { + return Err(RemoteBackendError::OutputLimit { limit: resources.max_output_bytes() }); + } + send_event(events, RemoteBackendEvent::Stdout(data)).await?; + } + WorkerMessage::Stderr { request_id: received_request, session_id: received_session, data } => { + check_correlation(received_request, received_session, request_id, session_id, "worker stderr")?; + output_bytes = output_bytes.saturating_add(data.len() as u64); + if output_bytes > resources.max_output_bytes() { + return Err(RemoteBackendError::OutputLimit { limit: resources.max_output_bytes() }); + } + send_event(events, RemoteBackendEvent::Stderr(data)).await?; + } + WorkerMessage::Completed { request_id: received_request, session_id: received_session, operation: WorkerOperation::Build, exit_code } => { + check_correlation(received_request, received_session, request_id, session_id, "worker completion")?; + break exit_code; + } + WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Build, kind, message)), + _ => return Err(worker_protocol("unexpected worker build response")), + } + }; + connection.cleanup(request_id, session_id, upload_id).await?; + connection.finish().await?; + send_event(events, RemoteBackendEvent::Completed { exit_code }).await +} + +fn worker_entry(entry: &crate::snapshot::SnapshotEntry) -> Result { + match entry.kind() { + SnapshotEntryKind::Directory => WorkerUploadEntry::directory(entry.path().as_str(), entry.mode() as u32).map_err(worker_io_error), + SnapshotEntryKind::RegularFile => WorkerUploadEntry::file( + entry.path().as_str(), + entry.mode() as u32, + entry.size(), + *entry.content_digest().ok_or_else(|| worker_protocol("regular snapshot entry has no digest"))?, + ) + .map_err(worker_io_error), + } +} + +fn check_correlation( + received_request: WorkerRequestId, received_session: WorkerSessionId, request_id: WorkerRequestId, session_id: WorkerSessionId, label: &str, +) -> Result<(), RemoteBackendError> { + if received_request != request_id || received_session != session_id { + return Err(worker_protocol(format!("{label} correlation mismatch"))); + } + Ok(()) +} + +fn send_event<'a>( + events: &'a tokio::sync::mpsc::Sender, event: RemoteBackendEvent, +) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + Box::pin(async move { events.send(event).await.map_err(|_| RemoteBackendError::Cancelled) }) +} + +fn classify_spawn_error(message: String) -> RemoteBackendError { + RemoteBackendError::Transport { class: RemoteFailureClass::Connect, message } +} + +fn classify_exit(status: i32, diagnostic: &[u8]) -> RemoteBackendError { + let text = String::from_utf8_lossy(diagnostic).to_ascii_lowercase(); + let class = if text.contains("permission denied") || text.contains("authentication failed") { + RemoteFailureClass::Authentication + } else if text.contains("host key") || text.contains("offending") || text.contains("known_hosts") { + RemoteFailureClass::HostIdentity + } else if text.contains("could not resolve") || text.contains("name or service not known") { + RemoteFailureClass::Dns + } else if text.contains("connection refused") || text.contains("connection timed out") { + RemoteFailureClass::Connect + } else if text.contains("worker") || text.contains("no such file") { + RemoteFailureClass::WorkerUnavailable + } else { + RemoteFailureClass::Disconnect + }; + RemoteBackendError::Transport { class, message: format!("SSH transport exited with status {status}") } +} + +fn worker_io_error(error: worker_protocol::WorkerProtocolError) -> RemoteBackendError { + if error.is_invalid() { + let class = if error.contains("version") { RemoteFailureClass::WorkerVersion } else { RemoteFailureClass::WorkerProtocol }; + RemoteBackendError::Transport { class, message: error.to_string() } + } else { + disconnected(error.to_string()) + } +} + +fn worker_error(operation: WorkerOperation, kind: WorkerErrorKind, message: String) -> RemoteBackendError { + let class = match kind { + WorkerErrorKind::WorkerProtocol => RemoteFailureClass::WorkerProtocol, + WorkerErrorKind::Upload | WorkerErrorKind::Sync => RemoteFailureClass::SnapshotTransfer, + WorkerErrorKind::Cleanup => RemoteFailureClass::Cleanup, + WorkerErrorKind::Build => return RemoteBackendError::Failed(format!("remote worker build failure ({operation:?}): {message}")), + }; + RemoteBackendError::Transport { class, message } +} + +fn worker_protocol(message: impl Into) -> RemoteBackendError { + RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, message: message.into() } +} + +fn unavailable(message: impl Into) -> RemoteBackendError { + RemoteBackendError::Transport { class: RemoteFailureClass::WorkerUnavailable, message: message.into() } +} + +fn disconnected(message: impl Into) -> RemoteBackendError { + RemoteBackendError::Transport { class: RemoteFailureClass::Disconnect, message: message.into() } +} + +async fn read_diagnostic(mut reader: WorkerReader) -> Vec { + let mut result = Vec::new(); + let mut buffer = [0u8; 4096]; + loop { + match reader.read(&mut buffer).await { + Ok(0) | Err(_) => break, + Ok(count) => { + if result.len() < MAX_SSH_DIAGNOSTIC_BYTES { + let remaining = MAX_SSH_DIAGNOSTIC_BYTES - result.len(); + result.extend_from_slice(&buffer[..count.min(remaining)]); + } + } + } + } + result +} + +fn random_upload_id() -> WorkerUploadId { + loop { + let mut bytes = [0u8; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + if bytes != [0; 16] { + return WorkerUploadId(bytes); + } + } +} + +fn shell_quote(value: &str) -> String { + format!("'{}'", value.replace('\'', "'\\''")) +} + +fn signal_group(pid: Option, signal: libc::c_int) { + if let Some(pid) = pid.filter(|pid| *pid > 0) { + unsafe { + libc::kill(-pid, signal); + } + } +} + +#[cfg(test)] +#[path = "ssh_ut.rs"] +mod tests; diff --git a/src/worker_protocol.rs b/src/worker_protocol.rs new file mode 100644 index 0000000..6547448 --- /dev/null +++ b/src/worker_protocol.rs @@ -0,0 +1,1448 @@ +//! A bounded binary protocol for a host-side worker reached over SSH. +//! +//! This protocol is deliberately independent from `vscomm`. A frame is: +//! +//! ```text +//! magic[4] version[u16] kind[u8] payload_length[u32] payload[payload_length] +//! ``` +//! +//! All integer fields are little-endian. The payload length is checked before +//! allocating a payload buffer, and every length and count inside a payload is +//! checked before it can drive an allocation. + +use std::collections::BTreeSet; +use std::fmt; +use std::io; + +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; + +pub const WORKER_PROTOCOL_MAGIC: [u8; 4] = *b"BBWK"; +pub const WORKER_MAGIC: [u8; 4] = WORKER_PROTOCOL_MAGIC; +pub const WORKER_PROTOCOL_VERSION: u16 = 1; +pub const WORKER_VERSION: u16 = WORKER_PROTOCOL_VERSION; +pub const WORKER_FRAME_HEADER_LEN: usize = 4 + 2 + 1 + 4; +pub const WORKER_FRAME_HEADER_SIZE: usize = WORKER_FRAME_HEADER_LEN; +pub const WORKER_ID_LEN: usize = 16; +pub const WORKER_DIGEST_LEN: usize = 32; + +/// Maximum bytes in one worker payload. The declared length is rejected +/// before a buffer of this size is allocated. +pub const MAX_WORKER_FRAME_PAYLOAD: usize = 1024 * 1024; +pub const MAX_WORKER_PAYLOAD: usize = MAX_WORKER_FRAME_PAYLOAD; +pub const MAX_WORKER_FRAME_BYTES: usize = WORKER_FRAME_HEADER_LEN + MAX_WORKER_FRAME_PAYLOAD; +pub const MAX_WORKER_FRAME_LENGTH: usize = MAX_WORKER_FRAME_BYTES; +pub const MAX_WORKER_STRING_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_PATH_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ENTRY_PATH_BYTES: usize = MAX_WORKER_PATH_BYTES; +pub const MAX_WORKER_PATH_COMPONENT_BYTES: usize = 255; +pub const MAX_WORKER_PATH_DEPTH: usize = 64; +pub const MAX_WORKER_CWD_BYTES: usize = MAX_WORKER_PATH_BYTES; +pub const MAX_WORKER_TOOL_BYTES: usize = 256; +pub const MAX_WORKER_EXECUTABLE_PATH_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_EXECUTABLE_BYTES: usize = MAX_WORKER_EXECUTABLE_PATH_BYTES; +pub const MAX_WORKER_ARG_COUNT: usize = 256; +pub const MAX_WORKER_ARGUMENT_COUNT: usize = MAX_WORKER_ARG_COUNT; +pub const MAX_WORKER_ARG_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ARG_TOTAL_BYTES: usize = 256 * 1024; +pub const MAX_WORKER_ENV_COUNT: usize = 128; +pub const MAX_WORKER_ENVIRONMENT_COUNT: usize = MAX_WORKER_ENV_COUNT; +pub const MAX_WORKER_ENV_KEY_BYTES: usize = 256; +pub const MAX_WORKER_ENV_VALUE_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ENV_TOTAL_BYTES: usize = 256 * 1024; +pub const MAX_WORKER_UPLOAD_ENTRIES: usize = 4096; +pub const MAX_WORKER_ENTRY_COUNT: usize = MAX_WORKER_UPLOAD_ENTRIES; +pub const MAX_WORKER_UPLOAD_ENTRY_COUNT: usize = MAX_WORKER_UPLOAD_ENTRIES; +pub const MAX_WORKER_FILE_BYTES: u64 = 256 * 1024 * 1024; +pub const MAX_WORKER_TOTAL_UPLOAD_BYTES: u64 = 512 * 1024 * 1024; +pub const MAX_WORKER_MANIFEST_BYTES: usize = 512 * 1024; +pub const MAX_WORKER_CHUNK_BYTES: usize = 64 * 1024; +pub const MAX_WORKER_OUTPUT_BYTES: usize = 64 * 1024; +pub const MAX_WORKER_ERROR_BYTES: usize = 4 * 1024; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum WorkerProtocolError { + Io(String), + Invalid(String), +} + +pub type WorkerResult = Result; + +impl WorkerProtocolError { + pub fn is_invalid(&self) -> bool { + matches!(self, Self::Invalid(_)) + } + + pub fn contains(&self, needle: &str) -> bool { + match self { + Self::Io(message) | Self::Invalid(message) => message.contains(needle), + } + } +} + +impl fmt::Display for WorkerProtocolError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Io(message) => write!(formatter, "worker I/O error: {message}"), + Self::Invalid(message) => write!(formatter, "invalid worker protocol: {message}"), + } + } +} + +impl std::error::Error for WorkerProtocolError {} + +impl From for WorkerProtocolError { + fn from(error: io::Error) -> Self { + Self::Io(error.to_string()) + } +} + +macro_rules! worker_id { + ($name:ident) => { + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Default)] + pub struct $name(pub [u8; WORKER_ID_LEN]); + + impl $name { + pub const fn new(bytes: [u8; WORKER_ID_LEN]) -> Self { + Self(bytes) + } + + pub const fn as_bytes(&self) -> &[u8; WORKER_ID_LEN] { + &self.0 + } + + pub const fn into_bytes(self) -> [u8; WORKER_ID_LEN] { + self.0 + } + } + + impl From<[u8; WORKER_ID_LEN]> for $name { + fn from(bytes: [u8; WORKER_ID_LEN]) -> Self { + Self(bytes) + } + } + }; +} + +worker_id!(WorkerRequestId); +worker_id!(WorkerSessionId); +worker_id!(WorkerUploadId); + +pub type RequestId = WorkerRequestId; +pub type SessionId = WorkerSessionId; +pub type UploadId = WorkerUploadId; +pub type UploadToken = WorkerUploadId; +pub type WorkerUploadToken = WorkerUploadId; +pub type WorkerDigest = [u8; WORKER_DIGEST_LEN]; + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerFrameKind { + Hello = 1, + UploadBegin = 2, + UploadEntry = 3, + UploadFileChunk = 4, + UploadComplete = 5, + Build = 6, + Cleanup = 7, + SyncProgress = 8, + Stdout = 9, + Stderr = 10, + Completed = 11, + Error = 12, +} + +impl WorkerFrameKind { + pub const fn as_u8(self) -> u8 { + self as u8 + } + + pub const fn from_u8(value: u8) -> Option { + match value { + 1 => Some(Self::Hello), + 2 => Some(Self::UploadBegin), + 3 => Some(Self::UploadEntry), + 4 => Some(Self::UploadFileChunk), + 5 => Some(Self::UploadComplete), + 6 => Some(Self::Build), + 7 => Some(Self::Cleanup), + 8 => Some(Self::SyncProgress), + 9 => Some(Self::Stdout), + 10 => Some(Self::Stderr), + 11 => Some(Self::Completed), + 12 => Some(Self::Error), + _ => None, + } + } +} + +impl TryFrom for WorkerFrameKind { + type Error = WorkerProtocolError; + + fn try_from(value: u8) -> WorkerResult { + Self::from_u8(value).ok_or_else(|| invalid(format!("unknown worker frame kind: {value}"))) + } +} + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerOperation { + Protocol = 0, + Upload = 1, + Build = 2, + Cleanup = 3, + Sync = 4, +} + +impl WorkerOperation { + pub const fn as_u8(self) -> u8 { + self as u8 + } + + pub fn from_u8(value: u8) -> WorkerResult { + match value { + 0 => Ok(Self::Protocol), + 1 => Ok(Self::Upload), + 2 => Ok(Self::Build), + 3 => Ok(Self::Cleanup), + 4 => Ok(Self::Sync), + _ => Err(invalid(format!("unknown worker operation: {value}"))), + } + } +} + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerErrorKind { + WorkerProtocol = 1, + Upload = 2, + Build = 3, + Cleanup = 4, + Sync = 5, +} + +pub type WorkerErrorClass = WorkerErrorKind; + +impl WorkerErrorKind { + #[allow(non_upper_case_globals)] + pub const Protocol: Self = Self::WorkerProtocol; + + pub const fn as_u8(self) -> u8 { + self as u8 + } + + pub fn from_u8(value: u8) -> WorkerResult { + match value { + 1 => Ok(Self::WorkerProtocol), + 2 => Ok(Self::Upload), + 3 => Ok(Self::Build), + 4 => Ok(Self::Cleanup), + 5 => Ok(Self::Sync), + _ => Err(invalid(format!("unknown worker error kind: {value}"))), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerRelativePath(String); + +impl WorkerRelativePath { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_relative_path("worker relative path", &value, true)?; + Ok(Self(value)) + } + + fn for_entry(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_relative_path("worker entry path", &value, false)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn is_empty(&self) -> bool { + self.0.is_empty() + } +} + +impl AsRef for WorkerRelativePath { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerRelativePath { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerTool(String); + +impl WorkerTool { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_tool(&value)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl AsRef for WorkerTool { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerTool { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerExecutablePath(String); + +impl WorkerExecutablePath { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_executable_path(&value)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl AsRef for WorkerExecutablePath { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerExecutablePath { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerEntryKind { + Directory = 1, + File = 2, +} + +impl WorkerEntryKind { + pub fn from_u8(value: u8) -> WorkerResult { + match value { + 1 => Ok(Self::Directory), + 2 => Ok(Self::File), + _ => Err(invalid(format!("unknown worker entry kind: {value}"))), + } + } + + pub const fn as_u8(self) -> u8 { + self as u8 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkerUploadEntry { + pub path: WorkerRelativePath, + pub kind: WorkerEntryKind, + pub mode: u32, + pub size: u64, + pub digest: Option, +} + +impl WorkerUploadEntry { + pub fn new(path: impl Into, kind: WorkerEntryKind, mode: u32, size: u64, digest: Option) -> WorkerResult { + let entry = Self { path: WorkerRelativePath::for_entry(path)?, kind, mode, size, digest }; + entry.validate()?; + Ok(entry) + } + + pub fn directory(path: impl Into, mode: u32) -> WorkerResult { + Self::new(path, WorkerEntryKind::Directory, mode, 0, None) + } + + pub fn file(path: impl Into, mode: u32, size: u64, digest: WorkerDigest) -> WorkerResult { + Self::new(path, WorkerEntryKind::File, mode, size, Some(digest)) + } + + pub fn path(&self) -> &WorkerRelativePath { + &self.path + } + + pub fn kind(&self) -> WorkerEntryKind { + self.kind + } + + pub fn mode(&self) -> u32 { + self.mode + } + + pub fn size(&self) -> u64 { + self.size + } + + pub fn digest(&self) -> Option<&WorkerDigest> { + self.digest.as_ref() + } + + pub fn validate(&self) -> WorkerResult<()> { + validate_relative_path("worker entry path", self.path.as_str(), false)?; + if self.mode & !0o7777 != 0 { + return Err(invalid(format!("worker entry mode has unsupported bits: {:o}", self.mode))); + } + + match self.kind { + WorkerEntryKind::Directory => { + if self.size != 0 { + return Err(invalid("worker directory entry must have zero size")); + } + if self.digest.is_some() { + return Err(invalid("worker directory entry must not have a digest")); + } + } + WorkerEntryKind::File => { + if self.size > MAX_WORKER_FILE_BYTES { + return Err(invalid(format!("worker file exceeds maximum size {MAX_WORKER_FILE_BYTES}"))); + } + if self.digest.is_none() { + return Err(invalid("worker file entry is missing a digest")); + } + } + } + Ok(()) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkerBuild { + pub tool: WorkerTool, + pub trusted_executable: WorkerExecutablePath, + pub argv: Vec, + pub cwd: WorkerRelativePath, + pub guest_env: Vec<(String, String)>, + pub target_env: Vec<(String, String)>, + pub upload_token: WorkerUploadId, +} + +impl WorkerBuild { + pub fn new( + tool: impl Into, trusted_executable: impl Into, argv: Vec, cwd: impl Into, guest_env: Vec<(String, String)>, + target_env: Vec<(String, String)>, upload_token: WorkerUploadId, + ) -> WorkerResult { + let build = Self { + tool: WorkerTool::new(tool)?, + trusted_executable: WorkerExecutablePath::new(trusted_executable)?, + argv, + cwd: WorkerRelativePath::new(cwd)?, + guest_env, + target_env, + upload_token, + }; + build.validate()?; + Ok(build) + } + + pub fn validate(&self) -> WorkerResult<()> { + validate_tool(self.tool.as_str())?; + validate_executable_path(self.trusted_executable.as_str())?; + validate_relative_path("worker cwd", self.cwd.as_str(), true)?; + validate_argv(&self.argv)?; + validate_environment("worker guest environment", &self.guest_env)?; + validate_environment("worker target environment", &self.target_env)?; + Ok(()) + } + + pub fn tool(&self) -> &WorkerTool { + &self.tool + } + + pub fn trusted_executable(&self) -> &WorkerExecutablePath { + &self.trusted_executable + } + + pub fn trusted_executable_path(&self) -> &str { + self.trusted_executable.as_str() + } + + pub fn argv(&self) -> &[String] { + &self.argv + } + + pub fn cwd(&self) -> &WorkerRelativePath { + &self.cwd + } + + pub fn guest_env(&self) -> &[(String, String)] { + &self.guest_env + } + + pub fn target_env(&self) -> &[(String, String)] { + &self.target_env + } + + pub fn upload_token(&self) -> WorkerUploadId { + self.upload_token + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum WorkerMessage { + Hello { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + version: u16, + response: bool, + }, + UploadBegin { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entries: Vec, + }, + UploadEntry { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entry_index: u32, + entry: WorkerUploadEntry, + }, + UploadFileChunk { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + path: WorkerRelativePath, + offset: u64, + data: Vec, + }, + UploadComplete { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + }, + Build { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + build: WorkerBuild, + }, + Cleanup { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_token: WorkerUploadId, + }, + SyncProgress { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + completed_bytes: u64, + total_bytes: Option, + }, + Stdout { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + data: Vec, + }, + Stderr { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + data: Vec, + }, + Completed { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + operation: WorkerOperation, + exit_code: i32, + }, + Error { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + operation: WorkerOperation, + kind: WorkerErrorKind, + message: String, + }, +} + +impl WorkerMessage { + pub fn hello(request_id: WorkerRequestId, session_id: WorkerSessionId, response: bool) -> Self { + Self::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response } + } + + pub fn build(request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild) -> Self { + Self::Build { request_id, session_id, build } + } + + pub fn stdout(request_id: WorkerRequestId, session_id: WorkerSessionId, data: Vec) -> Self { + Self::Stdout { request_id, session_id, data } + } + + pub fn stderr(request_id: WorkerRequestId, session_id: WorkerSessionId, data: Vec) -> Self { + Self::Stderr { request_id, session_id, data } + } + + pub fn completed(request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, exit_code: i32) -> Self { + Self::Completed { request_id, session_id, operation, exit_code } + } + + pub fn error( + request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, kind: WorkerErrorKind, message: impl Into, + ) -> Self { + Self::Error { request_id, session_id, operation, kind, message: message.into() } + } + + pub const fn kind(&self) -> WorkerFrameKind { + match self { + Self::Hello { .. } => WorkerFrameKind::Hello, + Self::UploadBegin { .. } => WorkerFrameKind::UploadBegin, + Self::UploadEntry { .. } => WorkerFrameKind::UploadEntry, + Self::UploadFileChunk { .. } => WorkerFrameKind::UploadFileChunk, + Self::UploadComplete { .. } => WorkerFrameKind::UploadComplete, + Self::Build { .. } => WorkerFrameKind::Build, + Self::Cleanup { .. } => WorkerFrameKind::Cleanup, + Self::SyncProgress { .. } => WorkerFrameKind::SyncProgress, + Self::Stdout { .. } => WorkerFrameKind::Stdout, + Self::Stderr { .. } => WorkerFrameKind::Stderr, + Self::Completed { .. } => WorkerFrameKind::Completed, + Self::Error { .. } => WorkerFrameKind::Error, + } + } + + pub fn request_id(&self) -> WorkerRequestId { + match self { + Self::Hello { request_id, .. } + | Self::UploadBegin { request_id, .. } + | Self::UploadEntry { request_id, .. } + | Self::UploadFileChunk { request_id, .. } + | Self::UploadComplete { request_id, .. } + | Self::Build { request_id, .. } + | Self::Cleanup { request_id, .. } + | Self::SyncProgress { request_id, .. } + | Self::Stdout { request_id, .. } + | Self::Stderr { request_id, .. } + | Self::Completed { request_id, .. } + | Self::Error { request_id, .. } => *request_id, + } + } + + pub fn session_id(&self) -> WorkerSessionId { + match self { + Self::Hello { session_id, .. } + | Self::UploadBegin { session_id, .. } + | Self::UploadEntry { session_id, .. } + | Self::UploadFileChunk { session_id, .. } + | Self::UploadComplete { session_id, .. } + | Self::Build { session_id, .. } + | Self::Cleanup { session_id, .. } + | Self::SyncProgress { session_id, .. } + | Self::Stdout { session_id, .. } + | Self::Stderr { session_id, .. } + | Self::Completed { session_id, .. } + | Self::Error { session_id, .. } => *session_id, + } + } + + pub fn upload_id(&self) -> Option { + match self { + Self::UploadBegin { upload_id, .. } + | Self::UploadEntry { upload_id, .. } + | Self::UploadFileChunk { upload_id, .. } + | Self::UploadComplete { upload_id, .. } + | Self::SyncProgress { upload_id, .. } => Some(*upload_id), + Self::Build { build, .. } => Some(build.upload_token), + Self::Cleanup { upload_token, .. } => Some(*upload_token), + Self::Hello { .. } | Self::Stdout { .. } | Self::Stderr { .. } | Self::Completed { .. } | Self::Error { .. } => None, + } + } + + pub fn validate(&self) -> WorkerResult<()> { + match self { + Self::Hello { version, .. } => { + if *version != WORKER_PROTOCOL_VERSION { + return Err(invalid(format!("unsupported worker hello version: {version}"))); + } + } + Self::UploadBegin { entries, .. } => { + validate_upload_manifest(entries)?; + } + Self::UploadEntry { entry_index, entry, .. } => { + validate_entry_index(*entry_index)?; + entry.validate()?; + } + Self::UploadFileChunk { path, offset, data, .. } => validate_chunk(path, *offset, data)?, + Self::UploadComplete { .. } => {} + Self::Build { build, .. } => build.validate()?, + Self::Cleanup { .. } => {} + Self::SyncProgress { completed_bytes, total_bytes, .. } => validate_progress(*completed_bytes, *total_bytes)?, + Self::Stdout { data, .. } | Self::Stderr { data, .. } => { + if data.len() > MAX_WORKER_OUTPUT_BYTES { + return Err(invalid(format!("worker output exceeds maximum length {MAX_WORKER_OUTPUT_BYTES}"))); + } + } + Self::Completed { operation, .. } => { + if *operation == WorkerOperation::Protocol { + return Err(invalid("worker completion cannot use protocol operation")); + } + } + Self::Error { message, .. } => validate_error_message(message)?, + } + Ok(()) + } + + pub fn encode(&self) -> WorkerResult> { + self.validate()?; + let mut payload = WireWriter::new(); + encode_payload(self, &mut payload)?; + let payload = payload.finish()?; + + let declared = u32::try_from(payload.len()).map_err(|_| invalid("worker payload length does not fit in u32"))?; + let mut frame = Vec::with_capacity(WORKER_FRAME_HEADER_LEN + payload.len()); + frame.extend_from_slice(&WORKER_PROTOCOL_MAGIC); + frame.extend_from_slice(&WORKER_PROTOCOL_VERSION.to_le_bytes()); + frame.push(self.kind().as_u8()); + frame.extend_from_slice(&declared.to_le_bytes()); + frame.extend_from_slice(&payload); + Ok(frame) + } + + pub fn decode(frame: &[u8]) -> WorkerResult { + let (kind, payload) = split_frame(frame)?; + decode_payload(kind, payload) + } + + pub async fn read_async(reader: &mut R) -> WorkerResult { + let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; + reader.read_exact(&mut header).await.map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + let (kind, payload_len) = decode_header(&header)?; + let mut payload = vec![0u8; payload_len]; + reader.read_exact(&mut payload).await.map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + decode_payload(kind, &payload) + } + + pub async fn write_async(&self, writer: &mut W) -> WorkerResult<()> { + let frame = self.encode()?; + writer.write_all(&frame).await.map_err(WorkerProtocolError::from)?; + writer.flush().await.map_err(WorkerProtocolError::from) + } +} + +pub async fn read_worker_message(reader: &mut R) -> WorkerResult { + WorkerMessage::read_async(reader).await +} + +pub async fn write_worker_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { + message.write_async(writer).await +} + +pub async fn read_message(reader: &mut R) -> WorkerResult { + read_worker_message(reader).await +} + +pub async fn write_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { + write_worker_message(writer, message).await +} + +pub fn encode_worker_message(message: &WorkerMessage) -> WorkerResult> { + message.encode() +} + +pub fn decode_worker_message(frame: &[u8]) -> WorkerResult { + WorkerMessage::decode(frame) +} + +pub fn validate_worker_relative_path(value: &str) -> WorkerResult<()> { + validate_relative_path("worker relative path", value, false) +} + +pub fn validate_worker_tool(value: &str) -> WorkerResult<()> { + validate_tool(value) +} + +pub fn validate_worker_executable_path(value: &str) -> WorkerResult<()> { + validate_executable_path(value) +} + +pub fn validate_upload_manifest(entries: &[WorkerUploadEntry]) -> WorkerResult { + validate_count(entries.len(), MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; + let mut previous: Option<&str> = None; + let mut total_bytes = 0u64; + let mut manifest_bytes = 4usize; + for entry in entries { + entry.validate()?; + if let Some(previous) = previous { + if previous >= entry.path.as_str() { + return Err(invalid("worker upload manifest must be strictly sorted by path")); + } + } + previous = Some(entry.path.as_str()); + let encoded_entry_bytes = 4usize + .checked_add(entry.path.as_str().len()) + .and_then(|bytes| bytes.checked_add(1 + 4 + 8 + 1)) + .and_then(|bytes| bytes.checked_add(if entry.digest.is_some() { WORKER_DIGEST_LEN } else { 0 })) + .ok_or_else(|| invalid("worker manifest length overflow"))?; + manifest_bytes = manifest_bytes.checked_add(encoded_entry_bytes).ok_or_else(|| invalid("worker manifest length overflow"))?; + if manifest_bytes > MAX_WORKER_MANIFEST_BYTES { + return Err(invalid(format!("worker manifest exceeds maximum length {MAX_WORKER_MANIFEST_BYTES}"))); + } + total_bytes = total_bytes.checked_add(entry.size).ok_or_else(|| invalid("worker upload byte count overflow"))?; + if total_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { + return Err(invalid(format!("worker upload exceeds maximum size {MAX_WORKER_TOTAL_UPLOAD_BYTES}"))); + } + } + Ok(total_bytes) +} + +fn encode_payload(message: &WorkerMessage, writer: &mut WireWriter) -> WorkerResult<()> { + match message { + WorkerMessage::Hello { request_id, session_id, version, response } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.u16(*version)?; + writer.boolean(*response)?; + } + WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.count(entries.len(), MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; + for entry in entries { + encode_entry(writer, entry)?; + } + } + WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.u32(*entry_index)?; + encode_entry(writer, entry)?; + } + WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.string(path.as_str(), MAX_WORKER_PATH_BYTES, "worker chunk path")?; + writer.u64(*offset)?; + writer.blob(data, MAX_WORKER_CHUNK_BYTES, "worker file chunk")?; + } + WorkerMessage::UploadComplete { request_id, session_id, upload_id } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + } + WorkerMessage::Build { request_id, session_id, build } => { + encode_correlation(writer, *request_id, *session_id)?; + encode_build(writer, build)?; + } + WorkerMessage::Cleanup { request_id, session_id, upload_token } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_token.0)?; + } + WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes, total_bytes } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.u64(*completed_bytes)?; + match total_bytes { + Some(total_bytes) => { + writer.boolean(true)?; + writer.u64(*total_bytes)?; + } + None => writer.boolean(false)?, + } + } + WorkerMessage::Stdout { request_id, session_id, data } | WorkerMessage::Stderr { request_id, session_id, data } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.blob(data, MAX_WORKER_OUTPUT_BYTES, "worker output")?; + } + WorkerMessage::Completed { request_id, session_id, operation, exit_code } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.u8(operation.as_u8())?; + writer.i32(*exit_code)?; + } + WorkerMessage::Error { request_id, session_id, operation, kind, message } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.u8(operation.as_u8())?; + writer.u8(kind.as_u8())?; + writer.string(message, MAX_WORKER_ERROR_BYTES, "worker error")?; + } + } + Ok(()) +} + +fn decode_payload(kind: WorkerFrameKind, payload: &[u8]) -> WorkerResult { + let mut reader = WireReader::new(payload); + let message = match kind { + WorkerFrameKind::Hello => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let version = reader.u16()?; + let response = reader.boolean("worker hello response")?; + WorkerMessage::Hello { request_id, session_id, version, response } + } + WorkerFrameKind::UploadBegin => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let count = reader.count(MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; + let mut entries = Vec::with_capacity(count); + for _ in 0..count { + entries.push(decode_entry(&mut reader)?); + } + WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } + } + WorkerFrameKind::UploadEntry => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let entry_index = reader.u32()?; + let entry = decode_entry(&mut reader)?; + WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } + } + WorkerFrameKind::UploadFileChunk => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let path = WorkerRelativePath::for_entry(reader.string(MAX_WORKER_PATH_BYTES, "worker chunk path")?)?; + let offset = reader.u64()?; + let data = reader.blob(MAX_WORKER_CHUNK_BYTES, "worker file chunk")?; + WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } + } + WorkerFrameKind::UploadComplete => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + WorkerMessage::UploadComplete { request_id, session_id, upload_id } + } + WorkerFrameKind::Build => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + WorkerMessage::Build { request_id, session_id, build: decode_build(&mut reader)? } + } + WorkerFrameKind::Cleanup => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_token = WorkerUploadId(reader.array16()?); + WorkerMessage::Cleanup { request_id, session_id, upload_token } + } + WorkerFrameKind::SyncProgress => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let completed_bytes = reader.u64()?; + let total_bytes = if reader.boolean("worker progress total flag")? { Some(reader.u64()?) } else { None }; + WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes, total_bytes } + } + WorkerFrameKind::Stdout => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + WorkerMessage::Stdout { request_id, session_id, data: reader.blob(MAX_WORKER_OUTPUT_BYTES, "worker stdout")? } + } + WorkerFrameKind::Stderr => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + WorkerMessage::Stderr { request_id, session_id, data: reader.blob(MAX_WORKER_OUTPUT_BYTES, "worker stderr")? } + } + WorkerFrameKind::Completed => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let operation = WorkerOperation::from_u8(reader.u8()?)?; + let exit_code = reader.i32()?; + WorkerMessage::Completed { request_id, session_id, operation, exit_code } + } + WorkerFrameKind::Error => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let operation = WorkerOperation::from_u8(reader.u8()?)?; + let kind = WorkerErrorKind::from_u8(reader.u8()?)?; + let message = reader.string(MAX_WORKER_ERROR_BYTES, "worker error")?; + WorkerMessage::Error { request_id, session_id, operation, kind, message } + } + }; + reader.finish()?; + message.validate()?; + Ok(message) +} + +fn encode_build(writer: &mut WireWriter, build: &WorkerBuild) -> WorkerResult<()> { + build.validate()?; + writer.string(build.tool.as_str(), MAX_WORKER_TOOL_BYTES, "worker tool")?; + writer.string(build.trusted_executable.as_str(), MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?; + writer.count(build.argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; + for argument in &build.argv { + writer.string(argument, MAX_WORKER_ARG_BYTES, "worker argument")?; + } + writer.string(build.cwd.as_str(), MAX_WORKER_CWD_BYTES, "worker cwd")?; + encode_environment(writer, &build.guest_env, "worker guest environment")?; + encode_environment(writer, &build.target_env, "worker target environment")?; + writer.id(build.upload_token.0)?; + Ok(()) +} + +fn decode_build(reader: &mut WireReader<'_>) -> WorkerResult { + let tool = WorkerTool::new(reader.string(MAX_WORKER_TOOL_BYTES, "worker tool")?)?; + let trusted_executable = WorkerExecutablePath::new(reader.string(MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?)?; + let argument_count = reader.count(MAX_WORKER_ARG_COUNT, "worker argv")?; + let mut argv = Vec::with_capacity(argument_count); + for _ in 0..argument_count { + argv.push(reader.string(MAX_WORKER_ARG_BYTES, "worker argument")?); + } + let cwd = WorkerRelativePath::new(reader.string(MAX_WORKER_CWD_BYTES, "worker cwd")?)?; + let guest_env = decode_environment(reader, "worker guest environment")?; + let target_env = decode_environment(reader, "worker target environment")?; + let upload_token = WorkerUploadId(reader.array16()?); + let build = WorkerBuild { tool, trusted_executable, argv, cwd, guest_env, target_env, upload_token }; + build.validate()?; + Ok(build) +} + +fn encode_environment(writer: &mut WireWriter, environment: &[(String, String)], field: &str) -> WorkerResult<()> { + validate_environment(field, environment)?; + writer.count(environment.len(), MAX_WORKER_ENV_COUNT, field)?; + for (key, value) in environment { + writer.string(key, MAX_WORKER_ENV_KEY_BYTES, "worker environment key")?; + writer.string(value, MAX_WORKER_ENV_VALUE_BYTES, "worker environment value")?; + } + Ok(()) +} + +fn decode_environment(reader: &mut WireReader<'_>, field: &str) -> WorkerResult> { + let count = reader.count(MAX_WORKER_ENV_COUNT, field)?; + let mut environment = Vec::with_capacity(count); + for _ in 0..count { + let key = reader.string(MAX_WORKER_ENV_KEY_BYTES, "worker environment key")?; + let value = reader.string(MAX_WORKER_ENV_VALUE_BYTES, "worker environment value")?; + environment.push((key, value)); + } + validate_environment(field, &environment)?; + Ok(environment) +} + +fn encode_entry(writer: &mut WireWriter, entry: &WorkerUploadEntry) -> WorkerResult<()> { + entry.validate()?; + writer.string(entry.path.as_str(), MAX_WORKER_PATH_BYTES, "worker entry path")?; + writer.u8(entry.kind.as_u8())?; + writer.u32(entry.mode)?; + writer.u64(entry.size)?; + match entry.digest { + Some(digest) => { + writer.boolean(true)?; + writer.bytes(&digest)?; + } + None => writer.boolean(false)?, + } + Ok(()) +} + +fn decode_entry(reader: &mut WireReader<'_>) -> WorkerResult { + let path = reader.string(MAX_WORKER_PATH_BYTES, "worker entry path")?; + let kind = WorkerEntryKind::from_u8(reader.u8()?)?; + let mode = reader.u32()?; + let size = reader.u64()?; + let digest = match reader.boolean("worker entry digest flag")? { + true => Some(reader.array32()?), + false => None, + }; + WorkerUploadEntry::new(path, kind, mode, size, digest) +} + +fn encode_correlation(writer: &mut WireWriter, request_id: WorkerRequestId, session_id: WorkerSessionId) -> WorkerResult<()> { + writer.id(request_id.0)?; + writer.id(session_id.0) +} + +fn decode_correlation(reader: &mut WireReader<'_>) -> WorkerResult<(WorkerRequestId, WorkerSessionId)> { + Ok((WorkerRequestId(reader.array16()?), WorkerSessionId(reader.array16()?))) +} + +fn validate_entry_index(index: u32) -> WorkerResult<()> { + if usize::try_from(index).map_or(true, |index| index >= MAX_WORKER_UPLOAD_ENTRIES) { + return Err(invalid(format!("worker upload entry index exceeds maximum {MAX_WORKER_UPLOAD_ENTRIES}"))); + } + Ok(()) +} + +fn validate_chunk(path: &WorkerRelativePath, offset: u64, data: &[u8]) -> WorkerResult<()> { + validate_relative_path("worker chunk path", path.as_str(), false)?; + if data.is_empty() { + return Err(invalid("worker file chunk must not be empty")); + } + if data.len() > MAX_WORKER_CHUNK_BYTES { + return Err(invalid(format!("worker file chunk exceeds maximum length {MAX_WORKER_CHUNK_BYTES}"))); + } + let end = offset + .checked_add(u64::try_from(data.len()).map_err(|_| invalid("worker file chunk length does not fit in u64"))?) + .ok_or_else(|| invalid("worker file chunk offset overflow"))?; + if end > MAX_WORKER_FILE_BYTES { + return Err(invalid(format!("worker file chunk exceeds maximum file size {MAX_WORKER_FILE_BYTES}"))); + } + Ok(()) +} + +fn validate_progress(completed_bytes: u64, total_bytes: Option) -> WorkerResult<()> { + if completed_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { + return Err(invalid("worker progress exceeds the maximum upload size")); + } + if let Some(total_bytes) = total_bytes { + if total_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { + return Err(invalid("worker progress total exceeds the maximum upload size")); + } + if completed_bytes > total_bytes { + return Err(invalid("worker progress exceeds its total")); + } + } + Ok(()) +} + +fn validate_argv(argv: &[String]) -> WorkerResult<()> { + validate_count(argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; + let mut total_bytes = 0usize; + for argument in argv { + validate_text("worker argument", argument, MAX_WORKER_ARG_BYTES)?; + total_bytes = total_bytes.checked_add(argument.len()).ok_or_else(|| invalid("worker argv length overflow"))?; + if total_bytes > MAX_WORKER_ARG_TOTAL_BYTES { + return Err(invalid(format!("worker argv exceeds maximum length {MAX_WORKER_ARG_TOTAL_BYTES}"))); + } + } + Ok(()) +} + +fn validate_environment(field: &str, environment: &[(String, String)]) -> WorkerResult<()> { + validate_count(environment.len(), MAX_WORKER_ENV_COUNT, field)?; + let mut names = BTreeSet::new(); + let mut total_bytes = 0usize; + for (key, value) in environment { + validate_environment_key(key)?; + validate_environment_value(value)?; + total_bytes = total_bytes + .checked_add(key.len()) + .and_then(|bytes| bytes.checked_add(value.len())) + .ok_or_else(|| invalid(format!("{field} length overflow")))?; + if total_bytes > MAX_WORKER_ENV_TOTAL_BYTES { + return Err(invalid(format!("{field} exceeds maximum length {MAX_WORKER_ENV_TOTAL_BYTES}"))); + } + if !names.insert(key.as_str()) { + return Err(invalid(format!("duplicate worker environment key: {key}"))); + } + } + Ok(()) +} + +fn validate_environment_key(key: &str) -> WorkerResult<()> { + validate_text("worker environment key", key, MAX_WORKER_ENV_KEY_BYTES)?; + let mut bytes = key.bytes(); + let Some(first) = bytes.next() else { + return Err(invalid("worker environment key is empty")); + }; + if !(first == b'_' || first.is_ascii_alphabetic()) || !bytes.all(|byte| byte == b'_' || byte.is_ascii_alphanumeric()) { + return Err(invalid("worker environment key is not a valid variable name")); + } + Ok(()) +} + +fn validate_environment_value(value: &str) -> WorkerResult<()> { + validate_text("worker environment value", value, MAX_WORKER_ENV_VALUE_BYTES)?; + if value.chars().any(char::is_control) { + return Err(invalid("worker environment value contains control data")); + } + Ok(()) +} + +fn validate_error_message(message: &str) -> WorkerResult<()> { + validate_text("worker error", message, MAX_WORKER_ERROR_BYTES) +} + +fn validate_tool(value: &str) -> WorkerResult<()> { + validate_text("worker tool", value, MAX_WORKER_TOOL_BYTES)?; + if value.is_empty() + || value == "." + || value == ".." + || !value.bytes().all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b'+')) + { + return Err(invalid("worker tool must be a single safe executable identity")); + } + Ok(()) +} + +fn validate_executable_path(value: &str) -> WorkerResult<()> { + validate_text("worker executable path", value, MAX_WORKER_EXECUTABLE_PATH_BYTES)?; + if !value.starts_with('/') || value == "/" || value.starts_with("//") { + return Err(invalid("worker executable path must be a normalized absolute path")); + } + if value.chars().any(|character| character.is_whitespace() || character.is_control()) { + return Err(invalid("worker executable path must not contain whitespace or control data")); + } + if value.bytes().any(|byte| { + byte < 0x20 + || byte == 0x7f + || matches!(byte, b'\\' | b';' | b'|' | b'&' | b'$' | b'`' | b'<' | b'>' | b'\'' | b'"' | b'(' | b')' | b'[' | b']' | b'{' | b'}') + }) { + return Err(invalid("worker executable path contains unsafe identity data")); + } + for (depth, component) in value.split('/').skip(1).enumerate() { + if depth >= MAX_WORKER_PATH_DEPTH + || component.len() > MAX_WORKER_PATH_COMPONENT_BYTES + || component.is_empty() + || component == "." + || component == ".." + || component.contains(':') + { + return Err(invalid("worker executable path is not normalized")); + } + } + Ok(()) +} + +fn validate_relative_path(field: &str, value: &str, allow_empty: bool) -> WorkerResult<()> { + validate_text(field, value, MAX_WORKER_PATH_BYTES)?; + if value.is_empty() { + if allow_empty { + return Ok(()); + } + return Err(invalid(format!("{field} must not be empty"))); + } + if value.starts_with('/') || value.starts_with('\\') || value.contains('\\') || value.bytes().any(|byte| byte == b':') { + return Err(invalid(format!("{field} must be a normalized relative path"))); + } + for (depth, component) in value.split('/').enumerate() { + if depth >= MAX_WORKER_PATH_DEPTH + || component.len() > MAX_WORKER_PATH_COMPONENT_BYTES + || component.is_empty() + || component == "." + || component == ".." + || component.chars().any(char::is_control) + { + return Err(invalid(format!("{field} must be a normalized relative path"))); + } + } + Ok(()) +} + +fn validate_text(field: &str, value: &str, maximum: usize) -> WorkerResult<()> { + if value.len() > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + if value.as_bytes().contains(&0) { + return Err(invalid(format!("{field} contains a NUL byte"))); + } + Ok(()) +} + +fn validate_count(count: usize, maximum: usize, field: &str) -> WorkerResult<()> { + if count > maximum { + return Err(invalid(format!("{field} exceeds maximum count {maximum}"))); + } + Ok(()) +} + +fn invalid(message: impl Into) -> WorkerProtocolError { + WorkerProtocolError::Invalid(message.into()) +} + +fn split_frame(frame: &[u8]) -> WorkerResult<(WorkerFrameKind, &[u8])> { + if frame.len() < WORKER_FRAME_HEADER_LEN { + return Err(invalid("truncated worker frame header")); + } + let (kind, payload_len) = decode_header(&frame[..WORKER_FRAME_HEADER_LEN])?; + let expected = WORKER_FRAME_HEADER_LEN.checked_add(payload_len).ok_or_else(|| invalid("worker frame length overflow"))?; + if frame.len() < expected { + return Err(invalid("truncated worker frame payload")); + } + if frame.len() > expected { + return Err(invalid("extra bytes after worker frame")); + } + Ok((kind, &frame[WORKER_FRAME_HEADER_LEN..expected])) +} + +fn decode_header(header: &[u8]) -> WorkerResult<(WorkerFrameKind, usize)> { + if header.len() != WORKER_FRAME_HEADER_LEN { + return Err(invalid("invalid worker frame header length")); + } + if header[..4] != WORKER_PROTOCOL_MAGIC { + return Err(invalid("invalid worker frame magic")); + } + let version = u16::from_le_bytes([header[4], header[5]]); + if version != WORKER_PROTOCOL_VERSION { + return Err(invalid(format!("unsupported worker protocol version: {version}"))); + } + let kind = WorkerFrameKind::try_from(header[6])?; + let payload_len = u32::from_le_bytes([header[7], header[8], header[9], header[10]]) as usize; + if payload_len > MAX_WORKER_FRAME_PAYLOAD { + return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); + } + Ok((kind, payload_len)) +} + +struct WireWriter { + bytes: Vec, +} + +impl WireWriter { + fn new() -> Self { + Self { bytes: Vec::new() } + } + + fn finish(self) -> WorkerResult> { + if self.bytes.len() > MAX_WORKER_FRAME_PAYLOAD { + return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); + } + Ok(self.bytes) + } + + fn bytes(&mut self, value: &[u8]) -> WorkerResult<()> { + let new_length = self.bytes.len().checked_add(value.len()).ok_or_else(|| invalid("worker payload length overflow"))?; + if new_length > MAX_WORKER_FRAME_PAYLOAD { + return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); + } + self.bytes.extend_from_slice(value); + Ok(()) + } + + fn id(&mut self, value: [u8; WORKER_ID_LEN]) -> WorkerResult<()> { + self.bytes(&value) + } + + fn u8(&mut self, value: u8) -> WorkerResult<()> { + self.bytes(&[value]) + } + + fn boolean(&mut self, value: bool) -> WorkerResult<()> { + self.u8(u8::from(value)) + } + + fn u16(&mut self, value: u16) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn u32(&mut self, value: u32) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn u64(&mut self, value: u64) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn i32(&mut self, value: i32) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn count(&mut self, count: usize, maximum: usize, field: &str) -> WorkerResult<()> { + validate_count(count, maximum, field)?; + self.u32(u32::try_from(count).map_err(|_| invalid(format!("{field} count does not fit in u32")))?) + } + + fn string(&mut self, value: &str, maximum: usize, field: &str) -> WorkerResult<()> { + validate_text(field, value, maximum)?; + let length = u32::try_from(value.len()).map_err(|_| invalid(format!("{field} length does not fit in u32")))?; + self.u32(length)?; + self.bytes(value.as_bytes()) + } + + fn blob(&mut self, value: &[u8], maximum: usize, field: &str) -> WorkerResult<()> { + if value.len() > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + let length = u32::try_from(value.len()).map_err(|_| invalid(format!("{field} length does not fit in u32")))?; + self.u32(length)?; + self.bytes(value) + } +} + +struct WireReader<'a> { + bytes: &'a [u8], + offset: usize, +} + +impl<'a> WireReader<'a> { + fn new(bytes: &'a [u8]) -> Self { + Self { bytes, offset: 0 } + } + + fn take(&mut self, length: usize) -> WorkerResult<&'a [u8]> { + let end = self.offset.checked_add(length).ok_or_else(|| invalid("worker payload length overflow"))?; + if end > self.bytes.len() { + return Err(invalid("truncated worker payload")); + } + let value = &self.bytes[self.offset..end]; + self.offset = end; + Ok(value) + } + + fn u8(&mut self) -> WorkerResult { + Ok(self.take(1)?[0]) + } + + fn boolean(&mut self, field: &str) -> WorkerResult { + match self.u8()? { + 0 => Ok(false), + 1 => Ok(true), + value => Err(invalid(format!("invalid {field} flag: {value}"))), + } + } + + fn u16(&mut self) -> WorkerResult { + let value = self.take(2)?; + Ok(u16::from_le_bytes([value[0], value[1]])) + } + + fn u32(&mut self) -> WorkerResult { + let value = self.take(4)?; + Ok(u32::from_le_bytes([value[0], value[1], value[2], value[3]])) + } + + fn u64(&mut self) -> WorkerResult { + let value = self.take(8)?; + Ok(u64::from_le_bytes(value.try_into().map_err(|_| invalid("invalid worker integer"))?)) + } + + fn i32(&mut self) -> WorkerResult { + let value = self.take(4)?; + Ok(i32::from_le_bytes(value.try_into().map_err(|_| invalid("invalid worker integer"))?)) + } + + fn array16(&mut self) -> WorkerResult<[u8; WORKER_ID_LEN]> { + self.take(WORKER_ID_LEN)?.try_into().map_err(|_| invalid("invalid worker identifier")) + } + + fn array32(&mut self) -> WorkerResult<[u8; WORKER_DIGEST_LEN]> { + self.take(WORKER_DIGEST_LEN)?.try_into().map_err(|_| invalid("invalid worker digest")) + } + + fn count(&mut self, maximum: usize, field: &str) -> WorkerResult { + let count = self.u32()? as usize; + validate_count(count, maximum, field)?; + Ok(count) + } + + fn string(&mut self, maximum: usize, field: &str) -> WorkerResult { + let length = self.u32()? as usize; + if length > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + let value = std::str::from_utf8(self.take(length)?).map_err(|_| invalid(format!("{field} is not valid UTF-8")))?; + validate_text(field, value, maximum)?; + Ok(value.to_owned()) + } + + fn blob(&mut self, maximum: usize, field: &str) -> WorkerResult> { + let length = self.u32()? as usize; + if length > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + Ok(self.take(length)?.to_vec()) + } + + fn finish(self) -> WorkerResult<()> { + if self.offset != self.bytes.len() { + return Err(invalid("extra bytes in worker payload")); + } + Ok(()) + } +} + +#[cfg(test)] +#[path = "worker_protocol_ut.rs"] +mod worker_protocol_tests; From ee6a06ed444f565db70fff5526636e66dcd30f88 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 19:48:45 +0200 Subject: [PATCH 29/52] Add SSH worker transport unit tests --- src/loopback.rs | 79 ++++++++- src/loopback_ut.rs | 72 ++++++++ src/remote_target_ut.rs | 321 ++++++++++++++++++++++++++++++++++++ src/snapshot_ut.rs | 28 ++++ src/ssh_ut.rs | 336 ++++++++++++++++++++++++++++++++++++++ src/worker_protocol_ut.rs | 264 ++++++++++++++++++++++++++++++ 6 files changed, 1097 insertions(+), 3 deletions(-) create mode 100644 src/remote_target_ut.rs create mode 100644 src/ssh_ut.rs create mode 100644 src/worker_protocol_ut.rs diff --git a/src/loopback.rs b/src/loopback.rs index 8fd9fab..77c8730 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -2,7 +2,7 @@ use crate::remote::{ AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteResourcePolicy, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, WorkspaceSessionId, }; -use crate::snapshot::{SnapshotBuilder, SnapshotHandle, SnapshotStore}; +use crate::snapshot::{SnapshotBuilder, SnapshotEntry, SnapshotExport, SnapshotHandle, SnapshotStore}; use rand::RngCore; use std::collections::{BTreeMap, HashMap}; use std::fs; @@ -80,6 +80,11 @@ impl RunRemoteSession { self.snapshot_store.clone() } + #[cfg(test)] + pub(crate) fn snapshot_capability_count(&self) -> usize { + self.snapshot_capabilities.lock().map(|registry| registry.capabilities.len()).unwrap_or(0) + } + pub fn sync_snapshot(&self) -> Result { self.sync_snapshot_for_request(true)?.ok_or_else(|| "retained snapshot capability was not created".to_string()) } @@ -133,12 +138,44 @@ impl RunRemoteSession { } } - fn claim_snapshot(self: &Arc, snapshot_id: RemoteSnapshotId) -> Result { + #[allow(dead_code)] + pub(crate) fn abort_snapshot_capability(&self, snapshot_id: RemoteSnapshotId) -> Result<(), String> { + let handle = self + .snapshot_capabilities + .lock() + .map_err(|_| "remote snapshot registry lock poisoned".to_string())? + .capabilities + .remove(&snapshot_id) + .ok_or_else(|| "remote snapshot capability is unavailable".to_string())?; + self.release_snapshot(&handle); + Ok(()) + } + + pub(crate) fn claim_snapshot(self: &Arc, snapshot_id: RemoteSnapshotId) -> Result { let mut registry = self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; let handle = registry.capabilities.remove(&snapshot_id).ok_or_else(|| "remote snapshot capability is unavailable".to_string())?; Ok(SnapshotClaim { session: self.clone(), handle }) } + #[allow(dead_code)] + pub(crate) fn claim_snapshot_for_export(self: &Arc, snapshot_id: RemoteSnapshotId) -> Result { + let handle = { + let mut registry = self.snapshot_capabilities.lock().map_err(|_| "remote snapshot registry lock poisoned".to_string())?; + let handle = registry.capabilities.get(&snapshot_id).cloned().ok_or_else(|| "remote snapshot capability is unavailable".to_string())?; + let references = registry.references.entry(handle.clone()).or_insert(0); + *references = references.checked_add(1).ok_or_else(|| "remote snapshot capability reference count overflow".to_string())?; + handle + }; + let snapshot = match self.snapshot_store.resolve_export(&handle) { + Ok(snapshot) => snapshot, + Err(error) => { + self.release_snapshot(&handle); + return Err(error); + } + }; + Ok(SnapshotExportClaim { session: self.clone(), handle, snapshot }) + } + fn release_snapshot(&self, handle: &SnapshotHandle) { let should_remove = self .snapshot_capabilities @@ -182,7 +219,7 @@ struct SnapshotCapabilityRegistry { references: HashMap, } -struct SnapshotClaim { +pub(crate) struct SnapshotClaim { session: Arc, handle: SnapshotHandle, } @@ -207,6 +244,42 @@ impl Drop for SnapshotClaim { } } +#[allow(dead_code)] +pub(crate) struct SnapshotExportClaim { + session: Arc, + handle: SnapshotHandle, + snapshot: SnapshotExport, +} + +#[allow(dead_code)] +impl SnapshotExportClaim { + pub(crate) fn handle(&self) -> &SnapshotHandle { + &self.handle + } + + pub(crate) fn entries(&self) -> &[SnapshotEntry] { + self.snapshot.entries() + } + + pub(crate) fn total_file_bytes(&self) -> u64 { + self.snapshot.total_file_bytes() + } + + pub(crate) fn read_file(&self, entry: &SnapshotEntry) -> Result, String> { + self.snapshot.read_file(entry) + } + + pub(crate) fn read_file_bounded(&self, entry: &SnapshotEntry, max_bytes: u64) -> Result, String> { + self.snapshot.read_file_bounded(entry, max_bytes) + } +} + +impl Drop for SnapshotExportClaim { + fn drop(&mut self) { + self.session.release_snapshot(&self.handle); + } +} + impl RemoteSnapshotAuthority for RunRemoteSession { fn snapshot_available(&self, session: WorkspaceSessionId, snapshot_id: RemoteSnapshotId) -> bool { if session != self.session_id { diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index cf280ab..60f8fc2 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -163,6 +163,78 @@ fn outstanding_snapshot_capabilities_are_bounded_and_consumption_frees_a_slot() drop(temp); } +#[test] +fn export_claim_uses_the_exact_capability_and_releases_its_snapshot() { + let (_temp, session, _target, _session_id) = fixture(); + fs::write(session.workspace_root().join("src/input.txt"), b"A\n").unwrap(); + let snapshot_a = session.sync_snapshot().unwrap(); + fs::write(session.workspace_root().join("src/input.txt"), b"B\n").unwrap(); + let snapshot_b = session.sync_snapshot().unwrap(); + + let claim_a = session.claim_snapshot_for_export(snapshot_a).unwrap(); + let entry_a = claim_a.entries().iter().find(|entry| entry.path().as_str() == "src/input.txt").unwrap(); + assert_eq!(claim_a.handle().session_id(), session.session_id()); + assert_eq!(claim_a.read_file(entry_a).unwrap(), b"A\n"); + assert!(claim_a.read_file_bounded(entry_a, 1).is_err()); + assert_eq!(claim_a.total_file_bytes(), 2); + assert!(session.snapshot_capabilities.lock().unwrap().capabilities.contains_key(&snapshot_a)); + assert!(session.snapshot_capabilities.lock().unwrap().capabilities.contains_key(&snapshot_b)); + assert_eq!(published_snapshot_count(&session), 2); + + drop(claim_a); + session.abort_snapshot_capability(snapshot_a).unwrap(); + assert_eq!(published_snapshot_count(&session), 1); + + let claim_b = session.claim_snapshot_for_export(snapshot_b).unwrap(); + let entry_b = claim_b.entries().iter().find(|entry| entry.path().as_str() == "src/input.txt").unwrap(); + assert_eq!(claim_b.read_file(entry_b).unwrap(), b"B\n"); + drop(claim_b); + session.abort_snapshot_capability(snapshot_b).unwrap(); + assert_eq!(published_snapshot_count(&session), 0); +} + +#[test] +fn export_claim_rejects_missing_replay_and_cross_session_capabilities() { + let (_temp_a, session_a, _target_a, _session_id_a) = fixture(); + let snapshot_id = session_a.sync_snapshot().unwrap(); + let (_temp_b, session_b, _target_b, _session_id_b) = fixture(); + + assert!(session_a.claim_snapshot_for_export(RemoteSnapshotId::from_bytes([9; 16])).is_err()); + assert!(session_b.claim_snapshot_for_export(snapshot_id).is_err()); + assert!(session_a.snapshot_capabilities.lock().unwrap().capabilities.contains_key(&snapshot_id)); + + let claim = session_a.claim_snapshot_for_export(snapshot_id).unwrap(); + assert!(session_a.snapshot_capabilities.lock().unwrap().capabilities.contains_key(&snapshot_id)); + drop(claim); + session_a.abort_snapshot_capability(snapshot_id).unwrap(); + assert_eq!(published_snapshot_count(&session_a), 0); +} + +#[test] +fn aborting_an_unclaimed_capability_releases_snapshot_storage() { + let (_temp, session, _target, _session_id) = fixture(); + let snapshot_id = session.sync_snapshot().unwrap(); + + assert_eq!(published_snapshot_count(&session), 1); + session.abort_snapshot_capability(snapshot_id).unwrap(); + assert_eq!(published_snapshot_count(&session), 0); + assert!(session.claim_snapshot_for_export(snapshot_id).is_err()); + assert!(session.abort_snapshot_capability(snapshot_id).is_err()); +} + +#[test] +fn export_claim_keeps_session_cleanup_scoped_until_drop() { + let (_temp, session, _target, _session_id) = fixture(); + let snapshot_id = session.sync_snapshot().unwrap(); + let snapshot_root = session.snapshot_store_root(); + let claim = session.claim_snapshot_for_export(snapshot_id).unwrap(); + + drop(session); + assert!(snapshot_root.exists()); + drop(claim); + assert!(!snapshot_root.exists()); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn interleaved_syncs_build_their_own_snapshot_capabilities() { let (_temp, session, target, session_id) = fixture(); diff --git a/src/remote_target_ut.rs b/src/remote_target_ut.rs new file mode 100644 index 0000000..1753ffc --- /dev/null +++ b/src/remote_target_ut.rs @@ -0,0 +1,321 @@ +use super::*; +use std::fs; +use std::path::{Path, PathBuf}; +use std::time::Duration; +use tempfile::{tempdir, TempDir}; + +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; + +struct Fixture { + temp: TempDir, + project: PathBuf, + key: PathBuf, + known_hosts: PathBuf, + config: PathBuf, +} + +impl Fixture { + fn new() -> Self { + let temp = tempdir().unwrap(); + let project = temp.path().join("project"); + fs::create_dir(&project).unwrap(); + + let key = temp.path().join("id_ed25519"); + fs::write(&key, b"not-read-by-this-module").unwrap(); + set_mode(&key, 0o600); + + let known_hosts = temp.path().join("known_hosts"); + fs::write(&known_hosts, b"example ssh-ed25519 AAAA\n").unwrap(); + set_mode(&known_hosts, 0o644); + + Self { config: temp.path().join("remote-targets.yaml"), temp, project, key, known_hosts } + } +} + +fn set_mode(path: &Path, mode: u32) { + #[cfg(unix)] + fs::set_permissions(path, fs::Permissions::from_mode(mode)).unwrap(); + #[cfg(not(unix))] + let _ = (path, mode); +} + +fn scalar(value: &str) -> String { + serde_yaml::to_string(value).unwrap().trim().to_string() +} + +fn valid_yaml(fixture: &Fixture, project: &Path, backend: &str, target: Option<&str>) -> String { + let target_line = target.map(|target| format!(" target: {}\n", scalar(target))).unwrap_or_default(); + format!( + r#"version: 1 +targets: + ssh-one: + transport: ssh + host: build.example.test + port: 2222 + user: builder + identity-file: {} + known-hosts-file: {} + worker-path: /usr/local/libexec/bunkerbox-worker + workspace-root: /var/tmp/bunkerbox-workers + tools: + cargo: /usr/local/bin/cargo + make: /usr/bin/make + environment: + CC: clang + PRIVATE_VALUE: super-secret-value + resources: + connect-timeout-seconds: 5 + sync-timeout-seconds: 120 + build-timeout-seconds: 3600 + max-output-bytes: 67108864 +projects: + {}: + backend: {} +{}"#, + scalar(fixture.key.to_str().unwrap()), + scalar(fixture.known_hosts.to_str().unwrap()), + scalar(project.to_str().unwrap()), + backend, + target_line, + ) +} + +fn write_config(fixture: &Fixture, yaml: &str) { + fs::write(&fixture.config, yaml).unwrap(); +} + +#[test] +fn valid_config_loads_and_validates_an_ssh_target() { + let fixture = Fixture::new(); + write_config(&fixture, &valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one"))); + + let config = RemoteTargetConfig::load_from(&fixture.config).unwrap(); + let target = config.ssh_target("ssh-one").unwrap(); + assert_eq!(target.name(), "ssh-one"); + assert_eq!(target.host(), "build.example.test"); + assert_eq!(target.port(), 2222); + assert_eq!(target.user(), "builder"); + assert_eq!(target.identity_file(), fixture.key.as_path()); + assert_eq!(target.known_hosts_file(), fixture.known_hosts.as_path()); + assert_eq!(target.worker_path(), "/usr/local/libexec/bunkerbox-worker"); + assert_eq!(target.workspace_root(), "/var/tmp/bunkerbox-workers"); + assert_eq!(target.tools().get("cargo").map(String::as_str), Some("/usr/local/bin/cargo")); + assert_eq!(target.environment().get("CC").map(String::as_str), Some("clang")); + assert_eq!(target.resources().connect_timeout(), Duration::from_secs(5)); + assert_eq!(target.resources().sync_timeout(), Duration::from_secs(120)); + assert_eq!(target.resources().build_timeout(), Duration::from_secs(3600)); + assert_eq!(target.resources().max_output_bytes(), 64 * 1024 * 1024); + + let debug = format!("{target:?}"); + assert!(!debug.contains("super-secret-value")); +} + +#[test] +fn ssh_and_loopback_projects_resolve_explicitly() { + let fixture = Fixture::new(); + let second_project = fixture.temp.path().join("loopback-project"); + fs::create_dir(&second_project).unwrap(); + let ssh_yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")); + let loopback_key = scalar(second_project.to_str().unwrap()); + let yaml = format!("{ssh_yaml} {loopback_key}:\n backend: loopback\n target: target-that-does-not-exist\n"); + write_config(&fixture, &yaml); + + let config = RemoteTargetConfig::load_from(&fixture.config).unwrap(); + let ssh = config.resolve_for_project(&fixture.project).unwrap(); + assert_eq!(ssh.backend(), BackendMode::Ssh); + assert_eq!(ssh.target().unwrap().name(), "ssh-one"); + + let loopback = config.resolve_for_project(&second_project).unwrap(); + assert_eq!(loopback.backend(), BackendMode::Loopback); + assert!(loopback.target().is_none()); +} + +#[test] +fn missing_project_binding_is_an_error() { + let fixture = Fixture::new(); + let other = fixture.temp.path().join("other-project"); + fs::create_dir(&other).unwrap(); + write_config(&fixture, &valid_yaml(&fixture, &other, "loopback", None)); + + let config = RemoteTargetConfig::load_from(&fixture.config).unwrap(); + assert!(config.resolve_for_project(&fixture.project).is_err()); +} + +#[test] +fn ssh_requires_a_known_target_and_a_target_name() { + let fixture = Fixture::new(); + write_config(&fixture, &valid_yaml(&fixture, &fixture.project, "ssh", Some("missing"))); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + write_config(&fixture, &valid_yaml(&fixture, &fixture.project, "ssh", None)); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} + +#[test] +fn duplicate_yaml_map_keys_are_rejected() { + let fixture = Fixture::new(); + let project = scalar(fixture.project.to_str().unwrap()); + let binding = " backend: loopback\n".to_string(); + let yaml = format!("version: 1\ntargets: {{}}\nprojects:\n {project}:\n{binding} {project}:\n{binding}"); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} + +#[test] +fn duplicate_canonical_project_bindings_are_rejected() { + let fixture = Fixture::new(); + let alias = fixture.temp.path().join("project-alias"); + #[cfg(unix)] + std::os::unix::fs::symlink(&fixture.project, &alias).unwrap(); + #[cfg(not(unix))] + fs::create_dir(&alias).unwrap(); + + let project = scalar(fixture.project.to_str().unwrap()); + let alias = scalar(alias.to_str().unwrap()); + let yaml = format!("version: 1\ntargets: {{}}\nprojects:\n {project}:\n backend: loopback\n {alias}:\n backend: loopback\n"); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} + +#[test] +fn invalid_host_port_and_user_are_rejected() { + let fixture = Fixture::new(); + for (host, port, user) in [ + ("bad;host", "22", "builder"), + ("build.example.test", "0", "builder"), + ("build.example.test", "65536", "builder"), + ("build.example.test", "22", "bad user"), + ] { + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")) + .replace("build.example.test", &scalar(host)) + .replace("port: 2222", &format!("port: {port}")) + .replace("user: builder", &format!("user: {}", scalar(user))); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err(), "accepted invalid target fields"); + } +} + +#[test] +fn missing_unreadable_and_insecure_identity_files_are_rejected() { + let fixture = Fixture::new(); + let missing = fixture.temp.path().join("missing-key"); + let missing_yaml = + valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(fixture.key.to_str().unwrap(), missing.to_str().unwrap()); + write_config(&fixture, &missing_yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + set_mode(&fixture.key, 0o644); + write_config(&fixture, &valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one"))); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + set_mode(&fixture.key, 0o000); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} + +#[test] +fn known_hosts_must_be_a_readable_regular_file() { + let fixture = Fixture::new(); + let missing = fixture.temp.path().join("missing-known-hosts"); + let missing_yaml = + valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(fixture.known_hosts.to_str().unwrap(), missing.to_str().unwrap()); + write_config(&fixture, &missing_yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + let directory = fixture.temp.path().join("known-hosts-directory"); + fs::create_dir(&directory).unwrap(); + let directory_yaml = + valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(fixture.known_hosts.to_str().unwrap(), directory.to_str().unwrap()); + write_config(&fixture, &directory_yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + set_mode(&fixture.known_hosts, 0o000); + write_config(&fixture, &valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one"))); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} + +#[test] +fn worker_root_tool_and_local_paths_are_strict() { + let fixture = Fixture::new(); + for (worker, root) in [ + ("relative-worker", "/var/tmp/workers"), + ("/usr/bin/worker/../bad", "/var/tmp/workers"), + ("/usr/bin/worker", "relative-root"), + ("/usr/bin/worker", "/var/tmp/workers/../escape"), + ] { + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")) + .replace("/usr/local/libexec/bunkerbox-worker", worker) + .replace("/var/tmp/bunkerbox-workers", root); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + } + + let tool_path_yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace("/usr/local/bin/cargo", "relative-cargo"); + write_config(&fixture, &tool_path_yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + let tool_identity_yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(" cargo:", " bad/tool:"); + write_config(&fixture, &tool_identity_yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + let local_path_yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(fixture.key.to_str().unwrap(), "relative-key"); + write_config(&fixture, &local_path_yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} + +#[test] +fn environment_names_values_and_resource_bounds_are_validated() { + let fixture = Fixture::new(); + for replacement in [ + (" CC: clang", " BAD-NAME: clang"), + (" CC: clang", " CC: bad\nvalue"), + (" connect-timeout-seconds: 5", " connect-timeout-seconds: 0"), + (" max-output-bytes: 67108864", " max-output-bytes: 0"), + ] { + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(replacement.0, replacement.1); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + } +} + +#[test] +fn project_binding_paths_must_be_canonical_and_resolution_rejects_traversal() { + let fixture = Fixture::new(); + let injected = fixture.temp.path().join("project").join("..").join("project"); + let yaml = valid_yaml(&fixture, &injected, "loopback", None); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + write_config(&fixture, &valid_yaml(&fixture, &fixture.project, "loopback", None)); + let config = RemoteTargetConfig::load_from(&fixture.config).unwrap(); + assert!(config.resolve_for_project(&injected).is_err()); +} + +#[test] +fn default_path_resolution_is_injectable_and_uses_xdg_then_home_fallback() { + let fixture = Fixture::new(); + let xdg = fixture.temp.path().join("xdg"); + let home = fixture.temp.path().join("home"); + let helper = ConfigPathHelper::new(Some(xdg.clone()), Some(home.clone())); + assert_eq!(helper.config_path().unwrap(), xdg.join("bunkerbox").join(CONFIG_FILE_NAME)); + + let fallback = ConfigPathHelper::new(None, Some(home.clone())); + assert_eq!(fallback.config_path().unwrap(), home.join(".config").join("bunkerbox").join(CONFIG_FILE_NAME)); + + fs::create_dir_all(xdg.join("bunkerbox")).unwrap(); + let config_path = helper.config_path().unwrap(); + fs::write(&config_path, valid_yaml(&fixture, &fixture.project, "loopback", None)).unwrap(); + assert_eq!(RemoteTargetConfig::load_default_with_path_helper(&helper).unwrap().source_path(), config_path.as_path()); +} + +#[test] +fn config_version_and_transport_are_strict() { + let fixture = Fixture::new(); + let version = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace("version: 1", "version: 2"); + write_config(&fixture, &version); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + + let transport = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace("transport: ssh", "transport: loopback"); + write_config(&fixture, &transport); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); +} diff --git a/src/snapshot_ut.rs b/src/snapshot_ut.rs index 5c520ce..26cf80a 100644 --- a/src/snapshot_ut.rs +++ b/src/snapshot_ut.rs @@ -341,3 +341,31 @@ fn materialization_rejects_destination_symlink() { assert!(SnapshotStore::new(store_dir.path()).materialize(snapshot.handle(), &destination).is_err()); assert!(outside.path().read_dir().unwrap().next().is_none()); } + +#[test] +fn export_reads_validated_manifest_files_with_bounded_no_follow_access() { + let source = TempDir::new().unwrap(); + let store_dir = TempDir::new().unwrap(); + write_file(source.path(), "nested/file", b"abcdef"); + let snapshot = build_at(&source, &store_dir, SnapshotLimits::default(), &[]).unwrap(); + let store = SnapshotStore::new(store_dir.path()); + let export = store.resolve_export(snapshot.handle()).unwrap(); + let entry = snapshot.entries().iter().find(|entry| entry.path().as_str() == "nested/file").unwrap(); + + assert_eq!(export.entries(), snapshot.entries()); + assert_eq!(export.read_file(entry).unwrap(), b"abcdef"); + assert!(export.read_file_bounded(entry, 5).is_err()); + + let staged_file = store.snapshot_path(snapshot.handle()).join("files/nested/file"); + fs::write(&staged_file, b"ghijkl").unwrap(); + assert!(export.read_file(entry).is_err()); + + fs::write(&staged_file, b"abcdef-longer").unwrap(); + assert!(export.read_file(entry).is_err()); + + fs::remove_file(&staged_file).unwrap(); + let outside = TempDir::new().unwrap(); + fs::write(outside.path().join("file"), b"abcdef").unwrap(); + symlink(outside.path().join("file"), &staged_file).unwrap(); + assert!(export.read_file(entry).is_err()); +} diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs new file mode 100644 index 0000000..60284dc --- /dev/null +++ b/src/ssh_ut.rs @@ -0,0 +1,336 @@ +use super::*; +use crate::remote::{ + RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteSnapshotId, RemoteTool, RequestId, WorkspaceRelativePath, + WorkspaceSessionId, +}; +use crate::remote_target::RemoteTargetConfig; +use crate::snapshot::{SnapshotBuilder, SnapshotExclusionPolicy, SnapshotLimits, SnapshotStore}; +use crate::worker_protocol::{self, WorkerErrorKind, WorkerMessage, WorkerOperation}; +use std::fs; +use std::path::Path; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; +use tempfile::{tempdir, TempDir}; +use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; +use tokio::sync::oneshot; + +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; + +#[derive(Clone, Copy)] +enum ScriptMode { + Success, + WrongVersion, + ProtocolError, +} + +struct ScriptedFactory { + modes: Vec, + next: AtomicUsize, + specs: Arc>>, + messages: Arc>>, +} + +impl ScriptedFactory { + fn new(modes: Vec) -> Arc { + Arc::new(Self { modes, next: AtomicUsize::new(0), specs: Arc::new(Mutex::new(Vec::new())), messages: Arc::new(Mutex::new(Vec::new())) }) + } +} + +impl SshProcessFactory for ScriptedFactory { + fn spawn(&self, spec: &SshLaunchSpec) -> Result, String> { + let index = self.next.fetch_add(1, Ordering::Relaxed); + let mode = *self.modes.get(index).ok_or_else(|| "unexpected scripted SSH process".to_string())?; + self.specs.lock().unwrap().push(spec.clone()); + + let (host_writer, worker_reader) = tokio::io::duplex(8192); + let (worker_writer, host_reader) = tokio::io::duplex(8192); + let (host_stderr, _worker_stderr) = tokio::io::duplex(256); + let (status_tx, status_rx) = oneshot::channel(); + let messages = self.messages.clone(); + tokio::spawn(async move { + let status = scripted_worker(worker_reader, worker_writer, mode, messages).await; + let _ = status_tx.send(status); + }); + + Ok(Box::new(ScriptedProcess { + stdin: Some(Box::new(host_writer)), + stdout: Some(Box::new(host_reader)), + stderr: Some(Box::new(host_stderr)), + status: Some(status_rx), + killed: Arc::new(AtomicBool::new(false)), + })) + } +} + +struct ScriptedProcess { + stdin: Option, + stdout: Option, + stderr: Option, + status: Option>, + killed: Arc, +} + +impl SshProcess for ScriptedProcess { + fn take_stdin(&mut self) -> Option { + self.stdin.take() + } + + fn take_stdout(&mut self) -> Option { + self.stdout.take() + } + + fn take_stderr(&mut self) -> Option { + self.stderr.take() + } + + fn terminate_group(&mut self) { + self.killed.store(true, Ordering::Relaxed); + } + + fn kill_group(&mut self) { + self.killed.store(true, Ordering::Relaxed); + } + + fn wait<'a>(&'a mut self) -> crate::remote::RemoteFuture<'a, Result> { + Box::pin(async move { + let status = self.status.take().ok_or_else(|| "scripted process was already waited".to_string())?; + let code = status.await.map_err(|_| "scripted worker stopped without an exit status".to_string())?; + Ok(code) + }) + } +} + +async fn scripted_worker(mut reader: R, mut writer: W, mode: ScriptMode, messages: Arc>>) -> i32 +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + let Ok(WorkerMessage::Hello { request_id, session_id, .. }) = worker_protocol::read_message(&mut reader).await else { + return 71; + }; + if matches!(mode, ScriptMode::WrongVersion) { + let mut frame = WorkerMessage::hello(request_id, session_id, true).encode().unwrap(); + frame[4..6].copy_from_slice(&(worker_protocol::WORKER_PROTOCOL_VERSION + 1).to_le_bytes()); + writer.write_all(&frame).await.unwrap(); + return 0; + } + worker_protocol::write_message(&mut writer, &WorkerMessage::hello(request_id, session_id, true)).await.unwrap(); + + let Ok(first) = worker_protocol::read_message(&mut reader).await else { return 72 }; + messages.lock().unwrap().push(first.clone()); + if matches!(mode, ScriptMode::ProtocolError) { + worker_protocol::write_message( + &mut writer, + &WorkerMessage::error(request_id, session_id, WorkerOperation::Upload, WorkerErrorKind::WorkerProtocol, "scripted protocol failure"), + ) + .await + .unwrap(); + while worker_protocol::read_message(&mut reader).await.is_ok() {} + return 0; + } + match first { + WorkerMessage::UploadBegin { request_id, session_id, upload_id, .. } => loop { + let Ok(message) = worker_protocol::read_message(&mut reader).await else { return 73 }; + messages.lock().unwrap().push(message.clone()); + match message { + WorkerMessage::UploadFileChunk { .. } => {} + WorkerMessage::UploadComplete { request_id: received_request, session_id: received_session, upload_id: received_upload } => { + if received_request != request_id || received_session != session_id || received_upload != upload_id { + return 74; + } + worker_protocol::write_message( + &mut writer, + &WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 1, total_bytes: Some(1) }, + ) + .await + .unwrap(); + worker_protocol::write_message(&mut writer, &WorkerMessage::UploadComplete { request_id, session_id, upload_id }).await.unwrap(); + return 0; + } + _ => return 75, + } + }, + WorkerMessage::Build { request_id, session_id, build } => { + messages.lock().unwrap().push(WorkerMessage::Build { request_id, session_id, build: build.clone() }); + worker_protocol::write_message(&mut writer, &WorkerMessage::stdout(request_id, session_id, b"remote stdout\n".to_vec())).await.unwrap(); + worker_protocol::write_message(&mut writer, &WorkerMessage::stderr(request_id, session_id, b"remote stderr\n".to_vec())).await.unwrap(); + worker_protocol::write_message(&mut writer, &WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, 7)).await.unwrap(); + let Ok(cleanup) = worker_protocol::read_message(&mut reader).await else { return 76 }; + messages.lock().unwrap().push(cleanup.clone()); + if !matches!(cleanup, WorkerMessage::Cleanup { request_id: received_request, session_id: received_session, upload_token } if received_request == request_id && received_session == session_id && upload_token == build.upload_token()) + { + return 77; + } + worker_protocol::write_message(&mut writer, &WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0)) + .await + .unwrap(); + 0 + } + _ => 78, + } +} + +struct Fixture { + _temp: TempDir, + session: Arc, + target: crate::remote::RemoteTargetId, + session_id: WorkspaceSessionId, + ssh_target: crate::remote_target::SshTarget, +} + +fn fixture() -> Fixture { + let temp = tempdir().unwrap(); + let project = temp.path().join("project"); + fs::create_dir(&project).unwrap(); + fs::create_dir(project.join("src")).unwrap(); + fs::write(project.join("src/input.txt"), b"snapshot\n").unwrap(); + + let key = temp.path().join("id_ed25519"); + fs::write(&key, b"test key").unwrap(); + set_mode(&key, 0o600); + let known_hosts = temp.path().join("known_hosts"); + fs::write(&known_hosts, b"build.example.test ssh-ed25519 AAAA\n").unwrap(); + set_mode(&known_hosts, 0o644); + let config_path = temp.path().join("remote-targets.yaml"); + let project_yaml = serde_yaml::to_string(&project).unwrap(); + let key_yaml = serde_yaml::to_string(&key).unwrap(); + let known_yaml = serde_yaml::to_string(&known_hosts).unwrap(); + let yaml = format!( + "version: 1\ntargets:\n test:\n transport: ssh\n host: build.example.test\n port: 22\n user: builder\n identity-file: {}\n known-hosts-file: {}\n worker-path: /usr/local/libexec/bunkerbox-worker\n workspace-root: /var/tmp/bunkerbox-workers\n tools:\n make: /usr/bin/make\n environment:\n PATH: /usr/bin\n resources:\n connect-timeout-seconds: 5\n sync-timeout-seconds: 5\n build-timeout-seconds: 5\n max-output-bytes: 67108864\nprojects:\n {}:\n backend: ssh\n target: test\n", + key_yaml.trim(), known_yaml.trim(), project_yaml.trim() + ); + fs::write(&config_path, yaml).unwrap(); + let config = RemoteTargetConfig::load_from(&config_path).unwrap(); + let ssh_target = config.resolve_for_project(&project).unwrap().target().unwrap().clone(); + + let session_id = WorkspaceSessionId([1; 16]); + let target = crate::remote::RemoteTargetId([2; 16]); + let store = SnapshotStore::new(temp.path().join("snapshots")); + let exclusions = SnapshotExclusionPolicy::from_patterns(Vec::::new()).unwrap(); + let builder = SnapshotBuilder::new(store.clone(), SnapshotLimits::default(), exclusions); + let session = Arc::new(RunRemoteSession::new(session_id, target, project, store, builder, temp.path().join("jobs")).unwrap()); + Fixture { _temp: temp, session, target, session_id, ssh_target } +} + +fn set_mode(path: &Path, mode: u32) { + #[cfg(unix)] + fs::set_permissions(path, fs::Permissions::from_mode(mode)).unwrap(); + #[cfg(not(unix))] + let _ = (path, mode); +} + +fn authorize_sync(fixture: &Fixture) -> crate::remote::AuthorizedRemoteRequest { + let policy = + RemoteAuthorizationPolicy::new(fixture.target, fixture.session_id, vec!["make".to_string()]).with_snapshot_authority(fixture.session.clone()); + policy + .authorize( + &RemoteExecutionContext { target: fixture.target, workspace_session_id: fixture.session_id }, + RemoteRequest::sync(RequestId([3; 16]), fixture.session_id), + ) + .unwrap() +} + +fn authorize_build(fixture: &Fixture, snapshot_id: RemoteSnapshotId) -> crate::remote::AuthorizedRemoteRequest { + let build = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("make").unwrap(), + vec!["release".to_string(), "literal $(arg)".to_string()], + vec![("CC".to_string(), "clang".to_string())], + snapshot_id, + ) + .unwrap(); + let policy = + RemoteAuthorizationPolicy::new(fixture.target, fixture.session_id, vec!["make".to_string()]).with_snapshot_authority(fixture.session.clone()); + policy + .authorize( + &RemoteExecutionContext { target: fixture.target, workspace_session_id: fixture.session_id }, + RemoteRequest::build(RequestId([4; 16]), fixture.session_id, build), + ) + .unwrap() +} + +async fn collect(mut receiver: tokio::sync::mpsc::Receiver) -> Vec { + let mut events = Vec::new(); + while let Some(event) = receiver.recv().await { + events.push(event); + } + events +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn scripted_worker_proves_end_to_end_ssh_transport_and_exact_build_result() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::Success, ScriptMode::Success]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory.clone()); + + let (sync_tx, sync_rx) = tokio::sync::mpsc::channel(64); + backend.execute(authorize_sync(&fixture), sync_tx).await.unwrap(); + let sync_events = collect(sync_rx).await; + let snapshot_id = sync_events.iter().find_map(|event| match event { + RemoteBackendEvent::SyncCompleted { snapshot_id } => Some(*snapshot_id), + _ => None, + }); + let snapshot_id = snapshot_id.expect("scripted sync must retain a capability"); + + let (build_tx, build_rx) = tokio::sync::mpsc::channel(64); + backend.execute(authorize_build(&fixture, snapshot_id), build_tx).await.unwrap(); + let build_events = collect(build_rx).await; + assert!(build_events.iter().any(|event| matches!(event, RemoteBackendEvent::Stdout(bytes) if bytes == b"remote stdout\n"))); + assert!(build_events.iter().any(|event| matches!(event, RemoteBackendEvent::Stderr(bytes) if bytes == b"remote stderr\n"))); + assert!(build_events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 7 }))); + assert_eq!(backend.pending_uploads(), 0); + + let messages = factory.messages.lock().unwrap(); + assert!(messages.iter().any(|message| matches!(message, WorkerMessage::UploadBegin { .. }))); + assert!(messages.iter().any(|message| matches!(message, WorkerMessage::Build { build, .. } if build.argv() == ["release", "literal $(arg)"]))); + assert!(messages.iter().any(|message| matches!(message, WorkerMessage::Cleanup { .. }))); + let specs = factory.specs.lock().unwrap(); + assert_eq!(specs.len(), 2); + for spec in specs.iter() { + let joined = spec.args().join(" "); + assert!(joined.contains("StrictHostKeyChecking=yes")); + assert!(joined.contains("IdentityAgent=none")); + assert!(!joined.contains("literal $(arg)")); + assert!(!joined.contains("snapshot_id")); + assert_eq!(spec.remote_command(), "exec '/usr/local/libexec/bunkerbox-worker' --stdio --workspace-root '/var/tmp/bunkerbox-workers'"); + } +} + +#[test] +fn launch_spec_contains_only_fixed_trusted_ssh_arguments() { + let fixture = fixture(); + let spec = SshLaunchSpec::from_target(&fixture.ssh_target).unwrap(); + assert_eq!(spec.program(), Path::new("/usr/bin/ssh")); + assert!(spec.args().windows(2).any(|pair| pair == ["-F", "/dev/null"])); + assert!(spec.args().contains(&"BatchMode=yes".to_string())); + assert!(spec.args().contains(&"RequestTTY=no".to_string())); + assert!(spec.args().contains(&"ControlMaster=no".to_string())); + assert!(spec.args().contains(&"EscapeChar=none".to_string())); + assert!(!spec.args().iter().any(|arg| arg == "-tt" || arg == "accept-new" || arg.contains("SSH_AUTH_SOCK"))); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn worker_version_failure_is_typed_and_never_falls_back() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::WrongVersion]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory); + let (tx, _rx) = tokio::sync::mpsc::channel(16); + + let error = backend.execute(authorize_sync(&fixture), tx).await.unwrap_err(); + assert!(matches!(error, RemoteBackendError::Transport { class: RemoteFailureClass::WorkerVersion, .. })); + assert_eq!(backend.pending_uploads(), 0); + assert_eq!(fixture.session.snapshot_capability_count(), 0); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn worker_protocol_failure_is_terminal_without_local_fallback() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::ProtocolError]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory); + let (tx, _rx) = tokio::sync::mpsc::channel(16); + + let error = backend.execute(authorize_sync(&fixture), tx).await.unwrap_err(); + assert!(matches!(error, RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, .. }), "{error:?}"); + assert_eq!(fixture.session.snapshot_capability_count(), 0); +} diff --git a/src/worker_protocol_ut.rs b/src/worker_protocol_ut.rs new file mode 100644 index 0000000..dfe4c23 --- /dev/null +++ b/src/worker_protocol_ut.rs @@ -0,0 +1,264 @@ +use super::*; +use tokio::io::AsyncWriteExt; + +fn ids() -> (WorkerRequestId, WorkerSessionId, WorkerUploadId) { + (WorkerRequestId([1; 16]), WorkerSessionId([2; 16]), WorkerUploadId([3; 16])) +} + +fn file_entry(path: &str) -> WorkerUploadEntry { + WorkerUploadEntry::file(path, 0o644, 4, [9; WORKER_DIGEST_LEN]).unwrap() +} + +fn build() -> WorkerBuild { + WorkerBuild::new( + "make", + "/usr/bin/make", + vec!["release mode".into(), "$(literal); still one arg".into()], + "src", + vec![("CC".into(), "gcc".into())], + vec![("PATH".into(), "/usr/bin".into())], + WorkerUploadId([3; 16]), + ) + .unwrap() +} + +fn messages() -> Vec { + let (request_id, session_id, upload_id) = ids(); + vec![ + WorkerMessage::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response: false }, + WorkerMessage::UploadBegin { + request_id, + session_id, + upload_id, + entries: vec![WorkerUploadEntry::directory("src", 0o755).unwrap(), file_entry("src/main.rs")], + }, + WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index: 1, entry: file_entry("src/main.rs") }, + WorkerMessage::UploadFileChunk { + request_id, + session_id, + upload_id, + path: WorkerRelativePath::new("src/main.rs").unwrap(), + offset: 0, + data: b"data".to_vec(), + }, + WorkerMessage::UploadComplete { request_id, session_id, upload_id }, + WorkerMessage::Build { request_id, session_id, build: build() }, + WorkerMessage::Cleanup { request_id, session_id, upload_token: upload_id }, + WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 4, total_bytes: Some(8) }, + WorkerMessage::Stdout { request_id, session_id, data: b"out".to_vec() }, + WorkerMessage::Stderr { request_id, session_id, data: b"err".to_vec() }, + WorkerMessage::Completed { request_id, session_id, operation: WorkerOperation::Build, exit_code: -17 }, + WorkerMessage::Error { + request_id, + session_id, + operation: WorkerOperation::Cleanup, + kind: WorkerErrorKind::Cleanup, + message: "cleanup failed".into(), + }, + ] +} + +#[test] +fn all_message_variants_round_trip() { + for message in messages() { + let frame = message.encode().unwrap(); + assert_eq!(&frame[..4], &WORKER_PROTOCOL_MAGIC); + assert_eq!(frame[4..6], WORKER_PROTOCOL_VERSION.to_le_bytes()); + assert_eq!(frame[6], message.kind().as_u8()); + assert_eq!(WorkerMessage::decode(&frame).unwrap(), message); + } +} + +#[tokio::test] +async fn async_helpers_handle_fragmented_duplex_frames() { + let message = WorkerMessage::Build { request_id: ids().0, session_id: ids().1, build: build() }; + let frame = message.encode().unwrap(); + let (mut reader, mut writer) = tokio::io::duplex(3); + let sender = tokio::spawn(async move { + for part in frame.chunks(2) { + writer.write_all(part).await.unwrap(); + tokio::task::yield_now().await; + } + }); + let decoded = WorkerMessage::read_async(&mut reader).await.unwrap(); + sender.await.unwrap(); + assert_eq!(decoded, message); +} + +#[tokio::test] +async fn async_write_helper_emits_a_decodable_frame() { + let message = WorkerMessage::Stdout { request_id: ids().0, session_id: ids().1, data: b"exact".to_vec() }; + let (mut reader, mut writer) = tokio::io::duplex(128); + let expected = message.clone(); + let sender = tokio::spawn(async move { message.write_async(&mut writer).await.unwrap() }); + assert_eq!(WorkerMessage::read_async(&mut reader).await.unwrap(), expected); + sender.await.unwrap(); +} + +#[test] +fn malformed_headers_and_lengths_are_rejected() { + let frame = + WorkerMessage::Hello { request_id: ids().0, session_id: ids().1, version: WORKER_PROTOCOL_VERSION, response: false }.encode().unwrap(); + + let mut bad_magic = frame.clone(); + bad_magic[0] ^= 1; + assert!(WorkerMessage::decode(&bad_magic).is_err()); + + let mut bad_version = frame.clone(); + bad_version[4..6].copy_from_slice(&(WORKER_PROTOCOL_VERSION + 1).to_le_bytes()); + assert!(WorkerMessage::decode(&bad_version).is_err()); + + let mut bad_kind = frame.clone(); + bad_kind[6] = 255; + assert!(WorkerMessage::decode(&bad_kind).is_err()); + + let mut oversized = frame[..WORKER_FRAME_HEADER_LEN].to_vec(); + oversized[7..11].copy_from_slice(&((MAX_WORKER_FRAME_PAYLOAD as u32) + 1).to_le_bytes()); + assert!(WorkerMessage::decode(&oversized).is_err()); + + assert!(WorkerMessage::decode(&frame[..frame.len() - 1]).is_err()); + let mut extra = frame.clone(); + extra.push(0); + assert!(WorkerMessage::decode(&extra).is_err()); +} + +#[test] +fn invalid_utf8_and_trailing_payload_are_rejected() { + let (request_id, session_id, _) = ids(); + let mut frame = + WorkerMessage::Error { request_id, session_id, operation: WorkerOperation::Build, kind: WorkerErrorKind::Build, message: "x".into() } + .encode() + .unwrap(); + let message_start = WORKER_FRAME_HEADER_LEN + 32 + 1 + 1 + 4; + frame[message_start] = 0xff; + assert!(WorkerMessage::decode(&frame).is_err()); + + let mut hello = WorkerMessage::hello(request_id, session_id, false).encode().unwrap(); + hello.push(0); + assert!(WorkerMessage::decode(&hello).is_err()); +} + +#[test] +fn paths_tools_and_executables_are_strictly_validated() { + for path in ["/absolute", "../parent", "a/../b", "./name", "a//b", "a\\b", ""] { + assert!(validate_worker_relative_path(path).is_err(), "accepted path {path:?}"); + } + assert!(validate_worker_relative_path("src/main.rs").is_ok()); + assert!(validate_worker_relative_path(&"a".repeat(MAX_WORKER_PATH_COMPONENT_BYTES + 1)).is_err()); + assert!(WorkerTool::new("make;rm").is_err()); + assert!(WorkerTool::new("/usr/bin/make").is_err()); + assert!(WorkerExecutablePath::new("make").is_err()); + assert!(WorkerExecutablePath::new("/usr/bin/make -f").is_err()); + assert!(WorkerExecutablePath::new("/usr/bin/../bin/make").is_err()); + assert!(WorkerExecutablePath::new("/usr/bin/make").is_ok()); +} + +#[test] +fn manifests_require_sorted_paths_and_valid_metadata() { + assert!(validate_upload_manifest(&[file_entry("z"), file_entry("a")]).is_err()); + assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, None).is_err()); + assert!(WorkerUploadEntry::new("dir", WorkerEntryKind::Directory, 0o755, 1, None).is_err()); + assert!(WorkerUploadEntry::new("dir", WorkerEntryKind::Directory, 0o755, 0, Some([1; 32])).is_err()); + assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o100000, 1, Some([1; 32])).is_err()); + assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, Some([1; 32])).is_ok()); +} + +#[test] +fn duplicate_environment_keys_and_bad_argv_are_rejected() { + assert!(WorkerBuild::new( + "make", + "/usr/bin/make", + vec!["x".into()], + "", + vec![("CC".into(), "one".into()), ("CC".into(), "two".into())], + Vec::new(), + ids().2, + ) + .is_err()); + + assert!(WorkerBuild::new("make", "/usr/bin/make", vec!["x\0y".into()], "", Vec::new(), Vec::new(), ids().2,).is_err()); + + assert!(WorkerBuild::new("make", "/usr/bin/make", Vec::new(), "", vec![("BAD-NAME".into(), "value".into())], Vec::new(), ids().2,).is_err()); + + assert!(WorkerBuild::new("make", "/usr/bin/make", vec!["arg".into(); MAX_WORKER_ARG_COUNT + 1], "", Vec::new(), Vec::new(), ids().2,).is_err()); + + assert!(WorkerBuild::new( + "make", + "/usr/bin/make", + Vec::new(), + "", + vec![("A".into(), "value".into()); MAX_WORKER_ENV_COUNT + 1], + Vec::new(), + ids().2, + ) + .is_err()); +} + +#[test] +fn bounded_chunks_output_errors_and_progress_are_rejected() { + let (request_id, session_id, upload_id) = ids(); + let path = WorkerRelativePath::new("file").unwrap(); + assert!(WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path: path.clone(), offset: 0, data: Vec::new() }.encode().is_err()); + assert!(WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset: 0, data: vec![0; MAX_WORKER_CHUNK_BYTES + 1] } + .encode() + .is_err()); + assert!(WorkerMessage::Stdout { request_id, session_id, data: vec![0; MAX_WORKER_OUTPUT_BYTES + 1] }.encode().is_err()); + assert!(WorkerMessage::Error { + request_id, + session_id, + operation: WorkerOperation::Build, + kind: WorkerErrorKind::Build, + message: "x".repeat(MAX_WORKER_ERROR_BYTES + 1), + } + .encode() + .is_err()); + assert!(WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 2, total_bytes: Some(1) }.encode().is_err()); +} + +#[test] +fn stdout_stderr_and_exact_completion_preserve_bytes_and_exit_code() { + let (request_id, session_id, _) = ids(); + let stdout = WorkerMessage::Stdout { request_id, session_id, data: b"out\0with bytes".to_vec() }; + let stderr = WorkerMessage::Stderr { request_id, session_id, data: b"err".to_vec() }; + let completed = WorkerMessage::Completed { request_id, session_id, operation: WorkerOperation::Build, exit_code: i32::MIN }; + assert_eq!(WorkerMessage::decode(&stdout.encode().unwrap()).unwrap(), stdout); + assert_eq!(WorkerMessage::decode(&stderr.encode().unwrap()).unwrap(), stderr); + assert_eq!(WorkerMessage::decode(&completed.encode().unwrap()).unwrap(), completed); +} + +#[test] +fn build_has_no_shell_packing_and_keeps_literal_arguments() { + let message = WorkerMessage::Build { request_id: ids().0, session_id: ids().1, build: build() }; + let decoded = WorkerMessage::decode(&message.encode().unwrap()).unwrap(); + let WorkerMessage::Build { build, .. } = decoded else { panic!("expected build") }; + assert_eq!(build.argv, vec!["release mode", "$(literal); still one arg"]); + assert_eq!(build.argv.len(), 2); + assert_eq!(build.tool.as_str(), "make"); + assert_eq!(build.trusted_executable.as_str(), "/usr/bin/make"); +} + +#[test] +fn error_kinds_preserve_protocol_build_and_cleanup_failures() { + let (request_id, session_id, _) = ids(); + for (operation, kind) in [ + (WorkerOperation::Protocol, WorkerErrorKind::WorkerProtocol), + (WorkerOperation::Build, WorkerErrorKind::Build), + (WorkerOperation::Cleanup, WorkerErrorKind::Cleanup), + ] { + let message = WorkerMessage::Error { request_id, session_id, operation, kind, message: "failure".into() }; + assert_eq!(WorkerMessage::decode(&message.encode().unwrap()).unwrap(), message); + } +} + +#[test] +fn unknown_nested_kinds_and_flags_are_rejected() { + let message = WorkerMessage::UploadBegin { request_id: ids().0, session_id: ids().1, upload_id: ids().2, entries: Vec::new() }; + let mut frame = message.encode().unwrap(); + frame[WORKER_FRAME_HEADER_LEN + 48..WORKER_FRAME_HEADER_LEN + 52].copy_from_slice(&u32::MAX.to_le_bytes()); + assert!(WorkerMessage::decode(&frame).is_err()); + + let hello = WorkerMessage::hello(ids().0, ids().1, false); + let mut frame = hello.encode().unwrap(); + frame[WORKER_FRAME_HEADER_LEN + 34] = 9; + assert!(WorkerMessage::decode(&frame).is_err()); +} From d385b5605509935af7199c39245a267d294aeb4a Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 21:04:08 +0200 Subject: [PATCH 30/52] Implement remote worker for other platforms --- Cargo.lock | 18 + Cargo.toml | 9 + Makefile | 8 +- crates/bunkerbox-worker-protocol/Cargo.toml | 14 + crates/bunkerbox-worker-protocol/src/lib.rs | 1498 +++++++++++++++++ .../bunkerbox-worker-protocol/src/lib_ut.rs | 299 ++++ crates/bunkerbox-worker/Cargo.toml | 16 + crates/bunkerbox-worker/src/main.rs | 66 + crates/bunkerbox-worker/src/platform.rs | 222 +++ crates/bunkerbox-worker/src/platform_ut.rs | 39 + crates/bunkerbox-worker/src/process.rs | 362 ++++ crates/bunkerbox-worker/src/process_ut.rs | 71 + crates/bunkerbox-worker/src/storage.rs | 627 +++++++ crates/bunkerbox-worker/src/storage_ut.rs | 104 ++ crates/bunkerbox-worker/src/worker.rs | 307 ++++ crates/bunkerbox-worker/src/worker_ut.rs | 185 ++ src/ssh.rs | 68 +- src/ssh_ut.rs | 11 +- src/worker_protocol.rs | 1445 +--------------- src/worker_protocol_ut.rs | 308 +--- 20 files changed, 3972 insertions(+), 1705 deletions(-) create mode 100644 crates/bunkerbox-worker-protocol/Cargo.toml create mode 100644 crates/bunkerbox-worker-protocol/src/lib.rs create mode 100644 crates/bunkerbox-worker-protocol/src/lib_ut.rs create mode 100644 crates/bunkerbox-worker/Cargo.toml create mode 100644 crates/bunkerbox-worker/src/main.rs create mode 100644 crates/bunkerbox-worker/src/platform.rs create mode 100644 crates/bunkerbox-worker/src/platform_ut.rs create mode 100644 crates/bunkerbox-worker/src/process.rs create mode 100644 crates/bunkerbox-worker/src/process_ut.rs create mode 100644 crates/bunkerbox-worker/src/storage.rs create mode 100644 crates/bunkerbox-worker/src/storage_ut.rs create mode 100644 crates/bunkerbox-worker/src/worker.rs create mode 100644 crates/bunkerbox-worker/src/worker_ut.rs diff --git a/Cargo.lock b/Cargo.lock index 4ee7c91..d401d0a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -195,6 +195,7 @@ version = "0.4.1" dependencies = [ "aes-gcm", "base64", + "bunkerbox-worker-protocol", "clap", "colored", "crossterm 0.28.1", @@ -214,6 +215,23 @@ dependencies = [ "vt100", ] +[[package]] +name = "bunkerbox-worker" +version = "0.1.0" +dependencies = [ + "bunkerbox-worker-protocol", + "libc", + "sha2", + "tempfile", +] + +[[package]] +name = "bunkerbox-worker-protocol" +version = "0.1.0" +dependencies = [ + "tokio", +] + [[package]] name = "by_address" version = "1.2.1" diff --git a/Cargo.toml b/Cargo.toml index ea873e0..d40a77b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,9 +3,18 @@ name = "bunkerbox" version = "0.4.1" edition = "2021" +[workspace] +members = [ + ".", + "crates/bunkerbox-worker-protocol", + "crates/bunkerbox-worker", +] +resolver = "2" + [dependencies] aes-gcm = "0.10" base64 = "0.22" +bunkerbox-worker-protocol = { path = "crates/bunkerbox-worker-protocol", features = ["async"] } clap = "4" colored = "3" glob = "0.3" diff --git a/Makefile b/Makefile index 2cdebb6..8793fb5 100644 --- a/Makefile +++ b/Makefile @@ -1,9 +1,10 @@ .DEFAULT_GOAL := help -.PHONY: help ensure-toolchain dev release check test integration-test setup image install-image prepare config docs docs-dev docs-clean musl-vscomm clean +.PHONY: help ensure-toolchain dev release check test integration-test setup image install-image prepare config docs docs-dev docs-clean musl-vscomm worker-netbsd clean DOCS_VENV := .venv-docs DOCS_MKDOCS := $(DOCS_VENV)/bin/mkdocs VSCOMM_TARGET := x86_64-unknown-linux-musl +WORKER_TARGET ?= x86_64-unknown-netbsd IMAGE ?= OCI ?= @@ -18,6 +19,7 @@ help: @printf " %-24s %s\n" "Toolchain" "" @printf " %-24s %s\n" " ensure-toolchain" "Install/update Rust stable and musl target" @printf " %-24s %s\n" " musl-vscomm" "Build static vscomm binary only" + @printf " %-24s %s\n" " worker-netbsd" "Build the portable worker for NetBSD" @printf " %-24s %s\n" "" "" @printf " %-24s %s\n" "Image" "" @printf " %-24s %s\n" " image" "Build OCI agent image (requires IMAGE=)" @@ -43,6 +45,7 @@ ensure-toolchain: dev: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image + cargo build -p bunkerbox-worker cargo build --bin bunkerbox-netrelay --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) cargo build --bin bunkerbox-status --target $(VSCOMM_TARGET) rm -rf target/dist @@ -86,6 +89,9 @@ musl-vscomm: ensure-toolchain cargo build --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) cargo build --bin bunkerbox-status --target $(VSCOMM_TARGET) +worker-netbsd: + cargo build -p bunkerbox-worker --target $(WORKER_TARGET) --release + image: dev @if [ -z "$(IMAGE)" ]; then echo "usage: make image IMAGE=images/name.conf" >&2; exit 1; fi target/debug/bunkerbox-image $(IMAGE) diff --git a/crates/bunkerbox-worker-protocol/Cargo.toml b/crates/bunkerbox-worker-protocol/Cargo.toml new file mode 100644 index 0000000..64be961 --- /dev/null +++ b/crates/bunkerbox-worker-protocol/Cargo.toml @@ -0,0 +1,14 @@ +[package] +name = "bunkerbox-worker-protocol" +version = "0.1.0" +edition = "2021" + +[features] +default = ["async"] +async = ["dep:tokio"] + +[dependencies] +tokio = { version = "1", features = ["io-util"], optional = true } + +[dev-dependencies] +tokio = { version = "1", features = ["io-util", "macros", "rt"] } diff --git a/crates/bunkerbox-worker-protocol/src/lib.rs b/crates/bunkerbox-worker-protocol/src/lib.rs new file mode 100644 index 0000000..39aa017 --- /dev/null +++ b/crates/bunkerbox-worker-protocol/src/lib.rs @@ -0,0 +1,1498 @@ +//! A bounded binary protocol for a host-side worker reached over SSH. +//! +//! This protocol is deliberately independent from `vscomm`. A frame is: +//! +//! ```text +//! magic[4] version[u16] kind[u8] payload_length[u32] payload[payload_length] +//! ``` +//! +//! All integer fields are little-endian. The payload length is checked before +//! allocating a payload buffer, and every length and count inside a payload is +//! checked before it can drive an allocation. + +use std::collections::BTreeSet; +use std::fmt; +use std::io::{self, Read, Write}; + +#[cfg(feature = "async")] +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; + +pub const WORKER_PROTOCOL_MAGIC: [u8; 4] = *b"BBWK"; +pub const WORKER_MAGIC: [u8; 4] = WORKER_PROTOCOL_MAGIC; +pub const WORKER_PROTOCOL_VERSION: u16 = 1; +pub const WORKER_VERSION: u16 = WORKER_PROTOCOL_VERSION; +pub const WORKER_FRAME_HEADER_LEN: usize = 4 + 2 + 1 + 4; +pub const WORKER_FRAME_HEADER_SIZE: usize = WORKER_FRAME_HEADER_LEN; +pub const WORKER_ID_LEN: usize = 16; +pub const WORKER_DIGEST_LEN: usize = 32; + +/// Maximum bytes in one worker payload. The declared length is rejected +/// before a buffer of this size is allocated. +pub const MAX_WORKER_FRAME_PAYLOAD: usize = 1024 * 1024; +pub const MAX_WORKER_PAYLOAD: usize = MAX_WORKER_FRAME_PAYLOAD; +pub const MAX_WORKER_FRAME_BYTES: usize = WORKER_FRAME_HEADER_LEN + MAX_WORKER_FRAME_PAYLOAD; +pub const MAX_WORKER_FRAME_LENGTH: usize = MAX_WORKER_FRAME_BYTES; +pub const MAX_WORKER_STRING_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_PATH_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ENTRY_PATH_BYTES: usize = MAX_WORKER_PATH_BYTES; +pub const MAX_WORKER_PATH_COMPONENT_BYTES: usize = 255; +pub const MAX_WORKER_PATH_DEPTH: usize = 64; +pub const MAX_WORKER_CWD_BYTES: usize = MAX_WORKER_PATH_BYTES; +pub const MAX_WORKER_TOOL_BYTES: usize = 256; +pub const MAX_WORKER_EXECUTABLE_PATH_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_EXECUTABLE_BYTES: usize = MAX_WORKER_EXECUTABLE_PATH_BYTES; +pub const MAX_WORKER_ARG_COUNT: usize = 256; +pub const MAX_WORKER_ARGUMENT_COUNT: usize = MAX_WORKER_ARG_COUNT; +pub const MAX_WORKER_ARG_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ARG_TOTAL_BYTES: usize = 256 * 1024; +pub const MAX_WORKER_ENV_COUNT: usize = 128; +pub const MAX_WORKER_ENVIRONMENT_COUNT: usize = MAX_WORKER_ENV_COUNT; +pub const MAX_WORKER_ENV_KEY_BYTES: usize = 256; +pub const MAX_WORKER_ENV_VALUE_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ENV_TOTAL_BYTES: usize = 256 * 1024; +pub const MAX_WORKER_UPLOAD_ENTRIES: usize = 4096; +pub const MAX_WORKER_ENTRY_COUNT: usize = MAX_WORKER_UPLOAD_ENTRIES; +pub const MAX_WORKER_UPLOAD_ENTRY_COUNT: usize = MAX_WORKER_UPLOAD_ENTRIES; +pub const MAX_WORKER_FILE_BYTES: u64 = 256 * 1024 * 1024; +pub const MAX_WORKER_TOTAL_UPLOAD_BYTES: u64 = 512 * 1024 * 1024; +pub const MAX_WORKER_MANIFEST_BYTES: usize = 512 * 1024; +pub const MAX_WORKER_CHUNK_BYTES: usize = 64 * 1024; +pub const MAX_WORKER_OUTPUT_BYTES: usize = 64 * 1024; +pub const MAX_WORKER_ERROR_BYTES: usize = 4 * 1024; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum WorkerProtocolError { + Io(String), + Invalid(String), +} + +pub type WorkerResult = Result; + +impl WorkerProtocolError { + pub fn is_invalid(&self) -> bool { + matches!(self, Self::Invalid(_)) + } + + pub fn contains(&self, needle: &str) -> bool { + match self { + Self::Io(message) | Self::Invalid(message) => message.contains(needle), + } + } +} + +impl fmt::Display for WorkerProtocolError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Io(message) => write!(formatter, "worker I/O error: {message}"), + Self::Invalid(message) => write!(formatter, "invalid worker protocol: {message}"), + } + } +} + +impl std::error::Error for WorkerProtocolError {} + +impl From for WorkerProtocolError { + fn from(error: io::Error) -> Self { + Self::Io(error.to_string()) + } +} + +macro_rules! worker_id { + ($name:ident) => { + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Default)] + pub struct $name(pub [u8; WORKER_ID_LEN]); + + impl $name { + pub const fn new(bytes: [u8; WORKER_ID_LEN]) -> Self { + Self(bytes) + } + + pub const fn as_bytes(&self) -> &[u8; WORKER_ID_LEN] { + &self.0 + } + + pub const fn into_bytes(self) -> [u8; WORKER_ID_LEN] { + self.0 + } + } + + impl From<[u8; WORKER_ID_LEN]> for $name { + fn from(bytes: [u8; WORKER_ID_LEN]) -> Self { + Self(bytes) + } + } + }; +} + +worker_id!(WorkerRequestId); +worker_id!(WorkerSessionId); +worker_id!(WorkerUploadId); + +pub type RequestId = WorkerRequestId; +pub type SessionId = WorkerSessionId; +pub type UploadId = WorkerUploadId; +pub type UploadToken = WorkerUploadId; +pub type WorkerUploadToken = WorkerUploadId; +pub type WorkerDigest = [u8; WORKER_DIGEST_LEN]; + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerFrameKind { + Hello = 1, + UploadBegin = 2, + UploadEntry = 3, + UploadFileChunk = 4, + UploadComplete = 5, + Build = 6, + Cleanup = 7, + SyncProgress = 8, + Stdout = 9, + Stderr = 10, + Completed = 11, + Error = 12, +} + +impl WorkerFrameKind { + pub const fn as_u8(self) -> u8 { + self as u8 + } + + pub const fn from_u8(value: u8) -> Option { + match value { + 1 => Some(Self::Hello), + 2 => Some(Self::UploadBegin), + 3 => Some(Self::UploadEntry), + 4 => Some(Self::UploadFileChunk), + 5 => Some(Self::UploadComplete), + 6 => Some(Self::Build), + 7 => Some(Self::Cleanup), + 8 => Some(Self::SyncProgress), + 9 => Some(Self::Stdout), + 10 => Some(Self::Stderr), + 11 => Some(Self::Completed), + 12 => Some(Self::Error), + _ => None, + } + } +} + +impl TryFrom for WorkerFrameKind { + type Error = WorkerProtocolError; + + fn try_from(value: u8) -> WorkerResult { + Self::from_u8(value).ok_or_else(|| invalid(format!("unknown worker frame kind: {value}"))) + } +} + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerOperation { + Protocol = 0, + Upload = 1, + Build = 2, + Cleanup = 3, + Sync = 4, +} + +impl WorkerOperation { + pub const fn as_u8(self) -> u8 { + self as u8 + } + + pub fn from_u8(value: u8) -> WorkerResult { + match value { + 0 => Ok(Self::Protocol), + 1 => Ok(Self::Upload), + 2 => Ok(Self::Build), + 3 => Ok(Self::Cleanup), + 4 => Ok(Self::Sync), + _ => Err(invalid(format!("unknown worker operation: {value}"))), + } + } +} + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerErrorKind { + WorkerProtocol = 1, + Upload = 2, + Build = 3, + Cleanup = 4, + Sync = 5, +} + +pub type WorkerErrorClass = WorkerErrorKind; + +impl WorkerErrorKind { + #[allow(non_upper_case_globals)] + pub const Protocol: Self = Self::WorkerProtocol; + + pub const fn as_u8(self) -> u8 { + self as u8 + } + + pub fn from_u8(value: u8) -> WorkerResult { + match value { + 1 => Ok(Self::WorkerProtocol), + 2 => Ok(Self::Upload), + 3 => Ok(Self::Build), + 4 => Ok(Self::Cleanup), + 5 => Ok(Self::Sync), + _ => Err(invalid(format!("unknown worker error kind: {value}"))), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerRelativePath(String); + +impl WorkerRelativePath { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_relative_path("worker relative path", &value, true)?; + Ok(Self(value)) + } + + fn for_entry(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_relative_path("worker entry path", &value, false)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } + + pub fn is_empty(&self) -> bool { + self.0.is_empty() + } +} + +impl AsRef for WorkerRelativePath { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerRelativePath { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerTool(String); + +impl WorkerTool { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_tool(&value)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl AsRef for WorkerTool { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerTool { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerExecutablePath(String); + +impl WorkerExecutablePath { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_executable_path(&value)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl AsRef for WorkerExecutablePath { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerExecutablePath { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[repr(u8)] +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerEntryKind { + Directory = 1, + File = 2, +} + +impl WorkerEntryKind { + pub fn from_u8(value: u8) -> WorkerResult { + match value { + 1 => Ok(Self::Directory), + 2 => Ok(Self::File), + _ => Err(invalid(format!("unknown worker entry kind: {value}"))), + } + } + + pub const fn as_u8(self) -> u8 { + self as u8 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkerUploadEntry { + pub path: WorkerRelativePath, + pub kind: WorkerEntryKind, + pub mode: u32, + pub size: u64, + pub digest: Option, +} + +impl WorkerUploadEntry { + pub fn new(path: impl Into, kind: WorkerEntryKind, mode: u32, size: u64, digest: Option) -> WorkerResult { + let entry = Self { path: WorkerRelativePath::for_entry(path)?, kind, mode, size, digest }; + entry.validate()?; + Ok(entry) + } + + pub fn directory(path: impl Into, mode: u32) -> WorkerResult { + Self::new(path, WorkerEntryKind::Directory, mode, 0, None) + } + + pub fn file(path: impl Into, mode: u32, size: u64, digest: WorkerDigest) -> WorkerResult { + Self::new(path, WorkerEntryKind::File, mode, size, Some(digest)) + } + + pub fn path(&self) -> &WorkerRelativePath { + &self.path + } + + pub fn kind(&self) -> WorkerEntryKind { + self.kind + } + + pub fn mode(&self) -> u32 { + self.mode + } + + pub fn size(&self) -> u64 { + self.size + } + + pub fn digest(&self) -> Option<&WorkerDigest> { + self.digest.as_ref() + } + + pub fn validate(&self) -> WorkerResult<()> { + validate_relative_path("worker entry path", self.path.as_str(), false)?; + if self.mode & !0o7777 != 0 { + return Err(invalid(format!("worker entry mode has unsupported bits: {:o}", self.mode))); + } + + match self.kind { + WorkerEntryKind::Directory => { + if self.size != 0 { + return Err(invalid("worker directory entry must have zero size")); + } + if self.digest.is_some() { + return Err(invalid("worker directory entry must not have a digest")); + } + } + WorkerEntryKind::File => { + if self.size > MAX_WORKER_FILE_BYTES { + return Err(invalid(format!("worker file exceeds maximum size {MAX_WORKER_FILE_BYTES}"))); + } + if self.digest.is_none() { + return Err(invalid("worker file entry is missing a digest")); + } + } + } + Ok(()) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkerBuild { + pub tool: WorkerTool, + pub trusted_executable: WorkerExecutablePath, + pub argv: Vec, + pub cwd: WorkerRelativePath, + pub guest_env: Vec<(String, String)>, + pub target_env: Vec<(String, String)>, + pub upload_token: WorkerUploadId, +} + +impl WorkerBuild { + pub fn new( + tool: impl Into, trusted_executable: impl Into, argv: Vec, cwd: impl Into, guest_env: Vec<(String, String)>, + target_env: Vec<(String, String)>, upload_token: WorkerUploadId, + ) -> WorkerResult { + let build = Self { + tool: WorkerTool::new(tool)?, + trusted_executable: WorkerExecutablePath::new(trusted_executable)?, + argv, + cwd: WorkerRelativePath::new(cwd)?, + guest_env, + target_env, + upload_token, + }; + build.validate()?; + Ok(build) + } + + pub fn validate(&self) -> WorkerResult<()> { + validate_tool(self.tool.as_str())?; + validate_executable_path(self.trusted_executable.as_str())?; + validate_relative_path("worker cwd", self.cwd.as_str(), true)?; + validate_argv(&self.argv)?; + validate_environment("worker guest environment", &self.guest_env)?; + validate_environment("worker target environment", &self.target_env)?; + Ok(()) + } + + pub fn tool(&self) -> &WorkerTool { + &self.tool + } + + pub fn trusted_executable(&self) -> &WorkerExecutablePath { + &self.trusted_executable + } + + pub fn trusted_executable_path(&self) -> &str { + self.trusted_executable.as_str() + } + + pub fn argv(&self) -> &[String] { + &self.argv + } + + pub fn cwd(&self) -> &WorkerRelativePath { + &self.cwd + } + + pub fn guest_env(&self) -> &[(String, String)] { + &self.guest_env + } + + pub fn target_env(&self) -> &[(String, String)] { + &self.target_env + } + + pub fn upload_token(&self) -> WorkerUploadId { + self.upload_token + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum WorkerMessage { + Hello { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + version: u16, + response: bool, + }, + UploadBegin { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entries: Vec, + }, + UploadEntry { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entry_index: u32, + entry: WorkerUploadEntry, + }, + UploadFileChunk { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + path: WorkerRelativePath, + offset: u64, + data: Vec, + }, + UploadComplete { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + }, + Build { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + build: WorkerBuild, + }, + Cleanup { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_token: WorkerUploadId, + }, + SyncProgress { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + completed_bytes: u64, + total_bytes: Option, + }, + Stdout { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + data: Vec, + }, + Stderr { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + data: Vec, + }, + Completed { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + operation: WorkerOperation, + exit_code: i32, + }, + Error { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + operation: WorkerOperation, + kind: WorkerErrorKind, + message: String, + }, +} + +impl WorkerMessage { + pub fn hello(request_id: WorkerRequestId, session_id: WorkerSessionId, response: bool) -> Self { + Self::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response } + } + + pub fn build(request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild) -> Self { + Self::Build { request_id, session_id, build } + } + + pub fn stdout(request_id: WorkerRequestId, session_id: WorkerSessionId, data: Vec) -> Self { + Self::Stdout { request_id, session_id, data } + } + + pub fn stderr(request_id: WorkerRequestId, session_id: WorkerSessionId, data: Vec) -> Self { + Self::Stderr { request_id, session_id, data } + } + + pub fn completed(request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, exit_code: i32) -> Self { + Self::Completed { request_id, session_id, operation, exit_code } + } + + pub fn error( + request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, kind: WorkerErrorKind, message: impl Into, + ) -> Self { + Self::Error { request_id, session_id, operation, kind, message: message.into() } + } + + pub const fn kind(&self) -> WorkerFrameKind { + match self { + Self::Hello { .. } => WorkerFrameKind::Hello, + Self::UploadBegin { .. } => WorkerFrameKind::UploadBegin, + Self::UploadEntry { .. } => WorkerFrameKind::UploadEntry, + Self::UploadFileChunk { .. } => WorkerFrameKind::UploadFileChunk, + Self::UploadComplete { .. } => WorkerFrameKind::UploadComplete, + Self::Build { .. } => WorkerFrameKind::Build, + Self::Cleanup { .. } => WorkerFrameKind::Cleanup, + Self::SyncProgress { .. } => WorkerFrameKind::SyncProgress, + Self::Stdout { .. } => WorkerFrameKind::Stdout, + Self::Stderr { .. } => WorkerFrameKind::Stderr, + Self::Completed { .. } => WorkerFrameKind::Completed, + Self::Error { .. } => WorkerFrameKind::Error, + } + } + + pub fn request_id(&self) -> WorkerRequestId { + match self { + Self::Hello { request_id, .. } + | Self::UploadBegin { request_id, .. } + | Self::UploadEntry { request_id, .. } + | Self::UploadFileChunk { request_id, .. } + | Self::UploadComplete { request_id, .. } + | Self::Build { request_id, .. } + | Self::Cleanup { request_id, .. } + | Self::SyncProgress { request_id, .. } + | Self::Stdout { request_id, .. } + | Self::Stderr { request_id, .. } + | Self::Completed { request_id, .. } + | Self::Error { request_id, .. } => *request_id, + } + } + + pub fn session_id(&self) -> WorkerSessionId { + match self { + Self::Hello { session_id, .. } + | Self::UploadBegin { session_id, .. } + | Self::UploadEntry { session_id, .. } + | Self::UploadFileChunk { session_id, .. } + | Self::UploadComplete { session_id, .. } + | Self::Build { session_id, .. } + | Self::Cleanup { session_id, .. } + | Self::SyncProgress { session_id, .. } + | Self::Stdout { session_id, .. } + | Self::Stderr { session_id, .. } + | Self::Completed { session_id, .. } + | Self::Error { session_id, .. } => *session_id, + } + } + + pub fn upload_id(&self) -> Option { + match self { + Self::UploadBegin { upload_id, .. } + | Self::UploadEntry { upload_id, .. } + | Self::UploadFileChunk { upload_id, .. } + | Self::UploadComplete { upload_id, .. } + | Self::SyncProgress { upload_id, .. } => Some(*upload_id), + Self::Build { build, .. } => Some(build.upload_token), + Self::Cleanup { upload_token, .. } => Some(*upload_token), + Self::Hello { .. } | Self::Stdout { .. } | Self::Stderr { .. } | Self::Completed { .. } | Self::Error { .. } => None, + } + } + + pub fn validate(&self) -> WorkerResult<()> { + match self { + Self::Hello { version, .. } => { + if *version != WORKER_PROTOCOL_VERSION { + return Err(invalid(format!("unsupported worker hello version: {version}"))); + } + } + Self::UploadBegin { entries, .. } => { + validate_upload_manifest(entries)?; + } + Self::UploadEntry { entry_index, entry, .. } => { + validate_entry_index(*entry_index)?; + entry.validate()?; + } + Self::UploadFileChunk { path, offset, data, .. } => validate_chunk(path, *offset, data)?, + Self::UploadComplete { .. } => {} + Self::Build { build, .. } => build.validate()?, + Self::Cleanup { .. } => {} + Self::SyncProgress { completed_bytes, total_bytes, .. } => validate_progress(*completed_bytes, *total_bytes)?, + Self::Stdout { data, .. } | Self::Stderr { data, .. } => { + if data.len() > MAX_WORKER_OUTPUT_BYTES { + return Err(invalid(format!("worker output exceeds maximum length {MAX_WORKER_OUTPUT_BYTES}"))); + } + } + Self::Completed { operation, .. } => { + if *operation == WorkerOperation::Protocol { + return Err(invalid("worker completion cannot use protocol operation")); + } + } + Self::Error { message, .. } => validate_error_message(message)?, + } + Ok(()) + } + + pub fn encode(&self) -> WorkerResult> { + self.validate()?; + let mut payload = WireWriter::new(); + encode_payload(self, &mut payload)?; + let payload = payload.finish()?; + + let declared = u32::try_from(payload.len()).map_err(|_| invalid("worker payload length does not fit in u32"))?; + let mut frame = Vec::with_capacity(WORKER_FRAME_HEADER_LEN + payload.len()); + frame.extend_from_slice(&WORKER_PROTOCOL_MAGIC); + frame.extend_from_slice(&WORKER_PROTOCOL_VERSION.to_le_bytes()); + frame.push(self.kind().as_u8()); + frame.extend_from_slice(&declared.to_le_bytes()); + frame.extend_from_slice(&payload); + Ok(frame) + } + + pub fn decode(frame: &[u8]) -> WorkerResult { + let (kind, payload) = split_frame(frame)?; + decode_payload(kind, payload) + } + + #[cfg(feature = "async")] + pub async fn read_async(reader: &mut R) -> WorkerResult { + let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; + reader.read_exact(&mut header).await.map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + let (kind, payload_len) = decode_header(&header)?; + let mut payload = vec![0u8; payload_len]; + reader.read_exact(&mut payload).await.map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + decode_payload(kind, &payload) + } + + pub fn read_blocking(reader: &mut R) -> WorkerResult { + let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; + reader.read_exact(&mut header).map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + let (kind, payload_len) = decode_header(&header)?; + let mut payload = vec![0u8; payload_len]; + reader.read_exact(&mut payload).map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + decode_payload(kind, &payload) + } + + pub fn read_blocking_optional(reader: &mut R) -> WorkerResult> { + let mut first = [0u8; 1]; + match reader.read_exact(&mut first) { + Ok(()) => {} + Err(error) if error.kind() == io::ErrorKind::UnexpectedEof => return Ok(None), + Err(error) => return Err(WorkerProtocolError::Io(error.to_string())), + } + let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; + header[0] = first[0]; + reader.read_exact(&mut header[1..]).map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + let (kind, payload_len) = decode_header(&header)?; + let mut payload = vec![0u8; payload_len]; + reader.read_exact(&mut payload).map_err(|error| match error.kind() { + io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), + _ => WorkerProtocolError::Io(error.to_string()), + })?; + decode_payload(kind, &payload).map(Some) + } + + #[cfg(feature = "async")] + pub async fn write_async(&self, writer: &mut W) -> WorkerResult<()> { + let frame = self.encode()?; + writer.write_all(&frame).await.map_err(WorkerProtocolError::from)?; + writer.flush().await.map_err(WorkerProtocolError::from) + } + + pub fn write_blocking(&self, writer: &mut W) -> WorkerResult<()> { + let frame = self.encode()?; + writer.write_all(&frame).map_err(WorkerProtocolError::from)?; + writer.flush().map_err(WorkerProtocolError::from) + } +} + +#[cfg(feature = "async")] +pub async fn read_worker_message(reader: &mut R) -> WorkerResult { + WorkerMessage::read_async(reader).await +} + +#[cfg(feature = "async")] +pub async fn write_worker_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { + message.write_async(writer).await +} + +#[cfg(feature = "async")] +pub async fn read_message(reader: &mut R) -> WorkerResult { + read_worker_message(reader).await +} + +#[cfg(feature = "async")] +pub async fn write_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { + write_worker_message(writer, message).await +} + +pub fn encode_worker_message(message: &WorkerMessage) -> WorkerResult> { + message.encode() +} + +pub fn decode_worker_message(frame: &[u8]) -> WorkerResult { + WorkerMessage::decode(frame) +} + +pub fn validate_worker_relative_path(value: &str) -> WorkerResult<()> { + validate_relative_path("worker relative path", value, false) +} + +pub fn validate_worker_tool(value: &str) -> WorkerResult<()> { + validate_tool(value) +} + +pub fn validate_worker_executable_path(value: &str) -> WorkerResult<()> { + validate_executable_path(value) +} + +pub fn validate_upload_manifest(entries: &[WorkerUploadEntry]) -> WorkerResult { + validate_count(entries.len(), MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; + let mut previous: Option<&str> = None; + let mut total_bytes = 0u64; + let mut manifest_bytes = 4usize; + for entry in entries { + entry.validate()?; + if let Some(previous) = previous { + if previous >= entry.path.as_str() { + return Err(invalid("worker upload manifest must be strictly sorted by path")); + } + } + previous = Some(entry.path.as_str()); + let encoded_entry_bytes = 4usize + .checked_add(entry.path.as_str().len()) + .and_then(|bytes| bytes.checked_add(1 + 4 + 8 + 1)) + .and_then(|bytes| bytes.checked_add(if entry.digest.is_some() { WORKER_DIGEST_LEN } else { 0 })) + .ok_or_else(|| invalid("worker manifest length overflow"))?; + manifest_bytes = manifest_bytes.checked_add(encoded_entry_bytes).ok_or_else(|| invalid("worker manifest length overflow"))?; + if manifest_bytes > MAX_WORKER_MANIFEST_BYTES { + return Err(invalid(format!("worker manifest exceeds maximum length {MAX_WORKER_MANIFEST_BYTES}"))); + } + total_bytes = total_bytes.checked_add(entry.size).ok_or_else(|| invalid("worker upload byte count overflow"))?; + if total_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { + return Err(invalid(format!("worker upload exceeds maximum size {MAX_WORKER_TOTAL_UPLOAD_BYTES}"))); + } + } + Ok(total_bytes) +} + +fn encode_payload(message: &WorkerMessage, writer: &mut WireWriter) -> WorkerResult<()> { + match message { + WorkerMessage::Hello { request_id, session_id, version, response } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.u16(*version)?; + writer.boolean(*response)?; + } + WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.count(entries.len(), MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; + for entry in entries { + encode_entry(writer, entry)?; + } + } + WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.u32(*entry_index)?; + encode_entry(writer, entry)?; + } + WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.string(path.as_str(), MAX_WORKER_PATH_BYTES, "worker chunk path")?; + writer.u64(*offset)?; + writer.blob(data, MAX_WORKER_CHUNK_BYTES, "worker file chunk")?; + } + WorkerMessage::UploadComplete { request_id, session_id, upload_id } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + } + WorkerMessage::Build { request_id, session_id, build } => { + encode_correlation(writer, *request_id, *session_id)?; + encode_build(writer, build)?; + } + WorkerMessage::Cleanup { request_id, session_id, upload_token } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_token.0)?; + } + WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes, total_bytes } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(upload_id.0)?; + writer.u64(*completed_bytes)?; + match total_bytes { + Some(total_bytes) => { + writer.boolean(true)?; + writer.u64(*total_bytes)?; + } + None => writer.boolean(false)?, + } + } + WorkerMessage::Stdout { request_id, session_id, data } | WorkerMessage::Stderr { request_id, session_id, data } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.blob(data, MAX_WORKER_OUTPUT_BYTES, "worker output")?; + } + WorkerMessage::Completed { request_id, session_id, operation, exit_code } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.u8(operation.as_u8())?; + writer.i32(*exit_code)?; + } + WorkerMessage::Error { request_id, session_id, operation, kind, message } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.u8(operation.as_u8())?; + writer.u8(kind.as_u8())?; + writer.string(message, MAX_WORKER_ERROR_BYTES, "worker error")?; + } + } + Ok(()) +} + +fn decode_payload(kind: WorkerFrameKind, payload: &[u8]) -> WorkerResult { + let mut reader = WireReader::new(payload); + let message = match kind { + WorkerFrameKind::Hello => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let version = reader.u16()?; + let response = reader.boolean("worker hello response")?; + WorkerMessage::Hello { request_id, session_id, version, response } + } + WorkerFrameKind::UploadBegin => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let count = reader.count(MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; + let mut entries = Vec::with_capacity(count); + for _ in 0..count { + entries.push(decode_entry(&mut reader)?); + } + WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } + } + WorkerFrameKind::UploadEntry => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let entry_index = reader.u32()?; + let entry = decode_entry(&mut reader)?; + WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } + } + WorkerFrameKind::UploadFileChunk => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let path = WorkerRelativePath::for_entry(reader.string(MAX_WORKER_PATH_BYTES, "worker chunk path")?)?; + let offset = reader.u64()?; + let data = reader.blob(MAX_WORKER_CHUNK_BYTES, "worker file chunk")?; + WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } + } + WorkerFrameKind::UploadComplete => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + WorkerMessage::UploadComplete { request_id, session_id, upload_id } + } + WorkerFrameKind::Build => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + WorkerMessage::Build { request_id, session_id, build: decode_build(&mut reader)? } + } + WorkerFrameKind::Cleanup => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_token = WorkerUploadId(reader.array16()?); + WorkerMessage::Cleanup { request_id, session_id, upload_token } + } + WorkerFrameKind::SyncProgress => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let upload_id = WorkerUploadId(reader.array16()?); + let completed_bytes = reader.u64()?; + let total_bytes = if reader.boolean("worker progress total flag")? { Some(reader.u64()?) } else { None }; + WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes, total_bytes } + } + WorkerFrameKind::Stdout => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + WorkerMessage::Stdout { request_id, session_id, data: reader.blob(MAX_WORKER_OUTPUT_BYTES, "worker stdout")? } + } + WorkerFrameKind::Stderr => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + WorkerMessage::Stderr { request_id, session_id, data: reader.blob(MAX_WORKER_OUTPUT_BYTES, "worker stderr")? } + } + WorkerFrameKind::Completed => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let operation = WorkerOperation::from_u8(reader.u8()?)?; + let exit_code = reader.i32()?; + WorkerMessage::Completed { request_id, session_id, operation, exit_code } + } + WorkerFrameKind::Error => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let operation = WorkerOperation::from_u8(reader.u8()?)?; + let kind = WorkerErrorKind::from_u8(reader.u8()?)?; + let message = reader.string(MAX_WORKER_ERROR_BYTES, "worker error")?; + WorkerMessage::Error { request_id, session_id, operation, kind, message } + } + }; + reader.finish()?; + message.validate()?; + Ok(message) +} + +fn encode_build(writer: &mut WireWriter, build: &WorkerBuild) -> WorkerResult<()> { + build.validate()?; + writer.string(build.tool.as_str(), MAX_WORKER_TOOL_BYTES, "worker tool")?; + writer.string(build.trusted_executable.as_str(), MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?; + writer.count(build.argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; + for argument in &build.argv { + writer.string(argument, MAX_WORKER_ARG_BYTES, "worker argument")?; + } + writer.string(build.cwd.as_str(), MAX_WORKER_CWD_BYTES, "worker cwd")?; + encode_environment(writer, &build.guest_env, "worker guest environment")?; + encode_environment(writer, &build.target_env, "worker target environment")?; + writer.id(build.upload_token.0)?; + Ok(()) +} + +fn decode_build(reader: &mut WireReader<'_>) -> WorkerResult { + let tool = WorkerTool::new(reader.string(MAX_WORKER_TOOL_BYTES, "worker tool")?)?; + let trusted_executable = WorkerExecutablePath::new(reader.string(MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?)?; + let argument_count = reader.count(MAX_WORKER_ARG_COUNT, "worker argv")?; + let mut argv = Vec::with_capacity(argument_count); + for _ in 0..argument_count { + argv.push(reader.string(MAX_WORKER_ARG_BYTES, "worker argument")?); + } + let cwd = WorkerRelativePath::new(reader.string(MAX_WORKER_CWD_BYTES, "worker cwd")?)?; + let guest_env = decode_environment(reader, "worker guest environment")?; + let target_env = decode_environment(reader, "worker target environment")?; + let upload_token = WorkerUploadId(reader.array16()?); + let build = WorkerBuild { tool, trusted_executable, argv, cwd, guest_env, target_env, upload_token }; + build.validate()?; + Ok(build) +} + +fn encode_environment(writer: &mut WireWriter, environment: &[(String, String)], field: &str) -> WorkerResult<()> { + validate_environment(field, environment)?; + writer.count(environment.len(), MAX_WORKER_ENV_COUNT, field)?; + for (key, value) in environment { + writer.string(key, MAX_WORKER_ENV_KEY_BYTES, "worker environment key")?; + writer.string(value, MAX_WORKER_ENV_VALUE_BYTES, "worker environment value")?; + } + Ok(()) +} + +fn decode_environment(reader: &mut WireReader<'_>, field: &str) -> WorkerResult> { + let count = reader.count(MAX_WORKER_ENV_COUNT, field)?; + let mut environment = Vec::with_capacity(count); + for _ in 0..count { + let key = reader.string(MAX_WORKER_ENV_KEY_BYTES, "worker environment key")?; + let value = reader.string(MAX_WORKER_ENV_VALUE_BYTES, "worker environment value")?; + environment.push((key, value)); + } + validate_environment(field, &environment)?; + Ok(environment) +} + +fn encode_entry(writer: &mut WireWriter, entry: &WorkerUploadEntry) -> WorkerResult<()> { + entry.validate()?; + writer.string(entry.path.as_str(), MAX_WORKER_PATH_BYTES, "worker entry path")?; + writer.u8(entry.kind.as_u8())?; + writer.u32(entry.mode)?; + writer.u64(entry.size)?; + match entry.digest { + Some(digest) => { + writer.boolean(true)?; + writer.bytes(&digest)?; + } + None => writer.boolean(false)?, + } + Ok(()) +} + +fn decode_entry(reader: &mut WireReader<'_>) -> WorkerResult { + let path = reader.string(MAX_WORKER_PATH_BYTES, "worker entry path")?; + let kind = WorkerEntryKind::from_u8(reader.u8()?)?; + let mode = reader.u32()?; + let size = reader.u64()?; + let digest = match reader.boolean("worker entry digest flag")? { + true => Some(reader.array32()?), + false => None, + }; + WorkerUploadEntry::new(path, kind, mode, size, digest) +} + +fn encode_correlation(writer: &mut WireWriter, request_id: WorkerRequestId, session_id: WorkerSessionId) -> WorkerResult<()> { + writer.id(request_id.0)?; + writer.id(session_id.0) +} + +fn decode_correlation(reader: &mut WireReader<'_>) -> WorkerResult<(WorkerRequestId, WorkerSessionId)> { + Ok((WorkerRequestId(reader.array16()?), WorkerSessionId(reader.array16()?))) +} + +fn validate_entry_index(index: u32) -> WorkerResult<()> { + if usize::try_from(index).map_or(true, |index| index >= MAX_WORKER_UPLOAD_ENTRIES) { + return Err(invalid(format!("worker upload entry index exceeds maximum {MAX_WORKER_UPLOAD_ENTRIES}"))); + } + Ok(()) +} + +fn validate_chunk(path: &WorkerRelativePath, offset: u64, data: &[u8]) -> WorkerResult<()> { + validate_relative_path("worker chunk path", path.as_str(), false)?; + if data.is_empty() { + return Err(invalid("worker file chunk must not be empty")); + } + if data.len() > MAX_WORKER_CHUNK_BYTES { + return Err(invalid(format!("worker file chunk exceeds maximum length {MAX_WORKER_CHUNK_BYTES}"))); + } + let end = offset + .checked_add(u64::try_from(data.len()).map_err(|_| invalid("worker file chunk length does not fit in u64"))?) + .ok_or_else(|| invalid("worker file chunk offset overflow"))?; + if end > MAX_WORKER_FILE_BYTES { + return Err(invalid(format!("worker file chunk exceeds maximum file size {MAX_WORKER_FILE_BYTES}"))); + } + Ok(()) +} + +fn validate_progress(completed_bytes: u64, total_bytes: Option) -> WorkerResult<()> { + if completed_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { + return Err(invalid("worker progress exceeds the maximum upload size")); + } + if let Some(total_bytes) = total_bytes { + if total_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { + return Err(invalid("worker progress total exceeds the maximum upload size")); + } + if completed_bytes > total_bytes { + return Err(invalid("worker progress exceeds its total")); + } + } + Ok(()) +} + +fn validate_argv(argv: &[String]) -> WorkerResult<()> { + validate_count(argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; + let mut total_bytes = 0usize; + for argument in argv { + validate_text("worker argument", argument, MAX_WORKER_ARG_BYTES)?; + total_bytes = total_bytes.checked_add(argument.len()).ok_or_else(|| invalid("worker argv length overflow"))?; + if total_bytes > MAX_WORKER_ARG_TOTAL_BYTES { + return Err(invalid(format!("worker argv exceeds maximum length {MAX_WORKER_ARG_TOTAL_BYTES}"))); + } + } + Ok(()) +} + +fn validate_environment(field: &str, environment: &[(String, String)]) -> WorkerResult<()> { + validate_count(environment.len(), MAX_WORKER_ENV_COUNT, field)?; + let mut names = BTreeSet::new(); + let mut total_bytes = 0usize; + for (key, value) in environment { + validate_environment_key(key)?; + validate_environment_value(value)?; + total_bytes = total_bytes + .checked_add(key.len()) + .and_then(|bytes| bytes.checked_add(value.len())) + .ok_or_else(|| invalid(format!("{field} length overflow")))?; + if total_bytes > MAX_WORKER_ENV_TOTAL_BYTES { + return Err(invalid(format!("{field} exceeds maximum length {MAX_WORKER_ENV_TOTAL_BYTES}"))); + } + if !names.insert(key.as_str()) { + return Err(invalid(format!("duplicate worker environment key: {key}"))); + } + } + Ok(()) +} + +fn validate_environment_key(key: &str) -> WorkerResult<()> { + validate_text("worker environment key", key, MAX_WORKER_ENV_KEY_BYTES)?; + let mut bytes = key.bytes(); + let Some(first) = bytes.next() else { + return Err(invalid("worker environment key is empty")); + }; + if !(first == b'_' || first.is_ascii_alphabetic()) || !bytes.all(|byte| byte == b'_' || byte.is_ascii_alphanumeric()) { + return Err(invalid("worker environment key is not a valid variable name")); + } + Ok(()) +} + +fn validate_environment_value(value: &str) -> WorkerResult<()> { + validate_text("worker environment value", value, MAX_WORKER_ENV_VALUE_BYTES)?; + if value.chars().any(char::is_control) { + return Err(invalid("worker environment value contains control data")); + } + Ok(()) +} + +fn validate_error_message(message: &str) -> WorkerResult<()> { + validate_text("worker error", message, MAX_WORKER_ERROR_BYTES) +} + +fn validate_tool(value: &str) -> WorkerResult<()> { + validate_text("worker tool", value, MAX_WORKER_TOOL_BYTES)?; + if value.is_empty() + || value == "." + || value == ".." + || !value.bytes().all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b'+')) + { + return Err(invalid("worker tool must be a single safe executable identity")); + } + Ok(()) +} + +fn validate_executable_path(value: &str) -> WorkerResult<()> { + validate_text("worker executable path", value, MAX_WORKER_EXECUTABLE_PATH_BYTES)?; + if !value.starts_with('/') || value == "/" || value.starts_with("//") { + return Err(invalid("worker executable path must be a normalized absolute path")); + } + if value.chars().any(|character| character.is_whitespace() || character.is_control()) { + return Err(invalid("worker executable path must not contain whitespace or control data")); + } + if value.bytes().any(|byte| { + byte < 0x20 + || byte == 0x7f + || matches!(byte, b'\\' | b';' | b'|' | b'&' | b'$' | b'`' | b'<' | b'>' | b'\'' | b'"' | b'(' | b')' | b'[' | b']' | b'{' | b'}') + }) { + return Err(invalid("worker executable path contains unsafe identity data")); + } + for (depth, component) in value.split('/').skip(1).enumerate() { + if depth >= MAX_WORKER_PATH_DEPTH + || component.len() > MAX_WORKER_PATH_COMPONENT_BYTES + || component.is_empty() + || component == "." + || component == ".." + || component.contains(':') + { + return Err(invalid("worker executable path is not normalized")); + } + } + Ok(()) +} + +fn validate_relative_path(field: &str, value: &str, allow_empty: bool) -> WorkerResult<()> { + validate_text(field, value, MAX_WORKER_PATH_BYTES)?; + if value.is_empty() { + if allow_empty { + return Ok(()); + } + return Err(invalid(format!("{field} must not be empty"))); + } + if value.starts_with('/') || value.starts_with('\\') || value.contains('\\') || value.bytes().any(|byte| byte == b':') { + return Err(invalid(format!("{field} must be a normalized relative path"))); + } + for (depth, component) in value.split('/').enumerate() { + if depth >= MAX_WORKER_PATH_DEPTH + || component.len() > MAX_WORKER_PATH_COMPONENT_BYTES + || component.is_empty() + || component == "." + || component == ".." + || component.chars().any(char::is_control) + { + return Err(invalid(format!("{field} must be a normalized relative path"))); + } + } + Ok(()) +} + +fn validate_text(field: &str, value: &str, maximum: usize) -> WorkerResult<()> { + if value.len() > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + if value.as_bytes().contains(&0) { + return Err(invalid(format!("{field} contains a NUL byte"))); + } + Ok(()) +} + +fn validate_count(count: usize, maximum: usize, field: &str) -> WorkerResult<()> { + if count > maximum { + return Err(invalid(format!("{field} exceeds maximum count {maximum}"))); + } + Ok(()) +} + +fn invalid(message: impl Into) -> WorkerProtocolError { + WorkerProtocolError::Invalid(message.into()) +} + +fn split_frame(frame: &[u8]) -> WorkerResult<(WorkerFrameKind, &[u8])> { + if frame.len() < WORKER_FRAME_HEADER_LEN { + return Err(invalid("truncated worker frame header")); + } + let (kind, payload_len) = decode_header(&frame[..WORKER_FRAME_HEADER_LEN])?; + let expected = WORKER_FRAME_HEADER_LEN.checked_add(payload_len).ok_or_else(|| invalid("worker frame length overflow"))?; + if frame.len() < expected { + return Err(invalid("truncated worker frame payload")); + } + if frame.len() > expected { + return Err(invalid("extra bytes after worker frame")); + } + Ok((kind, &frame[WORKER_FRAME_HEADER_LEN..expected])) +} + +fn decode_header(header: &[u8]) -> WorkerResult<(WorkerFrameKind, usize)> { + if header.len() != WORKER_FRAME_HEADER_LEN { + return Err(invalid("invalid worker frame header length")); + } + if header[..4] != WORKER_PROTOCOL_MAGIC { + return Err(invalid("invalid worker frame magic")); + } + let version = u16::from_le_bytes([header[4], header[5]]); + if version != WORKER_PROTOCOL_VERSION { + return Err(invalid(format!("unsupported worker protocol version: {version}"))); + } + let kind = WorkerFrameKind::try_from(header[6])?; + let payload_len = u32::from_le_bytes([header[7], header[8], header[9], header[10]]) as usize; + if payload_len > MAX_WORKER_FRAME_PAYLOAD { + return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); + } + Ok((kind, payload_len)) +} + +struct WireWriter { + bytes: Vec, +} + +impl WireWriter { + fn new() -> Self { + Self { bytes: Vec::new() } + } + + fn finish(self) -> WorkerResult> { + if self.bytes.len() > MAX_WORKER_FRAME_PAYLOAD { + return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); + } + Ok(self.bytes) + } + + fn bytes(&mut self, value: &[u8]) -> WorkerResult<()> { + let new_length = self.bytes.len().checked_add(value.len()).ok_or_else(|| invalid("worker payload length overflow"))?; + if new_length > MAX_WORKER_FRAME_PAYLOAD { + return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); + } + self.bytes.extend_from_slice(value); + Ok(()) + } + + fn id(&mut self, value: [u8; WORKER_ID_LEN]) -> WorkerResult<()> { + self.bytes(&value) + } + + fn u8(&mut self, value: u8) -> WorkerResult<()> { + self.bytes(&[value]) + } + + fn boolean(&mut self, value: bool) -> WorkerResult<()> { + self.u8(u8::from(value)) + } + + fn u16(&mut self, value: u16) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn u32(&mut self, value: u32) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn u64(&mut self, value: u64) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn i32(&mut self, value: i32) -> WorkerResult<()> { + self.bytes(&value.to_le_bytes()) + } + + fn count(&mut self, count: usize, maximum: usize, field: &str) -> WorkerResult<()> { + validate_count(count, maximum, field)?; + self.u32(u32::try_from(count).map_err(|_| invalid(format!("{field} count does not fit in u32")))?) + } + + fn string(&mut self, value: &str, maximum: usize, field: &str) -> WorkerResult<()> { + validate_text(field, value, maximum)?; + let length = u32::try_from(value.len()).map_err(|_| invalid(format!("{field} length does not fit in u32")))?; + self.u32(length)?; + self.bytes(value.as_bytes()) + } + + fn blob(&mut self, value: &[u8], maximum: usize, field: &str) -> WorkerResult<()> { + if value.len() > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + let length = u32::try_from(value.len()).map_err(|_| invalid(format!("{field} length does not fit in u32")))?; + self.u32(length)?; + self.bytes(value) + } +} + +struct WireReader<'a> { + bytes: &'a [u8], + offset: usize, +} + +impl<'a> WireReader<'a> { + fn new(bytes: &'a [u8]) -> Self { + Self { bytes, offset: 0 } + } + + fn take(&mut self, length: usize) -> WorkerResult<&'a [u8]> { + let end = self.offset.checked_add(length).ok_or_else(|| invalid("worker payload length overflow"))?; + if end > self.bytes.len() { + return Err(invalid("truncated worker payload")); + } + let value = &self.bytes[self.offset..end]; + self.offset = end; + Ok(value) + } + + fn u8(&mut self) -> WorkerResult { + Ok(self.take(1)?[0]) + } + + fn boolean(&mut self, field: &str) -> WorkerResult { + match self.u8()? { + 0 => Ok(false), + 1 => Ok(true), + value => Err(invalid(format!("invalid {field} flag: {value}"))), + } + } + + fn u16(&mut self) -> WorkerResult { + let value = self.take(2)?; + Ok(u16::from_le_bytes([value[0], value[1]])) + } + + fn u32(&mut self) -> WorkerResult { + let value = self.take(4)?; + Ok(u32::from_le_bytes([value[0], value[1], value[2], value[3]])) + } + + fn u64(&mut self) -> WorkerResult { + let value = self.take(8)?; + Ok(u64::from_le_bytes(value.try_into().map_err(|_| invalid("invalid worker integer"))?)) + } + + fn i32(&mut self) -> WorkerResult { + let value = self.take(4)?; + Ok(i32::from_le_bytes(value.try_into().map_err(|_| invalid("invalid worker integer"))?)) + } + + fn array16(&mut self) -> WorkerResult<[u8; WORKER_ID_LEN]> { + self.take(WORKER_ID_LEN)?.try_into().map_err(|_| invalid("invalid worker identifier")) + } + + fn array32(&mut self) -> WorkerResult<[u8; WORKER_DIGEST_LEN]> { + self.take(WORKER_DIGEST_LEN)?.try_into().map_err(|_| invalid("invalid worker digest")) + } + + fn count(&mut self, maximum: usize, field: &str) -> WorkerResult { + let count = self.u32()? as usize; + validate_count(count, maximum, field)?; + Ok(count) + } + + fn string(&mut self, maximum: usize, field: &str) -> WorkerResult { + let length = self.u32()? as usize; + if length > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + let value = std::str::from_utf8(self.take(length)?).map_err(|_| invalid(format!("{field} is not valid UTF-8")))?; + validate_text(field, value, maximum)?; + Ok(value.to_owned()) + } + + fn blob(&mut self, maximum: usize, field: &str) -> WorkerResult> { + let length = self.u32()? as usize; + if length > maximum { + return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); + } + Ok(self.take(length)?.to_vec()) + } + + fn finish(self) -> WorkerResult<()> { + if self.offset != self.bytes.len() { + return Err(invalid("extra bytes in worker payload")); + } + Ok(()) + } +} + +#[cfg(test)] +#[path = "lib_ut.rs"] +mod worker_protocol_tests; diff --git a/crates/bunkerbox-worker-protocol/src/lib_ut.rs b/crates/bunkerbox-worker-protocol/src/lib_ut.rs new file mode 100644 index 0000000..00d1bcd --- /dev/null +++ b/crates/bunkerbox-worker-protocol/src/lib_ut.rs @@ -0,0 +1,299 @@ +use super::*; +use std::io::{self, Cursor, Read}; +use tokio::io::AsyncWriteExt; + +fn ids() -> (WorkerRequestId, WorkerSessionId, WorkerUploadId) { + (WorkerRequestId([1; 16]), WorkerSessionId([2; 16]), WorkerUploadId([3; 16])) +} + +fn file_entry(path: &str) -> WorkerUploadEntry { + WorkerUploadEntry::file(path, 0o644, 4, [9; WORKER_DIGEST_LEN]).unwrap() +} + +fn build() -> WorkerBuild { + WorkerBuild::new( + "make", + "/usr/bin/make", + vec!["release mode".into(), "$(literal); still one arg".into()], + "src", + vec![("CC".into(), "gcc".into())], + vec![("PATH".into(), "/usr/bin".into())], + WorkerUploadId([3; 16]), + ) + .unwrap() +} + +fn messages() -> Vec { + let (request_id, session_id, upload_id) = ids(); + vec![ + WorkerMessage::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response: false }, + WorkerMessage::UploadBegin { + request_id, + session_id, + upload_id, + entries: vec![WorkerUploadEntry::directory("src", 0o755).unwrap(), file_entry("src/main.rs")], + }, + WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index: 1, entry: file_entry("src/main.rs") }, + WorkerMessage::UploadFileChunk { + request_id, + session_id, + upload_id, + path: WorkerRelativePath::new("src/main.rs").unwrap(), + offset: 0, + data: b"data".to_vec(), + }, + WorkerMessage::UploadComplete { request_id, session_id, upload_id }, + WorkerMessage::Build { request_id, session_id, build: build() }, + WorkerMessage::Cleanup { request_id, session_id, upload_token: upload_id }, + WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 4, total_bytes: Some(8) }, + WorkerMessage::Stdout { request_id, session_id, data: b"out".to_vec() }, + WorkerMessage::Stderr { request_id, session_id, data: b"err".to_vec() }, + WorkerMessage::Completed { request_id, session_id, operation: WorkerOperation::Build, exit_code: -17 }, + WorkerMessage::Error { + request_id, + session_id, + operation: WorkerOperation::Cleanup, + kind: WorkerErrorKind::Cleanup, + message: "cleanup failed".into(), + }, + ] +} + +#[test] +fn all_message_variants_round_trip() { + for message in messages() { + let frame = message.encode().unwrap(); + assert_eq!(&frame[..4], &WORKER_PROTOCOL_MAGIC); + assert_eq!(frame[4..6], WORKER_PROTOCOL_VERSION.to_le_bytes()); + assert_eq!(frame[6], message.kind().as_u8()); + assert_eq!(WorkerMessage::decode(&frame).unwrap(), message); + } +} + +#[test] +fn blocking_helpers_preserve_the_async_wire_encoding() { + let message = WorkerMessage::Stdout { request_id: ids().0, session_id: ids().1, data: b"blocking parity".to_vec() }; + let mut encoded = Vec::new(); + message.write_blocking(&mut encoded).unwrap(); + assert_eq!(WorkerMessage::decode(&encoded).unwrap(), message); + assert_eq!(WorkerMessage::read_blocking(&mut Cursor::new(encoded)).unwrap(), message); +} + +#[test] +fn blocking_reader_handles_fragmented_frames_and_clean_eof() { + let message = WorkerMessage::Stderr { request_id: ids().0, session_id: ids().1, data: b"fragmented".to_vec() }; + let reader = OneByteReader { bytes: message.encode().unwrap(), offset: 0 }; + let mut reader = reader; + assert_eq!(WorkerMessage::read_blocking_optional(&mut reader).unwrap(), Some(message)); + assert_eq!(WorkerMessage::read_blocking_optional(&mut reader).unwrap(), None); +} + +struct OneByteReader { + bytes: Vec, + offset: usize, +} + +impl Read for OneByteReader { + fn read(&mut self, buffer: &mut [u8]) -> io::Result { + if self.offset == self.bytes.len() { + return Ok(0); + } + buffer[0] = self.bytes[self.offset]; + self.offset += 1; + Ok(1) + } +} + +#[tokio::test] +async fn async_helpers_handle_fragmented_duplex_frames() { + let message = WorkerMessage::Build { request_id: ids().0, session_id: ids().1, build: build() }; + let frame = message.encode().unwrap(); + let (mut reader, mut writer) = tokio::io::duplex(3); + let sender = tokio::spawn(async move { + for part in frame.chunks(2) { + writer.write_all(part).await.unwrap(); + tokio::task::yield_now().await; + } + }); + let decoded = WorkerMessage::read_async(&mut reader).await.unwrap(); + sender.await.unwrap(); + assert_eq!(decoded, message); +} + +#[tokio::test] +async fn async_write_helper_emits_a_decodable_frame() { + let message = WorkerMessage::Stdout { request_id: ids().0, session_id: ids().1, data: b"exact".to_vec() }; + let (mut reader, mut writer) = tokio::io::duplex(128); + let expected = message.clone(); + let sender = tokio::spawn(async move { message.write_async(&mut writer).await.unwrap() }); + assert_eq!(WorkerMessage::read_async(&mut reader).await.unwrap(), expected); + sender.await.unwrap(); +} + +#[test] +fn malformed_headers_and_lengths_are_rejected() { + let frame = + WorkerMessage::Hello { request_id: ids().0, session_id: ids().1, version: WORKER_PROTOCOL_VERSION, response: false }.encode().unwrap(); + + let mut bad_magic = frame.clone(); + bad_magic[0] ^= 1; + assert!(WorkerMessage::decode(&bad_magic).is_err()); + + let mut bad_version = frame.clone(); + bad_version[4..6].copy_from_slice(&(WORKER_PROTOCOL_VERSION + 1).to_le_bytes()); + assert!(WorkerMessage::decode(&bad_version).is_err()); + + let mut bad_kind = frame.clone(); + bad_kind[6] = 255; + assert!(WorkerMessage::decode(&bad_kind).is_err()); + + let mut oversized = frame[..WORKER_FRAME_HEADER_LEN].to_vec(); + oversized[7..11].copy_from_slice(&((MAX_WORKER_FRAME_PAYLOAD as u32) + 1).to_le_bytes()); + assert!(WorkerMessage::decode(&oversized).is_err()); + + assert!(WorkerMessage::decode(&frame[..frame.len() - 1]).is_err()); + let mut extra = frame.clone(); + extra.push(0); + assert!(WorkerMessage::decode(&extra).is_err()); +} + +#[test] +fn invalid_utf8_and_trailing_payload_are_rejected() { + let (request_id, session_id, _) = ids(); + let mut frame = + WorkerMessage::Error { request_id, session_id, operation: WorkerOperation::Build, kind: WorkerErrorKind::Build, message: "x".into() } + .encode() + .unwrap(); + let message_start = WORKER_FRAME_HEADER_LEN + 32 + 1 + 1 + 4; + frame[message_start] = 0xff; + assert!(WorkerMessage::decode(&frame).is_err()); + + let mut hello = WorkerMessage::hello(request_id, session_id, false).encode().unwrap(); + hello.push(0); + assert!(WorkerMessage::decode(&hello).is_err()); +} + +#[test] +fn paths_tools_and_executables_are_strictly_validated() { + for path in ["/absolute", "../parent", "a/../b", "./name", "a//b", "a\\b", ""] { + assert!(validate_worker_relative_path(path).is_err(), "accepted path {path:?}"); + } + assert!(validate_worker_relative_path("src/main.rs").is_ok()); + assert!(validate_worker_relative_path(&"a".repeat(MAX_WORKER_PATH_COMPONENT_BYTES + 1)).is_err()); + assert!(WorkerTool::new("make;rm").is_err()); + assert!(WorkerTool::new("/usr/bin/make").is_err()); + assert!(WorkerExecutablePath::new("make").is_err()); + assert!(WorkerExecutablePath::new("/usr/bin/make -f").is_err()); + assert!(WorkerExecutablePath::new("/usr/bin/../bin/make").is_err()); + assert!(WorkerExecutablePath::new("/usr/bin/make").is_ok()); +} + +#[test] +fn manifests_require_sorted_paths_and_valid_metadata() { + assert!(validate_upload_manifest(&[file_entry("z"), file_entry("a")]).is_err()); + assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, None).is_err()); + assert!(WorkerUploadEntry::new("dir", WorkerEntryKind::Directory, 0o755, 1, None).is_err()); + assert!(WorkerUploadEntry::new("dir", WorkerEntryKind::Directory, 0o755, 0, Some([1; 32])).is_err()); + assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o100000, 1, Some([1; 32])).is_err()); + assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, Some([1; 32])).is_ok()); +} + +#[test] +fn duplicate_environment_keys_and_bad_argv_are_rejected() { + assert!(WorkerBuild::new( + "make", + "/usr/bin/make", + vec!["x".into()], + "", + vec![("CC".into(), "one".into()), ("CC".into(), "two".into())], + Vec::new(), + ids().2, + ) + .is_err()); + + assert!(WorkerBuild::new("make", "/usr/bin/make", vec!["x\0y".into()], "", Vec::new(), Vec::new(), ids().2,).is_err()); + + assert!(WorkerBuild::new("make", "/usr/bin/make", Vec::new(), "", vec![("BAD-NAME".into(), "value".into())], Vec::new(), ids().2,).is_err()); + + assert!(WorkerBuild::new("make", "/usr/bin/make", vec!["arg".into(); MAX_WORKER_ARG_COUNT + 1], "", Vec::new(), Vec::new(), ids().2,).is_err()); + + assert!(WorkerBuild::new( + "make", + "/usr/bin/make", + Vec::new(), + "", + vec![("A".into(), "value".into()); MAX_WORKER_ENV_COUNT + 1], + Vec::new(), + ids().2, + ) + .is_err()); +} + +#[test] +fn bounded_chunks_output_errors_and_progress_are_rejected() { + let (request_id, session_id, upload_id) = ids(); + let path = WorkerRelativePath::new("file").unwrap(); + assert!(WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path: path.clone(), offset: 0, data: Vec::new() }.encode().is_err()); + assert!(WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset: 0, data: vec![0; MAX_WORKER_CHUNK_BYTES + 1] } + .encode() + .is_err()); + assert!(WorkerMessage::Stdout { request_id, session_id, data: vec![0; MAX_WORKER_OUTPUT_BYTES + 1] }.encode().is_err()); + assert!(WorkerMessage::Error { + request_id, + session_id, + operation: WorkerOperation::Build, + kind: WorkerErrorKind::Build, + message: "x".repeat(MAX_WORKER_ERROR_BYTES + 1), + } + .encode() + .is_err()); + assert!(WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 2, total_bytes: Some(1) }.encode().is_err()); +} + +#[test] +fn stdout_stderr_and_exact_completion_preserve_bytes_and_exit_code() { + let (request_id, session_id, _) = ids(); + let stdout = WorkerMessage::Stdout { request_id, session_id, data: b"out\0with bytes".to_vec() }; + let stderr = WorkerMessage::Stderr { request_id, session_id, data: b"err".to_vec() }; + let completed = WorkerMessage::Completed { request_id, session_id, operation: WorkerOperation::Build, exit_code: i32::MIN }; + assert_eq!(WorkerMessage::decode(&stdout.encode().unwrap()).unwrap(), stdout); + assert_eq!(WorkerMessage::decode(&stderr.encode().unwrap()).unwrap(), stderr); + assert_eq!(WorkerMessage::decode(&completed.encode().unwrap()).unwrap(), completed); +} + +#[test] +fn build_has_no_shell_packing_and_keeps_literal_arguments() { + let message = WorkerMessage::Build { request_id: ids().0, session_id: ids().1, build: build() }; + let decoded = WorkerMessage::decode(&message.encode().unwrap()).unwrap(); + let WorkerMessage::Build { build, .. } = decoded else { panic!("expected build") }; + assert_eq!(build.argv, vec!["release mode", "$(literal); still one arg"]); + assert_eq!(build.argv.len(), 2); + assert_eq!(build.tool.as_str(), "make"); + assert_eq!(build.trusted_executable.as_str(), "/usr/bin/make"); +} + +#[test] +fn error_kinds_preserve_protocol_build_and_cleanup_failures() { + let (request_id, session_id, _) = ids(); + for (operation, kind) in [ + (WorkerOperation::Protocol, WorkerErrorKind::WorkerProtocol), + (WorkerOperation::Build, WorkerErrorKind::Build), + (WorkerOperation::Cleanup, WorkerErrorKind::Cleanup), + ] { + let message = WorkerMessage::Error { request_id, session_id, operation, kind, message: "failure".into() }; + assert_eq!(WorkerMessage::decode(&message.encode().unwrap()).unwrap(), message); + } +} + +#[test] +fn unknown_nested_kinds_and_flags_are_rejected() { + let message = WorkerMessage::UploadBegin { request_id: ids().0, session_id: ids().1, upload_id: ids().2, entries: Vec::new() }; + let mut frame = message.encode().unwrap(); + frame[WORKER_FRAME_HEADER_LEN + 48..WORKER_FRAME_HEADER_LEN + 52].copy_from_slice(&u32::MAX.to_le_bytes()); + assert!(WorkerMessage::decode(&frame).is_err()); + + let hello = WorkerMessage::hello(ids().0, ids().1, false); + let mut frame = hello.encode().unwrap(); + frame[WORKER_FRAME_HEADER_LEN + 34] = 9; + assert!(WorkerMessage::decode(&frame).is_err()); +} diff --git a/crates/bunkerbox-worker/Cargo.toml b/crates/bunkerbox-worker/Cargo.toml new file mode 100644 index 0000000..e60e28f --- /dev/null +++ b/crates/bunkerbox-worker/Cargo.toml @@ -0,0 +1,16 @@ +[package] +name = "bunkerbox-worker" +version = "0.1.0" +edition = "2021" + +[[bin]] +name = "bunkerbox-worker" +path = "src/main.rs" + +[dependencies] +bunkerbox-worker-protocol = { path = "../bunkerbox-worker-protocol", default-features = false } +libc = "0.2" +sha2 = "0.10" + +[dev-dependencies] +tempfile = "3" diff --git a/crates/bunkerbox-worker/src/main.rs b/crates/bunkerbox-worker/src/main.rs new file mode 100644 index 0000000..030ba64 --- /dev/null +++ b/crates/bunkerbox-worker/src/main.rs @@ -0,0 +1,66 @@ +mod platform; +mod process; +mod storage; +mod worker; + +use std::path::PathBuf; + +fn main() { + match parse_args(std::env::args().skip(1).collect()) { + Ok(root) => { + if let Err(error) = worker::run_stdio(&root) { + write_diagnostic(&error); + std::process::exit(70); + } + } + Err(error) => { + write_diagnostic(&error); + std::process::exit(64); + } + } +} + +fn parse_args(args: Vec) -> Result { + let mut stdio = false; + let mut root = None; + let mut index = 0; + while index < args.len() { + match args[index].as_str() { + "--stdio" => { + if stdio { + return Err("duplicate --stdio".to_string()); + } + stdio = true; + index += 1; + } + "--workspace-root" => { + if root.is_some() { + return Err("duplicate --workspace-root".to_string()); + } + let value = args.get(index + 1).ok_or_else(|| "--workspace-root requires a path".to_string())?; + if value.is_empty() || !std::path::Path::new(value).is_absolute() { + return Err("--workspace-root must be an absolute path".to_string()); + } + root = Some(PathBuf::from(value)); + index += 2; + } + value => return Err(format!("unknown worker argument: {value}")), + } + } + if !stdio { + return Err("--stdio is required".to_string()); + } + root.ok_or_else(|| "--workspace-root is required".to_string()) +} + +fn write_diagnostic(message: &str) { + let mut message = message.as_bytes().to_vec(); + const MAX_DIAGNOSTIC_BYTES: usize = 16 * 1024; + if message.len() > MAX_DIAGNOSTIC_BYTES { + message.truncate(MAX_DIAGNOSTIC_BYTES); + } + let stderr = std::io::stderr(); + let mut stderr = stderr.lock(); + let _ = std::io::Write::write_all(&mut stderr, &message); + let _ = std::io::Write::write_all(&mut stderr, b"\n"); +} diff --git a/crates/bunkerbox-worker/src/platform.rs b/crates/bunkerbox-worker/src/platform.rs new file mode 100644 index 0000000..dc1dc3b --- /dev/null +++ b/crates/bunkerbox-worker/src/platform.rs @@ -0,0 +1,222 @@ +use std::ffi::{CStr, CString, OsString}; +use std::fs::{File, OpenOptions}; +use std::io; +use std::os::fd::{AsRawFd, FromRawFd, RawFd}; +use std::os::unix::ffi::OsStringExt; +use std::os::unix::fs::OpenOptionsExt; + +pub const OPEN_DIRECTORY_FLAGS: i32 = libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC; +pub const OPEN_FILE_FLAGS: i32 = libc::O_RDONLY | libc::O_NOFOLLOW | libc::O_CLOEXEC; + +pub fn open_root(path: &std::path::Path) -> Result { + let file = OpenOptions::new() + .read(true) + .custom_flags(OPEN_DIRECTORY_FLAGS & !libc::O_RDONLY) + .open(path) + .map_err(|error| format!("open worker workspace root: {error}"))?; + validate_private_directory(&file, "worker workspace root")?; + Ok(file) +} + +pub fn validate_private_directory(file: &File, label: &str) -> Result<(), String> { + let metadata = stat_fd(file.as_raw_fd()).map_err(|error| format!("stat {label}: {error}"))?; + if metadata.st_mode & libc::S_IFMT != libc::S_IFDIR { + return Err(format!("{label} is not a directory")); + } + if metadata.st_uid != unsafe { libc::geteuid() } { + return Err(format!("{label} is not owned by the worker account")); + } + if metadata.st_mode & 0o077 != 0 { + return Err(format!("{label} is not private")); + } + Ok(()) +} + +pub fn open_dir_at(parent: &File, name: &str) -> io::Result { + open_at(parent.as_raw_fd(), name, OPEN_DIRECTORY_FLAGS) +} + +pub fn open_file_at(parent: &File, name: &str) -> io::Result { + open_at(parent.as_raw_fd(), name, OPEN_FILE_FLAGS) +} + +pub fn open_lock_at(parent: &File, name: &str) -> io::Result { + open_at(parent.as_raw_fd(), name, libc::O_RDWR | libc::O_NOFOLLOW | libc::O_CLOEXEC) +} + +pub fn open_at(parent: RawFd, name: &str, flags: i32) -> io::Result { + let name = c_string(name)?; + let fd = unsafe { libc::openat(parent, name.as_ptr(), flags, 0) }; + if fd < 0 { + return Err(io::Error::last_os_error()); + } + Ok(unsafe { File::from_raw_fd(fd) }) +} + +pub fn create_file_at(parent: &File, name: &str, mode: u32) -> io::Result { + let name = c_string(name)?; + let fd = unsafe { + libc::openat( + parent.as_raw_fd(), + name.as_ptr(), + libc::O_WRONLY | libc::O_CREAT | libc::O_EXCL | libc::O_NOFOLLOW | libc::O_CLOEXEC, + mode as libc::mode_t, + ) + }; + if fd < 0 { + return Err(io::Error::last_os_error()); + } + Ok(unsafe { File::from_raw_fd(fd) }) +} + +pub fn create_dir_at(parent: &File, name: &str, mode: u32) -> io::Result<()> { + let name = c_string(name)?; + if unsafe { libc::mkdirat(parent.as_raw_fd(), name.as_ptr(), mode as libc::mode_t) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn chmod_fd(file: &File, mode: u32) -> io::Result<()> { + if unsafe { libc::fchmod(file.as_raw_fd(), mode as libc::mode_t) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn sync_fd(file: &File) -> io::Result<()> { + if unsafe { libc::fsync(file.as_raw_fd()) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn stat_fd(fd: RawFd) -> io::Result { + let mut metadata = unsafe { std::mem::zeroed::() }; + if unsafe { libc::fstat(fd, &mut metadata) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(metadata) +} + +pub fn stat_at(parent: &File, name: &str) -> io::Result { + let name = c_string(name)?; + let mut metadata = unsafe { std::mem::zeroed::() }; + if unsafe { libc::fstatat(parent.as_raw_fd(), name.as_ptr(), &mut metadata, libc::AT_SYMLINK_NOFOLLOW) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(metadata) +} + +pub fn unlink_at(parent: &File, name: &str, flags: i32) -> io::Result<()> { + let name = c_string(name)?; + if unsafe { libc::unlinkat(parent.as_raw_fd(), name.as_ptr(), flags) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn remove_tree_at(parent: &File, name: &str) -> io::Result<()> { + let child = match open_dir_at(parent, name) { + Ok(child) => Some(child), + Err(error) if error.kind() == io::ErrorKind::NotADirectory || error.raw_os_error() == Some(libc::ELOOP) => None, + Err(error) => return Err(error), + }; + if let Some(child) = child { + for entry in list_names(&child)? { + let metadata = stat_at(&child, &entry)?; + if metadata.st_mode & libc::S_IFMT == libc::S_IFDIR { + remove_tree_at(&child, &entry)?; + } else { + unlink_at(&child, &entry, 0)?; + } + } + unlink_at(parent, name, libc::AT_REMOVEDIR) + } else { + unlink_at(parent, name, 0) + } +} + +pub fn list_names(directory: &File) -> io::Result> { + let duplicate = unsafe { libc::dup(directory.as_raw_fd()) }; + if duplicate < 0 { + return Err(io::Error::last_os_error()); + } + let stream = unsafe { libc::fdopendir(duplicate) }; + if stream.is_null() { + unsafe { libc::close(duplicate) }; + return Err(io::Error::last_os_error()); + } + let mut names = Vec::new(); + loop { + let entry = unsafe { libc::readdir(stream) }; + if entry.is_null() { + break; + } + let name = unsafe { CStr::from_ptr((*entry).d_name.as_ptr()) }.to_bytes(); + if name != b"." && name != b".." { + let name = + OsString::from_vec(name.to_vec()).into_string().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "non-UTF-8 state name"))?; + names.push(name); + } + } + unsafe { libc::closedir(stream) }; + names.sort(); + Ok(names) +} + +pub fn lock_exclusive(file: &File) -> io::Result { + let mut lock = unsafe { std::mem::zeroed::() }; + lock.l_type = libc::F_WRLCK as _; + lock.l_whence = libc::SEEK_SET as _; + if unsafe { libc::fcntl(file.as_raw_fd(), libc::F_SETLK, &lock) } == 0 { + return Ok(true); + } + let error = io::Error::last_os_error(); + if matches!(error.raw_os_error(), Some(code) if code == libc::EACCES || code == libc::EAGAIN) { + Ok(false) + } else { + Err(error) + } +} + +pub fn set_nonblocking_fd(fd: RawFd) -> io::Result<()> { + let flags = unsafe { libc::fcntl(fd, libc::F_GETFL) }; + if flags < 0 { + return Err(io::Error::last_os_error()); + } + if unsafe { libc::fcntl(fd, libc::F_SETFL, flags | libc::O_NONBLOCK) } < 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn change_directory(fd: RawFd) -> io::Result<()> { + if unsafe { libc::fchdir(fd) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn set_process_group() -> io::Result<()> { + if unsafe { libc::setpgid(0, 0) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + +pub fn signal_process_group(pgid: libc::pid_t, signal: libc::c_int) { + if pgid > 0 { + unsafe { + libc::kill(-pgid, signal); + } + } +} + +fn c_string(value: &str) -> io::Result { + CString::new(value.as_bytes()).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "NUL in worker path")) +} + +#[cfg(test)] +#[path = "platform_ut.rs"] +mod tests; diff --git a/crates/bunkerbox-worker/src/platform_ut.rs b/crates/bunkerbox-worker/src/platform_ut.rs new file mode 100644 index 0000000..6e3d047 --- /dev/null +++ b/crates/bunkerbox-worker/src/platform_ut.rs @@ -0,0 +1,39 @@ +use super::*; +use std::fs; +use std::os::unix::fs::PermissionsExt; +use std::path::Path; +use tempfile::tempdir; + +#[test] +fn root_requires_private_owned_directory_without_following_final_symlink() { + let temp = tempdir().unwrap(); + let root = temp.path().join("root"); + fs::create_dir(&root).unwrap(); + fs::set_permissions(&root, fs::Permissions::from_mode(0o700)).unwrap(); + assert!(open_root(&root).is_ok()); + + fs::set_permissions(&root, fs::Permissions::from_mode(0o755)).unwrap(); + assert!(open_root(&root).is_err()); + + let link = temp.path().join("link"); + #[cfg(unix)] + std::os::unix::fs::symlink(&root, &link).unwrap(); + assert!(open_root(Path::new(&link)).is_err()); +} + +#[test] +fn descriptor_relative_tree_removal_does_not_follow_symlinks() { + let temp = tempdir().unwrap(); + let root_path = temp.path().join("root"); + let outside = temp.path().join("outside"); + fs::create_dir(&root_path).unwrap(); + fs::create_dir(&outside).unwrap(); + fs::set_permissions(&root_path, fs::Permissions::from_mode(0o700)).unwrap(); + let root = open_root(&root_path).unwrap(); + create_dir_at(&root, "state", 0o700).unwrap(); + fs::write(outside.join("outside.txt"), b"outside").unwrap(); + #[cfg(unix)] + std::os::unix::fs::symlink(&outside, root_path.join("state/link")).unwrap(); + remove_tree_at(&root, "state").unwrap(); + assert!(outside.join("outside.txt").exists()); +} diff --git a/crates/bunkerbox-worker/src/process.rs b/crates/bunkerbox-worker/src/process.rs new file mode 100644 index 0000000..3841c60 --- /dev/null +++ b/crates/bunkerbox-worker/src/process.rs @@ -0,0 +1,362 @@ +use crate::platform; +use crate::storage; +use bunkerbox_worker_protocol::{WorkerBuild, WorkerMessage, WorkerRequestId, WorkerSessionId}; +use std::fs::{self, File}; +use std::io::{self, Read}; +use std::os::fd::AsRawFd; +use std::os::unix::fs::PermissionsExt; +use std::os::unix::process::CommandExt; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::mpsc::{self, RecvTimeoutError}; +use std::sync::Arc; +use std::thread; +use std::time::{Duration, Instant}; + +pub const WORKER_BUILD_TIMEOUT: Duration = Duration::from_secs(30); +pub const WORKER_MAX_OUTPUT_BYTES: u64 = 64 * 1024 * 1024; +const OUTPUT_BUFFER_BYTES: usize = 8192; +const POST_EXIT_DRAIN_TIMEOUT: Duration = Duration::from_millis(100); +const FINAL_DRAIN_TIMEOUT: Duration = Duration::from_millis(100); +const MAX_STALE_JOBS: usize = 256; + +pub trait OutputSink: Send + Sync { + fn send(&self, message: WorkerMessage) -> Result<(), String>; +} + +pub struct JobWorkspace { + parent: File, + name: String, + lock_name: String, + root: File, + lock: File, +} + +impl JobWorkspace { + pub fn create(parent: &File) -> Result { + for _ in 0..32 { + let name = format!("job-{}-{}", unsafe { libc::getpid() }, storage::next_job_id()); + let lock_name = format!("{name}.lock"); + let lock = match platform::create_file_at(parent, &lock_name, 0o600) { + Ok(lock) => lock, + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(format!("create worker job lock: {error}")), + }; + let locked = match platform::lock_exclusive(&lock) { + Ok(locked) => locked, + Err(error) => { + let _ = platform::unlink_at(parent, &lock_name, 0); + return Err(format!("lock worker job: {error}")); + } + }; + if !locked { + let _ = platform::unlink_at(parent, &lock_name, 0); + continue; + } + if let Err(error) = platform::create_dir_at(parent, &name, 0o700) { + let _ = platform::unlink_at(parent, &lock_name, 0); + if error.kind() == io::ErrorKind::AlreadyExists { + continue; + } + return Err(format!("create worker job workspace: {error}")); + } + let root = match platform::open_dir_at(parent, &name) { + Ok(root) => root, + Err(error) => { + let _ = platform::remove_tree_at(parent, &name); + let _ = platform::unlink_at(parent, &lock_name, 0); + return Err(format!("open worker job workspace: {error}")); + } + }; + return Ok(Self { + parent: parent.try_clone().map_err(|error| format!("clone worker jobs directory: {error}"))?, + name, + lock_name, + root, + lock, + }); + } + Err("could not reserve a unique worker job workspace".to_string()) + } + + pub fn root(&self) -> &File { + &self.root + } +} + +impl Drop for JobWorkspace { + fn drop(&mut self) { + let _ = &self.lock; + let _ = platform::remove_tree_at(&self.parent, &self.name); + let _ = platform::unlink_at(&self.parent, &self.lock_name, 0); + } +} + +pub fn cleanup_stale_jobs(parent: &File) -> io::Result<()> { + for lock_name in platform::list_names(parent)?.into_iter().take(MAX_STALE_JOBS) { + let Some(name) = lock_name.strip_suffix(".lock") else { continue }; + if !name.starts_with("job-") { + continue; + } + let Ok(lock) = platform::open_lock_at(parent, &lock_name) else { continue }; + if platform::lock_exclusive(&lock)? { + let _ = platform::remove_tree_at(parent, name); + let _ = platform::unlink_at(parent, &lock_name, 0); + } + } + Ok(()) +} + +pub fn execute_build( + job: &JobWorkspace, build: &WorkerBuild, request_id: WorkerRequestId, session_id: WorkerSessionId, sink: &S, disconnected: &dyn Fn() -> bool, +) -> Result { + validate_executable(build.trusted_executable_path())?; + let cwd = storage::open_relative_directory(job.root(), build.cwd().as_str())?; + let cwd_fd = cwd.as_raw_fd(); + let mut command = std::process::Command::new(build.trusted_executable_path()); + command.args(build.argv()).stdin(std::process::Stdio::null()).stdout(std::process::Stdio::piped()).stderr(std::process::Stdio::piped()); + command.env_clear(); + for (key, value) in build.guest_env() { + command.env(key, value); + } + for (key, value) in build.target_env() { + command.env(key, value); + } + unsafe { + command.pre_exec(move || { + platform::set_process_group()?; + platform::change_directory(cwd_fd)?; + Ok(()) + }); + } + + let mut child = command.spawn().map_err(|error| format!("spawn worker tool: {error}"))?; + let pgid = child.id() as libc::pid_t; + let stdout = match child.stdout.take() { + Some(stdout) => stdout, + None => { + kill_group(pgid); + let _ = child.wait(); + return Err("worker child has no stdout".to_string()); + } + }; + let stderr = match child.stderr.take() { + Some(stderr) => stderr, + None => { + kill_group(pgid); + let _ = child.wait(); + return Err("worker child has no stderr".to_string()); + } + }; + if let Err(error) = platform::set_nonblocking_fd(stdout.as_raw_fd()) { + kill_group(pgid); + let _ = child.wait(); + return Err(format!("set worker stdout nonblocking: {error}")); + } + if let Err(error) = platform::set_nonblocking_fd(stderr.as_raw_fd()) { + kill_group(pgid); + let _ = child.wait(); + return Err(format!("set worker stderr nonblocking: {error}")); + } + + let stop = Arc::new(AtomicBool::new(false)); + let (events_tx, events_rx) = mpsc::channel(); + let stdout_thread = spawn_pump(stdout, StreamKind::Stdout, events_tx.clone(), stop.clone()); + let stderr_thread = spawn_pump(stderr, StreamKind::Stderr, events_tx, stop.clone()); + + let started = Instant::now(); + let mut child_status = None; + let mut stdout_done = false; + let mut stderr_done = false; + let mut output_total = 0u64; + let mut failure = None; + let mut post_exit_deadline = None; + let mut final_deadline = None; + let mut group_killed = false; + + while child_status.is_none() || !stdout_done || !stderr_done { + if child_status.is_none() && failure.is_none() && disconnected() { + failure = Some("worker protocol input disconnected during build".to_string()); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + if child_status.is_none() && failure.is_none() && started.elapsed() >= WORKER_BUILD_TIMEOUT { + failure = Some("worker build timed out".to_string()); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + + if child_status.is_none() { + match child.try_wait() { + Ok(Some(status)) => { + child_status = Some(status); + if failure.is_none() { + platform::signal_process_group(pgid, libc::SIGTERM); + post_exit_deadline = Some(Instant::now() + POST_EXIT_DRAIN_TIMEOUT); + } else { + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + } + Ok(None) => {} + Err(error) => { + failure.get_or_insert(format!("wait for worker tool: {error}")); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + } + } + + if let Some(deadline) = post_exit_deadline { + if Instant::now() >= deadline { + kill_group(pgid); + group_killed = true; + post_exit_deadline = None; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + } + if let Some(deadline) = final_deadline { + if Instant::now() >= deadline { + stop.store(true, Ordering::Release); + break; + } + } + + match events_rx.recv_timeout(Duration::from_millis(10)) { + Ok(PumpEvent::Data(stream, bytes)) => { + if failure.is_some() { + continue; + } + output_total = output_total.saturating_add(bytes.len() as u64); + if output_total > WORKER_MAX_OUTPUT_BYTES { + failure = Some("worker combined output limit exceeded".to_string()); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + continue; + } + let message = match stream { + StreamKind::Stdout => WorkerMessage::stdout(request_id, session_id, bytes), + StreamKind::Stderr => WorkerMessage::stderr(request_id, session_id, bytes), + }; + if let Err(error) = sink.send(message) { + failure.get_or_insert(error); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + } + Ok(PumpEvent::End(stream)) => match stream { + StreamKind::Stdout => stdout_done = true, + StreamKind::Stderr => stderr_done = true, + }, + Ok(PumpEvent::Failed(stream, error)) => { + failure.get_or_insert(format!("read worker {}: {error}", stream.label())); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + Err(RecvTimeoutError::Timeout) => {} + Err(RecvTimeoutError::Disconnected) => { + failure.get_or_insert("worker output pump channel disconnected".to_string()); + kill_group(pgid); + group_killed = true; + final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); + } + } + } + + if !group_killed { + kill_group(pgid); + } + stop.store(true, Ordering::Release); + let status = match child_status { + Some(status) => status, + None => child.wait().map_err(|error| format!("reap worker tool: {error}"))?, + }; + let _ = stdout_thread.join(); + let _ = stderr_thread.join(); + if let Some(error) = failure { + return Err(error); + } + Ok(status.code().unwrap_or(-1)) +} + +fn spawn_pump( + mut reader: R, stream: StreamKind, sender: mpsc::Sender, stop: Arc, +) -> thread::JoinHandle<()> { + thread::spawn(move || { + let mut buffer = [0u8; OUTPUT_BUFFER_BYTES]; + let mut total = 0u64; + loop { + if stop.load(Ordering::Acquire) { + return; + } + match reader.read(&mut buffer) { + Ok(0) => { + let _ = sender.send(PumpEvent::End(stream)); + return; + } + Ok(count) => { + total = total.saturating_add(count as u64); + if total > WORKER_MAX_OUTPUT_BYTES { + let _ = sender.send(PumpEvent::Failed(stream, "worker output limit exceeded".to_string())); + return; + } + if sender.send(PumpEvent::Data(stream, buffer[..count].to_vec())).is_err() { + return; + } + } + Err(error) if error.kind() == io::ErrorKind::WouldBlock => thread::sleep(Duration::from_millis(5)), + Err(error) => { + let _ = sender.send(PumpEvent::Failed(stream, error.to_string())); + return; + } + } + } + }) +} + +fn validate_executable(path: &str) -> Result<(), String> { + let metadata = fs::metadata(path).map_err(|error| format!("inspect worker executable: {error}"))?; + if !metadata.file_type().is_file() { + return Err("worker executable is not a regular file".to_string()); + } + if metadata.permissions().mode() & 0o111 == 0 { + return Err("worker executable is not executable".to_string()); + } + Ok(()) +} + +fn kill_group(pgid: libc::pid_t) { + platform::signal_process_group(pgid, libc::SIGTERM); + platform::signal_process_group(pgid, libc::SIGKILL); +} + +#[derive(Clone, Copy)] +enum StreamKind { + Stdout, + Stderr, +} + +impl StreamKind { + fn label(self) -> &'static str { + match self { + Self::Stdout => "stdout", + Self::Stderr => "stderr", + } + } +} + +enum PumpEvent { + Data(StreamKind, Vec), + End(StreamKind), + Failed(StreamKind, String), +} + +#[cfg(test)] +#[path = "process_ut.rs"] +mod tests; diff --git a/crates/bunkerbox-worker/src/process_ut.rs b/crates/bunkerbox-worker/src/process_ut.rs new file mode 100644 index 0000000..f8cf12a --- /dev/null +++ b/crates/bunkerbox-worker/src/process_ut.rs @@ -0,0 +1,71 @@ +use super::*; +use crate::platform; +use crate::worker::FrameWriter; +use bunkerbox_worker_protocol::{WorkerBuild, WorkerMessage, WorkerRequestId, WorkerSessionId, WorkerUploadId}; +use std::fs; +use std::io::Cursor; +use std::os::unix::fs::PermissionsExt; +use std::time::{Duration, Instant}; +use tempfile::tempdir; + +#[test] +fn direct_process_execution_preserves_literal_arguments_and_reaps_group() { + let temp = tempdir().unwrap(); + fs::set_permissions(temp.path(), fs::Permissions::from_mode(0o700)).unwrap(); + let jobs = platform::open_root(temp.path()).unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + let build = WorkerBuild::new( + "echo", + "/bin/echo", + vec!["literal $(argument)".to_string()], + "", + Vec::new(), + vec![("PATH".to_string(), "/usr/bin:/bin".to_string())], + WorkerUploadId([3; 16]), + ) + .unwrap(); + let writer = FrameWriter::new(Vec::new()); + let status = execute_build(&job, &build, WorkerRequestId([1; 16]), WorkerSessionId([2; 16]), &writer, &|| false).unwrap(); + assert_eq!(status, 0); + let bytes = writer.into_inner().unwrap(); + let mut reader = Cursor::new(bytes); + let mut output = Vec::new(); + while let Some(message) = WorkerMessage::read_blocking_optional(&mut reader).unwrap() { + if let WorkerMessage::Stdout { data, .. } = message { + output.extend_from_slice(&data); + } + } + assert_eq!(output, b"literal $(argument)\n"); +} + +#[test] +fn disconnect_and_inherited_pipes_do_not_leave_a_build_running() { + let temp = tempdir().unwrap(); + fs::set_permissions(temp.path(), fs::Permissions::from_mode(0o700)).unwrap(); + let jobs = platform::open_root(temp.path()).unwrap(); + let script = temp.path().join("descendant.sh"); + fs::write(&script, b"#!/bin/sh\nprintf before\n(sleep 10) &\nexit 0\n").unwrap(); + fs::set_permissions(&script, fs::Permissions::from_mode(0o755)).unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + let build = WorkerBuild::new( + "descendant", + script.to_string_lossy(), + Vec::new(), + "", + Vec::new(), + vec![("PATH".to_string(), "/usr/bin:/bin".to_string())], + WorkerUploadId([4; 16]), + ) + .unwrap(); + let writer = FrameWriter::new(Vec::new()); + let started = Instant::now(); + let status = execute_build(&job, &build, WorkerRequestId([5; 16]), WorkerSessionId([6; 16]), &writer, &|| false).unwrap(); + assert_eq!(status, 0); + assert!(started.elapsed() < Duration::from_secs(2)); + + let job = JobWorkspace::create(&jobs).unwrap(); + let writer = FrameWriter::new(Vec::new()); + let error = execute_build(&job, &build, WorkerRequestId([7; 16]), WorkerSessionId([8; 16]), &writer, &|| true).unwrap_err(); + assert!(error.contains("disconnected")); + assert!(started.elapsed() < Duration::from_secs(2)); +} diff --git a/crates/bunkerbox-worker/src/storage.rs b/crates/bunkerbox-worker/src/storage.rs new file mode 100644 index 0000000..e6f9c63 --- /dev/null +++ b/crates/bunkerbox-worker/src/storage.rs @@ -0,0 +1,627 @@ +use crate::platform; +use bunkerbox_worker_protocol::{ + validate_upload_manifest, WorkerDigest, WorkerEntryKind, WorkerProtocolError, WorkerRelativePath, WorkerSessionId, WorkerUploadEntry, + WorkerUploadId, MAX_WORKER_MANIFEST_BYTES, +}; +use sha2::{Digest, Sha256}; +use std::collections::BTreeMap; +use std::fs::File; +use std::io::{self, Read, Write}; +use std::os::fd::AsRawFd; +use std::sync::atomic::{AtomicU64, Ordering}; + +const STATE_DIRECTORY: &str = ".bunkerbox-worker"; +const SESSIONS_DIRECTORY: &str = "sessions"; +const UPLOADS_DIRECTORY: &str = "uploads"; +const JOBS_DIRECTORY: &str = "jobs"; +const MANIFEST_FILE: &str = "manifest"; +const COMPLETE_FILE: &str = "complete"; +const LOCK_FILE: &str = "lock"; +const FILES_DIRECTORY: &str = "files"; +const MANIFEST_MAGIC: [u8; 4] = *b"BBWM"; +const MANIFEST_VERSION: u16 = 1; +const COPY_BUFFER_BYTES: usize = 64 * 1024; +const MAX_STALE_SESSIONS: usize = 64; +const MAX_STALE_UPLOADS_PER_SESSION: usize = 256; + +static NEXT_JOB_ID: AtomicU64 = AtomicU64::new(1); + +pub struct UploadStore { + sessions: File, + jobs: File, +} + +impl UploadStore { + pub fn new(root: &File) -> Result { + let state = private_directory(root, STATE_DIRECTORY)?; + let sessions = private_directory(&state, SESSIONS_DIRECTORY)?; + let jobs = private_directory(&state, JOBS_DIRECTORY)?; + let store = Self { sessions, jobs }; + store.cleanup_stale().map_err(|error| format!("clean stale worker state: {error}"))?; + Ok(store) + } + + pub fn begin( + &self, session_id: WorkerSessionId, upload_id: WorkerUploadId, entries: Vec, + ) -> Result { + require_nonzero_id(session_id.0, "worker session")?; + require_nonzero_id(upload_id.0, "worker upload")?; + validate_upload_manifest(&entries).map_err(protocol_error)?; + validate_manifest_structure(&entries)?; + + let session = private_directory(&self.sessions, &hex_id(session_id.0))?; + let uploads = private_directory(&session, UPLOADS_DIRECTORY)?; + let token_name = hex_id(upload_id.0); + platform::create_dir_at(&uploads, &token_name, 0o700).map_err(|error| { + if error.kind() == io::ErrorKind::AlreadyExists { + format!("worker upload token is already reserved: {token_name}") + } else { + format!("reserve worker upload token: {error}") + } + })?; + let token = match platform::open_dir_at(&uploads, &token_name) { + Ok(token) => token, + Err(error) => { + let _ = platform::remove_tree_at(&uploads, &token_name); + return Err(format!("open reserved worker upload: {error}")); + } + }; + if let Err(error) = platform::validate_private_directory(&token, "worker upload token") { + let _ = platform::remove_tree_at(&uploads, &token_name); + return Err(error); + } + let lock = match platform::create_file_at(&token, LOCK_FILE, 0o600) { + Ok(lock) => lock, + Err(error) => { + let _ = platform::remove_tree_at(&uploads, &token_name); + return Err(format!("create worker upload lock: {error}")); + } + }; + let locked = match platform::lock_exclusive(&lock) { + Ok(locked) => locked, + Err(error) => { + let _ = platform::remove_tree_at(&uploads, &token_name); + return Err(format!("lock worker upload: {error}")); + } + }; + if !locked { + let _ = platform::remove_tree_at(&uploads, &token_name); + return Err("worker upload token is active".to_string()); + } + let files = match private_directory(&token, FILES_DIRECTORY) { + Ok(files) => files, + Err(error) => { + let _ = platform::remove_tree_at(&uploads, &token_name); + return Err(error); + } + }; + + let mut transaction = + UploadTransaction { uploads, token_name, token, lock, files, session_id, upload_id, entries, states: BTreeMap::new(), committed: false }; + transaction.initialize()?; + Ok(transaction) + } + + pub fn open_completed(&self, session_id: WorkerSessionId, upload_id: WorkerUploadId) -> Result { + require_nonzero_id(session_id.0, "worker session")?; + require_nonzero_id(upload_id.0, "worker upload")?; + let session = open_existing_private_directory(&self.sessions, &hex_id(session_id.0), "worker session")?; + let uploads = open_existing_private_directory(&session, UPLOADS_DIRECTORY, "worker uploads")?; + let token_name = hex_id(upload_id.0); + let token = open_existing_private_directory(&uploads, &token_name, "worker upload token")?; + let lock = platform::open_lock_at(&token, LOCK_FILE).map_err(|error| format!("open worker upload lock: {error}"))?; + if !platform::lock_exclusive(&lock).map_err(|error| format!("lock worker upload: {error}"))? { + return Err("worker upload token is active".to_string()); + } + let complete = platform::open_file_at(&token, COMPLETE_FILE).map_err(|_| "worker upload is not complete".to_string())?; + let marker = read_bounded(complete, 1)?; + if !marker.is_empty() { + return Err("worker upload completion marker is invalid".to_string()); + } + let manifest_file = platform::open_file_at(&token, MANIFEST_FILE).map_err(|error| format!("open worker upload manifest: {error}"))?; + let manifest = StoredManifest::read(manifest_file)?; + if manifest.session_id != session_id || manifest.upload_id != upload_id { + return Err("worker upload identity mismatch".to_string()); + } + let files = open_existing_private_directory(&token, FILES_DIRECTORY, "worker upload files")?; + Ok(StoredUpload { token, lock, files, entries: manifest.entries }) + } + + pub fn cleanup(&self, session_id: WorkerSessionId, upload_id: WorkerUploadId) -> Result<(), String> { + require_nonzero_id(session_id.0, "worker session")?; + require_nonzero_id(upload_id.0, "worker upload")?; + let session = open_existing_private_directory(&self.sessions, &hex_id(session_id.0), "worker session")?; + let uploads = open_existing_private_directory(&session, UPLOADS_DIRECTORY, "worker uploads")?; + let token_name = hex_id(upload_id.0); + let token = open_existing_private_directory(&uploads, &token_name, "worker upload token")?; + let lock = platform::open_lock_at(&token, LOCK_FILE).map_err(|error| format!("open worker upload lock: {error}"))?; + if !platform::lock_exclusive(&lock).map_err(|error| format!("lock worker upload: {error}"))? { + return Err("worker upload token is active".to_string()); + } + platform::remove_tree_at(&uploads, &token_name).map_err(|error| format!("remove worker upload: {error}")) + } + + pub fn jobs_directory(&self) -> Result { + self.jobs.try_clone().map_err(|error| format!("clone worker jobs directory: {error}")) + } + + fn cleanup_stale(&self) -> io::Result<()> { + let sessions = platform::list_names(&self.sessions)?; + for session_name in sessions.into_iter().take(MAX_STALE_SESSIONS) { + if !is_hex_id(&session_name) { + continue; + } + let Ok(session) = platform::open_dir_at(&self.sessions, &session_name) else { continue }; + let Ok(uploads) = platform::open_dir_at(&session, UPLOADS_DIRECTORY) else { continue }; + for token_name in platform::list_names(&uploads)?.into_iter().take(MAX_STALE_UPLOADS_PER_SESSION) { + if !is_hex_id(&token_name) { + continue; + } + let Ok(token) = platform::open_dir_at(&uploads, &token_name) else { continue }; + let Ok(lock) = platform::open_lock_at(&token, LOCK_FILE) else { continue }; + if !platform::lock_exclusive(&lock)? { + continue; + } + if platform::open_file_at(&token, COMPLETE_FILE).is_err() { + let _ = platform::remove_tree_at(&uploads, &token_name); + } + } + } + Ok(()) + } +} + +pub struct UploadTransaction { + uploads: File, + token_name: String, + token: File, + lock: File, + files: File, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entries: Vec, + states: BTreeMap, + committed: bool, +} + +impl UploadTransaction { + fn initialize(&mut self) -> Result<(), String> { + let manifest = StoredManifest { session_id: self.session_id, upload_id: self.upload_id, entries: self.entries.clone() }; + let encoded = manifest.encode()?; + let mut manifest_file = + platform::create_file_at(&self.token, MANIFEST_FILE, 0o600).map_err(|error| format!("create worker upload manifest: {error}"))?; + manifest_file.write_all(&encoded).map_err(|error| format!("write worker upload manifest: {error}"))?; + platform::sync_fd(&manifest_file).map_err(|error| format!("flush worker upload manifest: {error}"))?; + + for entry in &self.entries { + match entry.kind() { + WorkerEntryKind::Directory => ensure_directory(&self.files, entry.path().as_str(), entry.mode())?, + WorkerEntryKind::File => { + let file = create_relative_file(&self.files, entry.path().as_str(), entry.mode())?; + self.states.insert( + entry.path().as_str().to_string(), + PendingFile { + file, + size: entry.size(), + digest: *entry.digest().ok_or_else(|| "worker regular file has no digest".to_string())?, + received: 0, + hasher: Sha256::new(), + }, + ); + } + } + } + platform::sync_fd(&self.files).map_err(|error| format!("flush worker upload files: {error}"))?; + Ok(()) + } + + pub fn accept_chunk(&mut self, path: &WorkerRelativePath, offset: u64, data: &[u8]) -> Result<(), String> { + let state = self.states.get_mut(path.as_str()).ok_or_else(|| format!("worker chunk path is not a declared file: {}", path.as_str()))?; + if offset != state.received { + return Err(format!("worker chunk offset is out of order for {}", path.as_str())); + } + let end = offset.checked_add(data.len() as u64).ok_or_else(|| "worker chunk offset overflow".to_string())?; + if end > state.size { + return Err(format!("worker chunk exceeds declared file size: {}", path.as_str())); + } + state.file.write_all(data).map_err(|error| format!("write worker upload file {}: {error}", path.as_str()))?; + state.hasher.update(data); + state.received = end; + Ok(()) + } + + pub fn commit(&mut self) -> Result<(), String> { + for entry in &self.entries { + if entry.kind() != WorkerEntryKind::File { + continue; + } + let state = + self.states.get_mut(entry.path().as_str()).ok_or_else(|| format!("worker file state is missing: {}", entry.path().as_str()))?; + if state.received != state.size { + return Err(format!("worker upload file is incomplete: {}", entry.path().as_str())); + } + if state.hasher.clone().finalize().as_slice() != state.digest { + return Err(format!("worker upload file digest mismatch: {}", entry.path().as_str())); + } + platform::sync_fd(&state.file).map_err(|error| format!("flush worker upload file {}: {error}", entry.path().as_str()))?; + let metadata = + platform::stat_fd(state.file.as_raw_fd()).map_err(|error| format!("stat worker upload file {}: {error}", entry.path().as_str()))?; + validate_regular_file(&metadata, state.size, entry.mode(), entry.path().as_str())?; + } + let marker = platform::create_file_at(&self.token, COMPLETE_FILE, 0o600).map_err(|error| format!("publish worker upload: {error}"))?; + platform::sync_fd(&marker).map_err(|error| format!("flush worker upload marker: {error}"))?; + platform::sync_fd(&self.token).map_err(|error| format!("flush worker upload directory: {error}"))?; + self.committed = true; + Ok(()) + } +} + +impl Drop for UploadTransaction { + fn drop(&mut self) { + if !self.committed { + let _ = platform::remove_tree_at(&self.uploads, &self.token_name); + } + let _ = &self.lock; + } +} + +pub struct StoredUpload { + token: File, + lock: File, + files: File, + entries: Vec, +} + +impl StoredUpload { + pub fn materialize(&self, destination: &File) -> Result<(), String> { + let _ = (&self.token, &self.lock); + for entry in &self.entries { + match entry.kind() { + WorkerEntryKind::Directory => ensure_directory(destination, entry.path().as_str(), entry.mode())?, + WorkerEntryKind::File => { + let source = open_relative_file(&self.files, entry.path().as_str())?; + let source_metadata = platform::stat_fd(source.as_raw_fd()).map_err(|error| format!("stat stored worker file: {error}"))?; + validate_regular_file(&source_metadata, entry.size(), entry.mode(), entry.path().as_str())?; + let target = create_relative_file(destination, entry.path().as_str(), entry.mode())?; + copy_and_verify(&source, &target, entry)?; + } + } + } + platform::sync_fd(destination).map_err(|error| format!("flush worker job workspace: {error}"))?; + Ok(()) + } +} + +struct PendingFile { + file: File, + size: u64, + digest: WorkerDigest, + received: u64, + hasher: Sha256, +} + +struct StoredManifest { + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entries: Vec, +} + +impl StoredManifest { + fn encode(&self) -> Result, String> { + validate_upload_manifest(&self.entries).map_err(protocol_error)?; + let mut bytes = Vec::new(); + bytes.extend_from_slice(&MANIFEST_MAGIC); + bytes.extend_from_slice(&MANIFEST_VERSION.to_le_bytes()); + bytes.extend_from_slice(&self.session_id.0); + bytes.extend_from_slice(&self.upload_id.0); + put_u32(&mut bytes, self.entries.len())?; + for entry in &self.entries { + put_string(&mut bytes, entry.path().as_str())?; + bytes.push(entry.kind().as_u8()); + bytes.extend_from_slice(&entry.mode().to_le_bytes()); + bytes.extend_from_slice(&entry.size().to_le_bytes()); + match entry.digest() { + Some(digest) => { + bytes.push(1); + bytes.extend_from_slice(digest); + } + None => bytes.push(0), + } + if bytes.len() > MAX_WORKER_MANIFEST_BYTES { + return Err("worker stored manifest exceeds maximum length".to_string()); + } + } + Ok(bytes) + } + + fn read(file: File) -> Result { + let bytes = read_bounded(file, MAX_WORKER_MANIFEST_BYTES)?; + let mut reader = ManifestReader { bytes: &bytes, offset: 0 }; + if reader.take(4)? != MANIFEST_MAGIC { + return Err("worker stored manifest has invalid magic".to_string()); + } + if reader.u16()? != MANIFEST_VERSION { + return Err("worker stored manifest has unsupported version".to_string()); + } + let session_id = WorkerSessionId(reader.array16()?); + let upload_id = WorkerUploadId(reader.array16()?); + let count = reader.count()?; + let mut entries = Vec::with_capacity(count); + for _ in 0..count { + let path = reader.string()?; + let kind = WorkerEntryKind::from_u8(reader.u8()?).map_err(protocol_error)?; + let mode = reader.u32()?; + let size = reader.u64()?; + let digest = match reader.u8()? { + 0 => None, + 1 => Some(reader.array32()?), + _ => return Err("worker stored manifest has invalid digest flag".to_string()), + }; + entries.push(WorkerUploadEntry::new(path, kind, mode, size, digest).map_err(protocol_error)?); + } + reader.finish()?; + validate_upload_manifest(&entries).map_err(protocol_error)?; + validate_manifest_structure(&entries)?; + Ok(Self { session_id, upload_id, entries }) + } +} + +struct ManifestReader<'a> { + bytes: &'a [u8], + offset: usize, +} + +impl<'a> ManifestReader<'a> { + fn take(&mut self, length: usize) -> Result<&'a [u8], String> { + let end = self.offset.checked_add(length).ok_or_else(|| "worker stored manifest length overflow".to_string())?; + if end > self.bytes.len() { + return Err("worker stored manifest is truncated".to_string()); + } + let result = &self.bytes[self.offset..end]; + self.offset = end; + Ok(result) + } + + fn u8(&mut self) -> Result { + Ok(self.take(1)?[0]) + } + + fn u16(&mut self) -> Result { + Ok(u16::from_le_bytes(self.take(2)?.try_into().map_err(|_| "invalid worker manifest integer".to_string())?)) + } + + fn u32(&mut self) -> Result { + Ok(u32::from_le_bytes(self.take(4)?.try_into().map_err(|_| "invalid worker manifest integer".to_string())?)) + } + + fn u64(&mut self) -> Result { + Ok(u64::from_le_bytes(self.take(8)?.try_into().map_err(|_| "invalid worker manifest integer".to_string())?)) + } + + fn array16(&mut self) -> Result<[u8; 16], String> { + self.take(16)?.try_into().map_err(|_| "invalid worker manifest identifier".to_string()) + } + + fn array32(&mut self) -> Result<[u8; 32], String> { + self.take(32)?.try_into().map_err(|_| "invalid worker manifest digest".to_string()) + } + + fn count(&mut self) -> Result { + let count = self.u32()? as usize; + if count > bunkerbox_worker_protocol::MAX_WORKER_UPLOAD_ENTRIES { + return Err("worker stored manifest has too many entries".to_string()); + } + Ok(count) + } + + fn string(&mut self) -> Result { + let length = self.u32()? as usize; + if length > bunkerbox_worker_protocol::MAX_WORKER_PATH_BYTES { + return Err("worker stored manifest path is too long".to_string()); + } + String::from_utf8(self.take(length)?.to_vec()).map_err(|_| "worker stored manifest path is not UTF-8".to_string()) + } + + fn finish(self) -> Result<(), String> { + if self.offset != self.bytes.len() { + return Err("worker stored manifest has trailing bytes".to_string()); + } + Ok(()) + } +} + +fn private_directory(parent: &File, name: &str) -> Result { + match platform::create_dir_at(parent, name, 0o700) { + Ok(()) => {} + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {} + Err(error) => return Err(format!("create private worker directory {name}: {error}")), + } + let directory = platform::open_dir_at(parent, name).map_err(|error| format!("open private worker directory {name}: {error}"))?; + platform::validate_private_directory(&directory, &format!("worker directory {name}"))?; + Ok(directory) +} + +fn open_existing_private_directory(parent: &File, name: &str, label: &str) -> Result { + let directory = platform::open_dir_at(parent, name).map_err(|error| format!("open {label}: {error}"))?; + platform::validate_private_directory(&directory, label)?; + Ok(directory) +} + +fn ensure_directory(root: &File, relative: &str, mode: u32) -> Result<(), String> { + let components = components(relative)?; + let mut current = root.try_clone().map_err(|error| format!("clone worker directory: {error}"))?; + for (index, component) in components.iter().enumerate() { + let created = match platform::create_dir_at(¤t, component, 0o700) { + Ok(()) => true, + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => false, + Err(error) => return Err(format!("create worker directory {relative}: {error}")), + }; + current = platform::open_dir_at(¤t, component).map_err(|error| format!("open worker directory {relative}: {error}"))?; + if created || index + 1 == components.len() { + platform::chmod_fd(¤t, mode & 0o777).map_err(|error| format!("set worker directory mode {relative}: {error}"))?; + } + } + Ok(()) +} + +fn create_relative_file(root: &File, relative: &str, mode: u32) -> Result { + let components = components(relative)?; + let (file_name, parents) = components.split_last().ok_or_else(|| "worker file path is empty".to_string())?; + let parent = open_relative_directory_components(root, parents)?; + let file = platform::create_file_at(&parent, file_name, mode & 0o777).map_err(|error| format!("create worker file {relative}: {error}"))?; + platform::chmod_fd(&file, mode & 0o777).map_err(|error| format!("set worker file mode {relative}: {error}"))?; + Ok(file) +} + +fn open_relative_file(root: &File, relative: &str) -> Result { + let components = components(relative)?; + let (file_name, parents) = components.split_last().ok_or_else(|| "worker file path is empty".to_string())?; + let parent = open_relative_directory_components(root, parents)?; + platform::open_file_at(&parent, file_name).map_err(|error| format!("open worker file {relative}: {error}")) +} + +fn open_relative_directory_components(root: &File, components: &[&str]) -> Result { + let mut current = root.try_clone().map_err(|error| format!("clone worker directory: {error}"))?; + for component in components { + current = platform::open_dir_at(¤t, component).map_err(|error| format!("open worker parent directory: {error}"))?; + } + Ok(current) +} + +fn components(relative: &str) -> Result, String> { + bunkerbox_worker_protocol::validate_worker_relative_path(relative).map_err(protocol_error)?; + let components = relative.split('/').collect::>(); + if components.iter().any(|component| component.is_empty() || *component == "." || *component == "..") { + return Err("worker path contains an invalid component".to_string()); + } + Ok(components) +} + +fn validate_manifest_structure(entries: &[WorkerUploadEntry]) -> Result<(), String> { + let mut kinds = BTreeMap::new(); + for entry in entries { + if kinds.insert(entry.path().as_str(), entry.kind()).is_some() { + return Err(format!("worker manifest has a duplicate path: {}", entry.path().as_str())); + } + let mut prefix = String::new(); + let parts = entry.path().as_str().split('/').collect::>(); + for (index, component) in parts.iter().enumerate().take(parts.len().saturating_sub(1)) { + if index > 0 { + prefix.push('/'); + } + prefix.push_str(component); + if matches!(kinds.get(prefix.as_str()), Some(WorkerEntryKind::File)) { + return Err(format!("worker manifest path collides with a file: {}", entry.path().as_str())); + } + } + } + for entry in entries { + let parts = entry.path().as_str().split('/').collect::>(); + let mut prefix = String::new(); + for component in parts.iter().take(parts.len().saturating_sub(1)) { + if !prefix.is_empty() { + prefix.push('/'); + } + prefix.push_str(component); + if kinds.get(prefix.as_str()) != Some(&WorkerEntryKind::Directory) { + return Err(format!("worker manifest is missing directory: {prefix}")); + } + } + } + Ok(()) +} + +fn copy_and_verify(source: &File, destination: &File, entry: &WorkerUploadEntry) -> Result<(), String> { + let mut source = source.try_clone().map_err(|error| format!("clone stored worker file: {error}"))?; + let mut destination = destination.try_clone().map_err(|error| format!("clone materialized worker file: {error}"))?; + let mut buffer = vec![0u8; COPY_BUFFER_BYTES]; + let mut hasher = Sha256::new(); + let mut copied = 0u64; + loop { + let count = source.read(&mut buffer).map_err(|error| format!("read stored worker file: {error}"))?; + if count == 0 { + break; + } + copied = copied.checked_add(count as u64).ok_or_else(|| "worker materialized size overflow".to_string())?; + if copied > entry.size() { + return Err(format!("stored worker file exceeds manifest: {}", entry.path().as_str())); + } + hasher.update(&buffer[..count]); + destination.write_all(&buffer[..count]).map_err(|error| format!("write materialized worker file: {error}"))?; + } + let expected = entry.digest().ok_or_else(|| "worker materialized file has no digest".to_string())?; + if copied != entry.size() || hasher.finalize().as_slice() != expected { + return Err(format!("worker materialized file digest mismatch: {}", entry.path().as_str())); + } + platform::sync_fd(&destination).map_err(|error| format!("flush materialized worker file: {error}"))?; + Ok(()) +} + +fn validate_regular_file(metadata: &libc::stat, size: u64, mode: u32, path: &str) -> Result<(), String> { + if metadata.st_mode & libc::S_IFMT != libc::S_IFREG || metadata.st_nlink != 1 { + return Err(format!("worker file is not a private regular file: {path}")); + } + if metadata.st_size < 0 || metadata.st_size as u64 != size { + return Err(format!("worker file size mismatch: {path}")); + } + if metadata.st_mode & 0o777 != mode as libc::mode_t & 0o777 { + return Err(format!("worker file mode mismatch: {path}")); + } + Ok(()) +} + +fn read_bounded(mut file: File, maximum: usize) -> Result, String> { + let mut bytes = Vec::new(); + let mut limited = (&mut file).take(maximum as u64 + 1); + limited.read_to_end(&mut bytes).map_err(|error| format!("read worker state: {error}"))?; + if bytes.len() > maximum { + return Err("worker state exceeds its maximum length".to_string()); + } + Ok(bytes) +} + +fn put_u32(bytes: &mut Vec, value: usize) -> Result<(), String> { + bytes.extend_from_slice(&u32::try_from(value).map_err(|_| "worker manifest count does not fit in u32".to_string())?.to_le_bytes()); + Ok(()) +} + +fn put_string(bytes: &mut Vec, value: &str) -> Result<(), String> { + put_u32(bytes, value.len())?; + bytes.extend_from_slice(value.as_bytes()); + Ok(()) +} + +fn require_nonzero_id(bytes: [u8; 16], label: &str) -> Result<(), String> { + if bytes == [0; 16] { + return Err(format!("{label} ID must be nonzero")); + } + Ok(()) +} + +fn hex_id(bytes: [u8; 16]) -> String { + bytes.iter().map(|byte| format!("{byte:02x}")).collect() +} + +fn is_hex_id(value: &str) -> bool { + value.len() == 32 && value.bytes().all(|byte| byte.is_ascii_hexdigit()) +} + +fn protocol_error(error: WorkerProtocolError) -> String { + error.to_string() +} + +pub(crate) fn open_relative_directory(root: &File, relative: &str) -> Result { + if relative.is_empty() { + return root.try_clone().map_err(|error| format!("clone worker cwd: {error}")); + } + let components = components(relative)?; + open_relative_directory_components(root, &components) +} + +pub(crate) fn next_job_id() -> u64 { + NEXT_JOB_ID.fetch_add(1, Ordering::Relaxed) +} + +#[cfg(test)] +#[path = "storage_ut.rs"] +mod tests; diff --git a/crates/bunkerbox-worker/src/storage_ut.rs b/crates/bunkerbox-worker/src/storage_ut.rs new file mode 100644 index 0000000..4b9984e --- /dev/null +++ b/crates/bunkerbox-worker/src/storage_ut.rs @@ -0,0 +1,104 @@ +use super::*; +use crate::platform; +use crate::process::JobWorkspace; +use bunkerbox_worker_protocol::{WorkerEntryKind, WorkerRelativePath, WorkerSessionId, WorkerUploadEntry, WorkerUploadId}; +use sha2::{Digest, Sha256}; +use std::fs; +use std::io::Read; +use std::os::unix::fs::PermissionsExt; +use tempfile::{tempdir, TempDir}; + +const SESSION: WorkerSessionId = WorkerSessionId([1; 16]); +const UPLOAD: WorkerUploadId = WorkerUploadId([2; 16]); + +fn store_fixture() -> (TempDir, UploadStore) { + let temp = tempdir().unwrap(); + let root_path = temp.path().join("root"); + fs::create_dir(&root_path).unwrap(); + fs::set_permissions(&root_path, fs::Permissions::from_mode(0o700)).unwrap(); + let root = platform::open_root(&root_path).unwrap(); + let store = UploadStore::new(&root).unwrap(); + (temp, store) +} + +fn entries(contents: &[u8]) -> Vec { + vec![ + WorkerUploadEntry::directory("src", 0o755).unwrap(), + WorkerUploadEntry::file("src/input", 0o644, contents.len() as u64, Sha256::digest(contents).into()).unwrap(), + ] +} + +#[test] +fn completed_upload_is_reopened_and_materialized_by_opaque_identity() { + let (_temp, store) = store_fixture(); + let contents = b"stored"; + let mut transaction = store.begin(SESSION, UPLOAD, entries(contents)).unwrap(); + transaction.accept_chunk(&WorkerRelativePath::new("src/input").unwrap(), 0, contents).unwrap(); + transaction.commit().unwrap(); + drop(transaction); + + let stored = store.open_completed(SESSION, UPLOAD).unwrap(); + let jobs = store.jobs_directory().unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + stored.materialize(job.root()).unwrap(); + let file = open_relative_file(job.root(), "src/input").unwrap(); + let mut actual = Vec::new(); + (&file).read_to_end(&mut actual).unwrap(); + assert_eq!(actual, contents); + assert!(store.open_completed(WorkerSessionId([9; 16]), UPLOAD).is_err()); +} + +#[test] +fn incomplete_and_digest_failed_uploads_are_removed_and_token_cannot_be_reused_while_active() { + let (_temp, store) = store_fixture(); + let contents = b"stored"; + { + let mut transaction = store.begin(SESSION, UPLOAD, entries(contents)).unwrap(); + transaction.accept_chunk(&WorkerRelativePath::new("src/input").unwrap(), 1, contents).unwrap_err(); + } + let mut transaction = store.begin(SESSION, UPLOAD, entries(contents)).unwrap(); + transaction.accept_chunk(&WorkerRelativePath::new("src/input").unwrap(), 0, b"wrong!").unwrap(); + assert!(store.begin(SESSION, UPLOAD, entries(contents)).is_err()); + assert!(transaction.commit().is_err()); + drop(transaction); + assert!(store.begin(SESSION, UPLOAD, entries(contents)).is_ok()); +} + +#[test] +fn manifest_requires_declared_parent_directories_and_cleanup_is_token_scoped() { + let (_temp, store) = store_fixture(); + let data = b"x"; + let missing_parent = vec![WorkerUploadEntry::file("missing/file", 0o644, 1, Sha256::digest(data).into()).unwrap()]; + assert!(store.begin(SESSION, UPLOAD, missing_parent).is_err()); + + let mut transaction = store.begin(SESSION, UPLOAD, entries(data)).unwrap(); + transaction.accept_chunk(&WorkerRelativePath::new("src/input").unwrap(), 0, data).unwrap(); + transaction.commit().unwrap(); + drop(transaction); + store.cleanup(SESSION, UPLOAD).unwrap(); + assert!(store.open_completed(SESSION, UPLOAD).is_err()); +} + +#[test] +fn replaced_stored_file_symlink_is_rejected_during_materialization() { + let (temp, store) = store_fixture(); + let data = b"x"; + let mut transaction = store.begin(SESSION, UPLOAD, entries(data)).unwrap(); + transaction.accept_chunk(&WorkerRelativePath::new("src/input").unwrap(), 0, data).unwrap(); + transaction.commit().unwrap(); + drop(transaction); + let files = + temp.path().join("root/.bunkerbox-worker/sessions/01010101010101010101010101010101/uploads/02020202020202020202020202020202/files/src/input"); + fs::remove_file(&files).unwrap(); + #[cfg(unix)] + std::os::unix::fs::symlink("/etc/passwd", &files).unwrap(); + let stored = store.open_completed(SESSION, UPLOAD).unwrap(); + let jobs = store.jobs_directory().unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + assert!(stored.materialize(job.root()).is_err()); +} + +#[test] +fn unsupported_entry_kind_is_not_accepted_by_the_manifest_constructor() { + assert!(WorkerUploadEntry::new("node", WorkerEntryKind::Directory, 0o755, 1, None).is_err()); +} diff --git a/crates/bunkerbox-worker/src/worker.rs b/crates/bunkerbox-worker/src/worker.rs new file mode 100644 index 0000000..cd0d6c2 --- /dev/null +++ b/crates/bunkerbox-worker/src/worker.rs @@ -0,0 +1,307 @@ +use crate::process::{self, JobWorkspace, OutputSink}; +use crate::storage::{UploadStore, UploadTransaction}; +use bunkerbox_worker_protocol::{ + WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadId, MAX_WORKER_ERROR_BYTES, +}; +use std::io::{self, Read, Write}; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::mpsc::{self, Receiver}; +use std::sync::{Arc, Mutex}; + +pub struct FrameWriter { + writer: Mutex, +} + +impl FrameWriter { + pub fn new(writer: W) -> Self { + Self { writer: Mutex::new(writer) } + } + + pub fn send_message(&self, message: &WorkerMessage) -> Result<(), String> { + let mut writer = self.writer.lock().map_err(|_| "worker protocol writer lock poisoned".to_string())?; + message.write_blocking(&mut *writer).map_err(|error| error.to_string()) + } + + #[cfg(test)] + pub(crate) fn into_inner(self) -> Result { + self.writer.into_inner().map_err(|_| "worker protocol writer lock poisoned".to_string()) + } +} + +impl OutputSink for FrameWriter { + fn send(&self, message: WorkerMessage) -> Result<(), String> { + self.send_message(&message) + } +} + +pub struct WorkerService { + store: UploadStore, +} + +impl WorkerService { + pub fn new(root: &std::fs::File) -> Result { + let store = UploadStore::new(root)?; + let jobs = store.jobs_directory()?; + process::cleanup_stale_jobs(&jobs).map_err(|error| format!("clean stale worker jobs: {error}"))?; + Ok(Self { store }) + } + + pub fn run(&self, input: R, writer: &FrameWriter) -> Result<(), String> { + let input = InputChannel::spawn(input); + let Some(hello) = input.next()? else { + return Ok(()); + }; + let (request_id, session_id, response, version) = match hello { + WorkerMessage::Hello { request_id, session_id, response, version } => (request_id, session_id, response, version), + _ => return Err("worker did not receive Hello first".to_string()), + }; + if response { + return Err("worker received a Hello response instead of a request".to_string()); + } + if version != bunkerbox_worker_protocol::WORKER_PROTOCOL_VERSION { + return Err(format!("unsupported worker protocol version: {version}")); + } + if session_id.0 == [0; 16] { + return Err("worker session ID must be nonzero".to_string()); + } + writer.send_message(&WorkerMessage::hello(request_id, session_id, true))?; + + let mut active_upload: Option = None; + loop { + let Some(message) = input.next()? else { + return Ok(()); + }; + if message.session_id() != session_id { + send_error( + writer, + message.request_id(), + session_id, + WorkerOperation::Protocol, + WorkerErrorKind::WorkerProtocol, + "worker session correlation mismatch", + )?; + return Ok(()); + } + + if let Some(active) = active_upload.as_mut() { + match message { + WorkerMessage::UploadFileChunk { request_id: received_request, session_id: received_session, upload_id, path, offset, data } => { + if received_request != active.request_id || received_session != session_id || upload_id != active.upload_id { + send_error( + writer, + received_request, + session_id, + WorkerOperation::Upload, + WorkerErrorKind::WorkerProtocol, + "worker upload correlation mismatch", + )?; + return Ok(()); + } + if let Err(error) = active.transaction.accept_chunk(&path, offset, &data) { + send_error(writer, received_request, session_id, WorkerOperation::Upload, WorkerErrorKind::Upload, &error)?; + return Ok(()); + } + } + WorkerMessage::UploadComplete { request_id: received_request, session_id: received_session, upload_id } => { + if received_request != active.request_id || received_session != session_id || upload_id != active.upload_id { + send_error( + writer, + received_request, + session_id, + WorkerOperation::Upload, + WorkerErrorKind::WorkerProtocol, + "worker upload completion correlation mismatch", + )?; + return Ok(()); + } + if let Err(error) = active.transaction.commit() { + send_error(writer, received_request, session_id, WorkerOperation::Upload, WorkerErrorKind::Upload, &error)?; + return Ok(()); + } + writer.send_message(&WorkerMessage::UploadComplete { request_id: received_request, session_id, upload_id })?; + active_upload = None; + } + _ => { + send_error( + writer, + message.request_id(), + session_id, + WorkerOperation::Upload, + WorkerErrorKind::WorkerProtocol, + "unexpected message during worker upload", + )?; + return Ok(()); + } + } + continue; + } + + match message { + WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } => match self.store.begin(session_id, upload_id, entries) { + Ok(transaction) => active_upload = Some(ActiveUpload { request_id, upload_id, transaction }), + Err(error) => { + send_error(writer, request_id, session_id, WorkerOperation::Upload, WorkerErrorKind::Upload, &error)?; + return Ok(()); + } + }, + WorkerMessage::Build { request_id, session_id, build } => { + self.handle_build(&input, writer, request_id, session_id, build)?; + return Ok(()); + } + _ => { + send_error( + writer, + message.request_id(), + session_id, + WorkerOperation::Protocol, + WorkerErrorKind::WorkerProtocol, + "unexpected worker message", + )?; + return Ok(()); + } + } + } + } + + fn handle_build( + &self, input: &InputChannel, writer: &FrameWriter, request_id: WorkerRequestId, session_id: WorkerSessionId, + build: bunkerbox_worker_protocol::WorkerBuild, + ) -> Result<(), String> { + let upload = match self.store.open_completed(session_id, build.upload_token()) { + Ok(upload) => upload, + Err(error) => { + send_error(writer, request_id, session_id, WorkerOperation::Upload, WorkerErrorKind::Upload, &error)?; + return Ok(()); + } + }; + let jobs = self.store.jobs_directory()?; + let job = match JobWorkspace::create(&jobs) { + Ok(job) => job, + Err(error) => { + send_error(writer, request_id, session_id, WorkerOperation::Build, WorkerErrorKind::Build, &error)?; + return Ok(()); + } + }; + if let Err(error) = upload.materialize(job.root()) { + send_error(writer, request_id, session_id, WorkerOperation::Upload, WorkerErrorKind::Upload, &error)?; + return Ok(()); + } + let exit_code = match process::execute_build(&job, &build, request_id, session_id, writer, &|| input.disconnected()) { + Ok(exit_code) => exit_code, + Err(error) => { + send_error(writer, request_id, session_id, WorkerOperation::Build, WorkerErrorKind::Build, &error)?; + return Ok(()); + } + }; + writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, exit_code))?; + + let Some(cleanup) = input.next()? else { + return Ok(()); + }; + match cleanup { + WorkerMessage::Cleanup { request_id: cleanup_request, session_id: cleanup_session, upload_token } + if cleanup_request == request_id && cleanup_session == session_id && upload_token == build.upload_token() => + { + match self.store.cleanup(session_id, upload_token) { + Ok(()) => writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0)), + Err(error) => send_error(writer, request_id, session_id, WorkerOperation::Cleanup, WorkerErrorKind::Cleanup, &error), + } + } + other => send_error( + writer, + other.request_id(), + session_id, + WorkerOperation::Cleanup, + WorkerErrorKind::WorkerProtocol, + "unexpected worker cleanup message", + ), + } + } +} + +struct InputChannel { + receiver: Receiver, String>>, + eof_seen: Arc, + pending: Arc, +} + +impl InputChannel { + fn spawn(mut input: R) -> Self { + let (sender, receiver) = mpsc::channel(); + let eof_seen = Arc::new(AtomicBool::new(false)); + let eof_seen_for_thread = eof_seen.clone(); + let pending = Arc::new(AtomicUsize::new(0)); + let pending_for_thread = pending.clone(); + std::thread::spawn(move || loop { + match bunkerbox_worker_protocol::WorkerMessage::read_blocking_optional(&mut input) { + Ok(Some(message)) => { + pending_for_thread.fetch_add(1, Ordering::Release); + if sender.send(Ok(Some(message))).is_err() { + return; + } + } + Ok(None) => { + eof_seen_for_thread.store(true, Ordering::Release); + let _ = sender.send(Ok(None)); + return; + } + Err(error) => { + eof_seen_for_thread.store(true, Ordering::Release); + let _ = sender.send(Err(error.to_string())); + return; + } + } + }); + Self { receiver, eof_seen, pending } + } + + fn next(&self) -> Result, String> { + let result = self.receiver.recv().map_err(|_| "worker input reader stopped".to_string())?; + if matches!(&result, Ok(Some(_))) { + self.pending.fetch_sub(1, Ordering::AcqRel); + } + result + } + + fn disconnected(&self) -> bool { + self.eof_seen.load(Ordering::Acquire) && self.pending.load(Ordering::Acquire) == 0 + } +} + +struct ActiveUpload { + request_id: WorkerRequestId, + upload_id: WorkerUploadId, + transaction: UploadTransaction, +} + +fn send_error( + writer: &FrameWriter, request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, kind: WorkerErrorKind, + message: &str, +) -> Result<(), String> { + let message = truncate_message(message); + writer.send_message(&WorkerMessage::error(request_id, session_id, operation, kind, message)) +} + +fn truncate_message(message: &str) -> &str { + if message.len() <= MAX_WORKER_ERROR_BYTES { + return message; + } + let mut end = MAX_WORKER_ERROR_BYTES; + while !message.is_char_boundary(end) { + end -= 1; + } + &message[..end] +} + +pub fn run_stdio(root: &std::path::Path) -> Result<(), String> { + let root = crate::platform::open_root(root)?; + let service = WorkerService::new(&root)?; + let stdin = io::stdin(); + let input = stdin; + let writer = Arc::new(FrameWriter::new(io::stdout())); + service.run(input, writer.as_ref()) +} + +#[cfg(test)] +#[path = "worker_ut.rs"] +mod tests; diff --git a/crates/bunkerbox-worker/src/worker_ut.rs b/crates/bunkerbox-worker/src/worker_ut.rs new file mode 100644 index 0000000..9798d3c --- /dev/null +++ b/crates/bunkerbox-worker/src/worker_ut.rs @@ -0,0 +1,185 @@ +use super::*; +use crate::platform; +use bunkerbox_worker_protocol::{ + WorkerBuild, WorkerEntryKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, WorkerUploadId, +}; +use sha2::{Digest, Sha256}; +use std::fs; +use std::io::Cursor; +use std::os::unix::fs::PermissionsExt; +use std::path::Path; +use tempfile::{tempdir, TempDir}; + +const REQUEST_ID: WorkerRequestId = WorkerRequestId([1; 16]); +const SESSION_ID: WorkerSessionId = WorkerSessionId([2; 16]); +const UPLOAD_ID: WorkerUploadId = WorkerUploadId([3; 16]); + +struct Fixture { + _temp: TempDir, + root: std::fs::File, + script: String, +} + +fn fixture() -> Fixture { + let temp = tempdir().unwrap(); + let root_path = temp.path().join("worker-root"); + fs::create_dir(&root_path).unwrap(); + fs::set_permissions(&root_path, fs::Permissions::from_mode(0o700)).unwrap(); + let script_path = temp.path().join("tool.sh"); + fs::write(&script_path, b"#!/bin/sh\nprintf '%s\\n' \"$CHECK\"\nprintf 'stdout:%s\\n' \"$1\"\nprintf 'stderr:%s\\n' \"$2\" >&2\nexit 7\n") + .unwrap(); + fs::set_permissions(&script_path, fs::Permissions::from_mode(0o755)).unwrap(); + let root = platform::open_root(&root_path).unwrap(); + Fixture { _temp: temp, root, script: script_path.to_string_lossy().into_owned() } +} + +fn file_entry(path: &str, contents: &[u8]) -> WorkerUploadEntry { + let digest: [u8; 32] = Sha256::digest(contents).into(); + WorkerUploadEntry::file(path, 0o644, contents.len() as u64, digest).unwrap() +} + +fn upload_messages(contents: &[u8]) -> Vec { + vec![ + WorkerMessage::hello(REQUEST_ID, SESSION_ID, false), + WorkerMessage::UploadBegin { + request_id: REQUEST_ID, + session_id: SESSION_ID, + upload_id: UPLOAD_ID, + entries: vec![WorkerUploadEntry::directory("src", 0o755).unwrap(), file_entry("src/input.txt", contents)], + }, + WorkerMessage::UploadFileChunk { + request_id: REQUEST_ID, + session_id: SESSION_ID, + upload_id: UPLOAD_ID, + path: bunkerbox_worker_protocol::WorkerRelativePath::new("src/input.txt").unwrap(), + offset: 0, + data: contents.to_vec(), + }, + WorkerMessage::UploadComplete { request_id: REQUEST_ID, session_id: SESSION_ID, upload_id: UPLOAD_ID }, + ] +} + +fn encode_messages(messages: &[WorkerMessage]) -> Vec { + messages.iter().flat_map(|message| message.encode().unwrap()).collect() +} + +fn run_messages(service: &WorkerService, messages: Vec, writer: &FrameWriter) -> Result<(), String> { + service.run(Cursor::new(encode_messages(&messages)), writer) +} + +fn decode_output(bytes: &[u8]) -> Vec { + let mut reader = Cursor::new(bytes); + let mut messages = Vec::new(); + loop { + match WorkerMessage::read_blocking_optional(&mut reader).unwrap() { + Some(message) => messages.push(message), + None => return messages, + } + } +} + +fn output_bytes(writer: FrameWriter>) -> Vec { + writer.into_inner().unwrap() +} + +fn build_message(script: &str) -> WorkerMessage { + WorkerMessage::Build { + request_id: WorkerRequestId([4; 16]), + session_id: SESSION_ID, + build: WorkerBuild::new( + "tool", + script, + vec!["literal space".to_string(), "$(not-a-shell-expansion)".to_string()], + "src", + vec![("CHECK".to_string(), "guest".to_string()), ("PATH".to_string(), "guest-path".to_string())], + vec![("CHECK".to_string(), "target".to_string()), ("PATH".to_string(), "/usr/bin:/bin".to_string())], + UPLOAD_ID, + ) + .unwrap(), + } +} + +#[test] +fn upload_persists_across_worker_instances_and_build_cleans_it() { + let fixture = fixture(); + let contents = b"snapshot contents\n"; + let upload_service = WorkerService::new(&fixture.root).unwrap(); + let upload_writer = FrameWriter::new(Vec::new()); + run_messages(&upload_service, upload_messages(contents), &upload_writer).unwrap(); + let upload_output = decode_output(&output_bytes(upload_writer)); + assert!(matches!(upload_output.as_slice(), [WorkerMessage::Hello { response: true, .. }, WorkerMessage::UploadComplete { .. }])); + drop(upload_service); + + let build_service = WorkerService::new(&fixture.root).unwrap(); + let build = build_message(&fixture.script); + let WorkerMessage::Build { request_id, session_id, build } = build.clone() else { unreachable!() }; + let build_writer = FrameWriter::new(Vec::new()); + let messages = vec![ + WorkerMessage::hello(WorkerRequestId([4; 16]), SESSION_ID, false), + WorkerMessage::Build { request_id, session_id, build: build.clone() }, + WorkerMessage::Cleanup { request_id, session_id, upload_token: UPLOAD_ID }, + ]; + run_messages(&build_service, messages, &build_writer).unwrap(); + let output = decode_output(&output_bytes(build_writer)); + assert!(output + .iter() + .any(|message| matches!(message, WorkerMessage::Stdout { data, .. } if data.windows(b"target".len()).any(|window| window == b"target")))); + assert!(output.iter().any(|message| matches!(message, WorkerMessage::Stdout { data, .. } if data.windows(b"stdout:literal space".len()).any(|window| window == b"stdout:literal space")))); + assert!(output.iter().any(|message| matches!(message, WorkerMessage::Stderr { data, .. } if data.windows(b"stderr:$(not-a-shell-expansion)".len()).any(|window| window == b"stderr:$(not-a-shell-expansion)")))); + assert!(output.iter().any(|message| matches!(message, WorkerMessage::Completed { operation: WorkerOperation::Build, exit_code: 7, .. }))); + assert!(output.iter().any(|message| matches!(message, WorkerMessage::Completed { operation: WorkerOperation::Cleanup, exit_code: 0, .. }))); + + let missing = build_service; + let missing_writer = FrameWriter::new(Vec::new()); + let missing_messages = vec![WorkerMessage::hello(WorkerRequestId([5; 16]), SESSION_ID, false), build_message(&fixture.script)]; + run_messages(&missing, missing_messages, &missing_writer).unwrap(); + let missing_output = decode_output(&output_bytes(missing_writer)); + assert!(missing_output.iter().any(|message| matches!( + message, + WorkerMessage::Error { operation: WorkerOperation::Upload, kind: bunkerbox_worker_protocol::WorkerErrorKind::Upload, .. } + ))); +} + +#[test] +fn malformed_uploads_are_not_buildable() { + let fixture = fixture(); + let service = WorkerService::new(&fixture.root).unwrap(); + let entry = file_entry("src/input.txt", b"expected"); + let messages = vec![ + WorkerMessage::hello(REQUEST_ID, SESSION_ID, false), + WorkerMessage::UploadBegin { + request_id: REQUEST_ID, + session_id: SESSION_ID, + upload_id: UPLOAD_ID, + entries: vec![WorkerUploadEntry::directory("src", 0o755).unwrap(), entry], + }, + WorkerMessage::UploadFileChunk { + request_id: REQUEST_ID, + session_id: SESSION_ID, + upload_id: UPLOAD_ID, + path: bunkerbox_worker_protocol::WorkerRelativePath::new("src/input.txt").unwrap(), + offset: 1, + data: b"wrong".to_vec(), + }, + ]; + let writer = FrameWriter::new(Vec::new()); + run_messages(&service, messages, &writer).unwrap(); + let output = decode_output(&output_bytes(writer)); + assert!(output.iter().any(|message| matches!(message, WorkerMessage::Error { operation: WorkerOperation::Upload, .. }))); + drop(service); + + let service = WorkerService::new(&fixture.root).unwrap(); + let writer = FrameWriter::new(Vec::new()); + let messages = vec![WorkerMessage::hello(WorkerRequestId([6; 16]), SESSION_ID, false), build_message(&fixture.script)]; + run_messages(&service, messages, &writer).unwrap(); + let output = decode_output(&output_bytes(writer)); + assert!(output.iter().any(|message| matches!(message, WorkerMessage::Error { operation: WorkerOperation::Upload, .. }))); +} + +#[test] +fn root_and_protocol_paths_are_confined() { + let fixture = fixture(); + assert!(platform::open_root(Path::new(&fixture.script)).is_err()); + assert!(bunkerbox_worker_protocol::WorkerUploadEntry::file("../escape", 0o644, 1, [0; 32]).is_err()); + assert!(bunkerbox_worker_protocol::WorkerUploadEntry::new("src/a", WorkerEntryKind::Directory, 0o755, 0, None).is_ok()); +} diff --git a/src/ssh.rs b/src/ssh.rs index 5334973..1bde16c 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -268,6 +268,19 @@ impl SshExecution { } }; let upload_id = random_upload_id(); + let entries = match export.entries().iter().map(worker_entry).collect::, _>>() { + Ok(entries) => entries, + Err(error) => { + drop(export); + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(error); + } + }; + if let Err(error) = preflight_upload(request_id, session_id_for(&self.session), upload_id, &entries) { + drop(export); + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(error); + } let mut connection = match WorkerConnection::spawn(&self.factory, &self.target) { Ok(connection) => connection, Err(error) => { @@ -276,8 +289,19 @@ impl SshExecution { return Err(error); } }; - let session_id = WorkerSessionId(self.session.session_id().0); - let operation = upload_and_finish(&mut connection, WorkerRequestId(request_id), session_id, upload_id, &export, &events, !retain_capability); + let session_id = session_id_for(&self.session); + let operation = upload_and_finish( + &mut connection, + UploadPlan { + request_id: WorkerRequestId(request_id), + session_id, + upload_id, + entries: &entries, + export: &export, + events: &events, + cleanup: !retain_capability, + }, + ); let result = match timeout(self.target.resources().sync_timeout(), operation).await { Ok(result) => result, Err(_) => { @@ -478,14 +502,21 @@ impl Drop for WorkerConnection { } } -async fn upload_and_finish( - connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, upload_id: WorkerUploadId, - export: &SnapshotExportClaim, events: &tokio::sync::mpsc::Sender, cleanup: bool, -) -> Result<(), RemoteBackendError> { +struct UploadPlan<'a> { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_id: WorkerUploadId, + entries: &'a [WorkerUploadEntry], + export: &'a SnapshotExportClaim, + events: &'a tokio::sync::mpsc::Sender, + cleanup: bool, +} + +async fn upload_and_finish(connection: &mut WorkerConnection, plan: UploadPlan<'_>) -> Result<(), RemoteBackendError> { + let UploadPlan { request_id, session_id, upload_id, entries, export, events, cleanup } = plan; connection.handshake(request_id, session_id).await?; - let entries = export.entries().iter().map(worker_entry).collect::, _>>()?; let total_bytes = export.total_file_bytes(); - connection.write(&WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries }).await?; + connection.write(&WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries: entries.to_vec() }).await?; let mut completed_bytes = 0u64; for entry in export.entries() { if entry.kind() != SnapshotEntryKind::RegularFile { @@ -592,6 +623,27 @@ fn worker_entry(entry: &crate::snapshot::SnapshotEntry) -> Result Result<(), RemoteBackendError> { + worker_protocol::validate_upload_manifest(entries).map_err(|error| upload_preflight_error(error.to_string()))?; + WorkerMessage::UploadBegin { request_id: WorkerRequestId(request_id), session_id, upload_id, entries: entries.to_vec() } + .encode() + .map_err(|error| upload_preflight_error(error.to_string()))?; + Ok(()) +} + +fn upload_preflight_error(message: String) -> RemoteBackendError { + RemoteBackendError::Transport { + class: RemoteFailureClass::SnapshotTransfer, + message: format!("snapshot cannot fit worker V1 upload protocol: {message}"), + } +} + +fn session_id_for(session: &RunRemoteSession) -> WorkerSessionId { + WorkerSessionId(session.session_id().0) +} + fn check_correlation( received_request: WorkerRequestId, received_session: WorkerSessionId, request_id: WorkerRequestId, session_id: WorkerSessionId, label: &str, ) -> Result<(), RemoteBackendError> { diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index 60284dc..e9bf160 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -5,7 +5,7 @@ use crate::remote::{ }; use crate::remote_target::RemoteTargetConfig; use crate::snapshot::{SnapshotBuilder, SnapshotExclusionPolicy, SnapshotLimits, SnapshotStore}; -use crate::worker_protocol::{self, WorkerErrorKind, WorkerMessage, WorkerOperation}; +use crate::worker_protocol::{self, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerSessionId, WorkerUploadEntry, WorkerUploadId}; use std::fs; use std::path::Path; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; @@ -310,6 +310,15 @@ fn launch_spec_contains_only_fixed_trusted_ssh_arguments() { assert!(!spec.args().iter().any(|arg| arg == "-tt" || arg == "accept-new" || arg.contains("SSH_AUTH_SOCK"))); } +#[test] +fn impossible_v1_uploads_fail_before_transport_spawn() { + let entries = (0..worker_protocol::MAX_WORKER_UPLOAD_ENTRIES + 1) + .map(|index| WorkerUploadEntry::directory(format!("d{index:04}"), 0o755).unwrap()) + .collect::>(); + let error = preflight_upload([1; 16], WorkerSessionId([2; 16]), WorkerUploadId([3; 16]), &entries).unwrap_err(); + assert!(matches!(error, RemoteBackendError::Transport { class: RemoteFailureClass::SnapshotTransfer, .. })); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn worker_version_failure_is_typed_and_never_falls_back() { let fixture = fixture(); diff --git a/src/worker_protocol.rs b/src/worker_protocol.rs index 6547448..5297090 100644 --- a/src/worker_protocol.rs +++ b/src/worker_protocol.rs @@ -1,1447 +1,6 @@ -//! A bounded binary protocol for a host-side worker reached over SSH. -//! -//! This protocol is deliberately independent from `vscomm`. A frame is: -//! -//! ```text -//! magic[4] version[u16] kind[u8] payload_length[u32] payload[payload_length] -//! ``` -//! -//! All integer fields are little-endian. The payload length is checked before -//! allocating a payload buffer, and every length and count inside a payload is -//! checked before it can drive an allocation. +//! Host-facing re-export of the neutral Ticket 11 worker protocol. -use std::collections::BTreeSet; -use std::fmt; -use std::io; - -use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; - -pub const WORKER_PROTOCOL_MAGIC: [u8; 4] = *b"BBWK"; -pub const WORKER_MAGIC: [u8; 4] = WORKER_PROTOCOL_MAGIC; -pub const WORKER_PROTOCOL_VERSION: u16 = 1; -pub const WORKER_VERSION: u16 = WORKER_PROTOCOL_VERSION; -pub const WORKER_FRAME_HEADER_LEN: usize = 4 + 2 + 1 + 4; -pub const WORKER_FRAME_HEADER_SIZE: usize = WORKER_FRAME_HEADER_LEN; -pub const WORKER_ID_LEN: usize = 16; -pub const WORKER_DIGEST_LEN: usize = 32; - -/// Maximum bytes in one worker payload. The declared length is rejected -/// before a buffer of this size is allocated. -pub const MAX_WORKER_FRAME_PAYLOAD: usize = 1024 * 1024; -pub const MAX_WORKER_PAYLOAD: usize = MAX_WORKER_FRAME_PAYLOAD; -pub const MAX_WORKER_FRAME_BYTES: usize = WORKER_FRAME_HEADER_LEN + MAX_WORKER_FRAME_PAYLOAD; -pub const MAX_WORKER_FRAME_LENGTH: usize = MAX_WORKER_FRAME_BYTES; -pub const MAX_WORKER_STRING_BYTES: usize = 4 * 1024; -pub const MAX_WORKER_PATH_BYTES: usize = 4 * 1024; -pub const MAX_WORKER_ENTRY_PATH_BYTES: usize = MAX_WORKER_PATH_BYTES; -pub const MAX_WORKER_PATH_COMPONENT_BYTES: usize = 255; -pub const MAX_WORKER_PATH_DEPTH: usize = 64; -pub const MAX_WORKER_CWD_BYTES: usize = MAX_WORKER_PATH_BYTES; -pub const MAX_WORKER_TOOL_BYTES: usize = 256; -pub const MAX_WORKER_EXECUTABLE_PATH_BYTES: usize = 4 * 1024; -pub const MAX_WORKER_EXECUTABLE_BYTES: usize = MAX_WORKER_EXECUTABLE_PATH_BYTES; -pub const MAX_WORKER_ARG_COUNT: usize = 256; -pub const MAX_WORKER_ARGUMENT_COUNT: usize = MAX_WORKER_ARG_COUNT; -pub const MAX_WORKER_ARG_BYTES: usize = 4 * 1024; -pub const MAX_WORKER_ARG_TOTAL_BYTES: usize = 256 * 1024; -pub const MAX_WORKER_ENV_COUNT: usize = 128; -pub const MAX_WORKER_ENVIRONMENT_COUNT: usize = MAX_WORKER_ENV_COUNT; -pub const MAX_WORKER_ENV_KEY_BYTES: usize = 256; -pub const MAX_WORKER_ENV_VALUE_BYTES: usize = 4 * 1024; -pub const MAX_WORKER_ENV_TOTAL_BYTES: usize = 256 * 1024; -pub const MAX_WORKER_UPLOAD_ENTRIES: usize = 4096; -pub const MAX_WORKER_ENTRY_COUNT: usize = MAX_WORKER_UPLOAD_ENTRIES; -pub const MAX_WORKER_UPLOAD_ENTRY_COUNT: usize = MAX_WORKER_UPLOAD_ENTRIES; -pub const MAX_WORKER_FILE_BYTES: u64 = 256 * 1024 * 1024; -pub const MAX_WORKER_TOTAL_UPLOAD_BYTES: u64 = 512 * 1024 * 1024; -pub const MAX_WORKER_MANIFEST_BYTES: usize = 512 * 1024; -pub const MAX_WORKER_CHUNK_BYTES: usize = 64 * 1024; -pub const MAX_WORKER_OUTPUT_BYTES: usize = 64 * 1024; -pub const MAX_WORKER_ERROR_BYTES: usize = 4 * 1024; - -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum WorkerProtocolError { - Io(String), - Invalid(String), -} - -pub type WorkerResult = Result; - -impl WorkerProtocolError { - pub fn is_invalid(&self) -> bool { - matches!(self, Self::Invalid(_)) - } - - pub fn contains(&self, needle: &str) -> bool { - match self { - Self::Io(message) | Self::Invalid(message) => message.contains(needle), - } - } -} - -impl fmt::Display for WorkerProtocolError { - fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::Io(message) => write!(formatter, "worker I/O error: {message}"), - Self::Invalid(message) => write!(formatter, "invalid worker protocol: {message}"), - } - } -} - -impl std::error::Error for WorkerProtocolError {} - -impl From for WorkerProtocolError { - fn from(error: io::Error) -> Self { - Self::Io(error.to_string()) - } -} - -macro_rules! worker_id { - ($name:ident) => { - #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Default)] - pub struct $name(pub [u8; WORKER_ID_LEN]); - - impl $name { - pub const fn new(bytes: [u8; WORKER_ID_LEN]) -> Self { - Self(bytes) - } - - pub const fn as_bytes(&self) -> &[u8; WORKER_ID_LEN] { - &self.0 - } - - pub const fn into_bytes(self) -> [u8; WORKER_ID_LEN] { - self.0 - } - } - - impl From<[u8; WORKER_ID_LEN]> for $name { - fn from(bytes: [u8; WORKER_ID_LEN]) -> Self { - Self(bytes) - } - } - }; -} - -worker_id!(WorkerRequestId); -worker_id!(WorkerSessionId); -worker_id!(WorkerUploadId); - -pub type RequestId = WorkerRequestId; -pub type SessionId = WorkerSessionId; -pub type UploadId = WorkerUploadId; -pub type UploadToken = WorkerUploadId; -pub type WorkerUploadToken = WorkerUploadId; -pub type WorkerDigest = [u8; WORKER_DIGEST_LEN]; - -#[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum WorkerFrameKind { - Hello = 1, - UploadBegin = 2, - UploadEntry = 3, - UploadFileChunk = 4, - UploadComplete = 5, - Build = 6, - Cleanup = 7, - SyncProgress = 8, - Stdout = 9, - Stderr = 10, - Completed = 11, - Error = 12, -} - -impl WorkerFrameKind { - pub const fn as_u8(self) -> u8 { - self as u8 - } - - pub const fn from_u8(value: u8) -> Option { - match value { - 1 => Some(Self::Hello), - 2 => Some(Self::UploadBegin), - 3 => Some(Self::UploadEntry), - 4 => Some(Self::UploadFileChunk), - 5 => Some(Self::UploadComplete), - 6 => Some(Self::Build), - 7 => Some(Self::Cleanup), - 8 => Some(Self::SyncProgress), - 9 => Some(Self::Stdout), - 10 => Some(Self::Stderr), - 11 => Some(Self::Completed), - 12 => Some(Self::Error), - _ => None, - } - } -} - -impl TryFrom for WorkerFrameKind { - type Error = WorkerProtocolError; - - fn try_from(value: u8) -> WorkerResult { - Self::from_u8(value).ok_or_else(|| invalid(format!("unknown worker frame kind: {value}"))) - } -} - -#[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum WorkerOperation { - Protocol = 0, - Upload = 1, - Build = 2, - Cleanup = 3, - Sync = 4, -} - -impl WorkerOperation { - pub const fn as_u8(self) -> u8 { - self as u8 - } - - pub fn from_u8(value: u8) -> WorkerResult { - match value { - 0 => Ok(Self::Protocol), - 1 => Ok(Self::Upload), - 2 => Ok(Self::Build), - 3 => Ok(Self::Cleanup), - 4 => Ok(Self::Sync), - _ => Err(invalid(format!("unknown worker operation: {value}"))), - } - } -} - -#[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum WorkerErrorKind { - WorkerProtocol = 1, - Upload = 2, - Build = 3, - Cleanup = 4, - Sync = 5, -} - -pub type WorkerErrorClass = WorkerErrorKind; - -impl WorkerErrorKind { - #[allow(non_upper_case_globals)] - pub const Protocol: Self = Self::WorkerProtocol; - - pub const fn as_u8(self) -> u8 { - self as u8 - } - - pub fn from_u8(value: u8) -> WorkerResult { - match value { - 1 => Ok(Self::WorkerProtocol), - 2 => Ok(Self::Upload), - 3 => Ok(Self::Build), - 4 => Ok(Self::Cleanup), - 5 => Ok(Self::Sync), - _ => Err(invalid(format!("unknown worker error kind: {value}"))), - } - } -} - -#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct WorkerRelativePath(String); - -impl WorkerRelativePath { - pub fn new(value: impl Into) -> WorkerResult { - let value = value.into(); - validate_relative_path("worker relative path", &value, true)?; - Ok(Self(value)) - } - - fn for_entry(value: impl Into) -> WorkerResult { - let value = value.into(); - validate_relative_path("worker entry path", &value, false)?; - Ok(Self(value)) - } - - pub fn as_str(&self) -> &str { - &self.0 - } - - pub fn is_empty(&self) -> bool { - self.0.is_empty() - } -} - -impl AsRef for WorkerRelativePath { - fn as_ref(&self) -> &str { - self.as_str() - } -} - -impl TryFrom for WorkerRelativePath { - type Error = WorkerProtocolError; - - fn try_from(value: String) -> WorkerResult { - Self::new(value) - } -} - -#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct WorkerTool(String); - -impl WorkerTool { - pub fn new(value: impl Into) -> WorkerResult { - let value = value.into(); - validate_tool(&value)?; - Ok(Self(value)) - } - - pub fn as_str(&self) -> &str { - &self.0 - } -} - -impl AsRef for WorkerTool { - fn as_ref(&self) -> &str { - self.as_str() - } -} - -impl TryFrom for WorkerTool { - type Error = WorkerProtocolError; - - fn try_from(value: String) -> WorkerResult { - Self::new(value) - } -} - -#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct WorkerExecutablePath(String); - -impl WorkerExecutablePath { - pub fn new(value: impl Into) -> WorkerResult { - let value = value.into(); - validate_executable_path(&value)?; - Ok(Self(value)) - } - - pub fn as_str(&self) -> &str { - &self.0 - } -} - -impl AsRef for WorkerExecutablePath { - fn as_ref(&self) -> &str { - self.as_str() - } -} - -impl TryFrom for WorkerExecutablePath { - type Error = WorkerProtocolError; - - fn try_from(value: String) -> WorkerResult { - Self::new(value) - } -} - -#[repr(u8)] -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum WorkerEntryKind { - Directory = 1, - File = 2, -} - -impl WorkerEntryKind { - pub fn from_u8(value: u8) -> WorkerResult { - match value { - 1 => Ok(Self::Directory), - 2 => Ok(Self::File), - _ => Err(invalid(format!("unknown worker entry kind: {value}"))), - } - } - - pub const fn as_u8(self) -> u8 { - self as u8 - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct WorkerUploadEntry { - pub path: WorkerRelativePath, - pub kind: WorkerEntryKind, - pub mode: u32, - pub size: u64, - pub digest: Option, -} - -impl WorkerUploadEntry { - pub fn new(path: impl Into, kind: WorkerEntryKind, mode: u32, size: u64, digest: Option) -> WorkerResult { - let entry = Self { path: WorkerRelativePath::for_entry(path)?, kind, mode, size, digest }; - entry.validate()?; - Ok(entry) - } - - pub fn directory(path: impl Into, mode: u32) -> WorkerResult { - Self::new(path, WorkerEntryKind::Directory, mode, 0, None) - } - - pub fn file(path: impl Into, mode: u32, size: u64, digest: WorkerDigest) -> WorkerResult { - Self::new(path, WorkerEntryKind::File, mode, size, Some(digest)) - } - - pub fn path(&self) -> &WorkerRelativePath { - &self.path - } - - pub fn kind(&self) -> WorkerEntryKind { - self.kind - } - - pub fn mode(&self) -> u32 { - self.mode - } - - pub fn size(&self) -> u64 { - self.size - } - - pub fn digest(&self) -> Option<&WorkerDigest> { - self.digest.as_ref() - } - - pub fn validate(&self) -> WorkerResult<()> { - validate_relative_path("worker entry path", self.path.as_str(), false)?; - if self.mode & !0o7777 != 0 { - return Err(invalid(format!("worker entry mode has unsupported bits: {:o}", self.mode))); - } - - match self.kind { - WorkerEntryKind::Directory => { - if self.size != 0 { - return Err(invalid("worker directory entry must have zero size")); - } - if self.digest.is_some() { - return Err(invalid("worker directory entry must not have a digest")); - } - } - WorkerEntryKind::File => { - if self.size > MAX_WORKER_FILE_BYTES { - return Err(invalid(format!("worker file exceeds maximum size {MAX_WORKER_FILE_BYTES}"))); - } - if self.digest.is_none() { - return Err(invalid("worker file entry is missing a digest")); - } - } - } - Ok(()) - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct WorkerBuild { - pub tool: WorkerTool, - pub trusted_executable: WorkerExecutablePath, - pub argv: Vec, - pub cwd: WorkerRelativePath, - pub guest_env: Vec<(String, String)>, - pub target_env: Vec<(String, String)>, - pub upload_token: WorkerUploadId, -} - -impl WorkerBuild { - pub fn new( - tool: impl Into, trusted_executable: impl Into, argv: Vec, cwd: impl Into, guest_env: Vec<(String, String)>, - target_env: Vec<(String, String)>, upload_token: WorkerUploadId, - ) -> WorkerResult { - let build = Self { - tool: WorkerTool::new(tool)?, - trusted_executable: WorkerExecutablePath::new(trusted_executable)?, - argv, - cwd: WorkerRelativePath::new(cwd)?, - guest_env, - target_env, - upload_token, - }; - build.validate()?; - Ok(build) - } - - pub fn validate(&self) -> WorkerResult<()> { - validate_tool(self.tool.as_str())?; - validate_executable_path(self.trusted_executable.as_str())?; - validate_relative_path("worker cwd", self.cwd.as_str(), true)?; - validate_argv(&self.argv)?; - validate_environment("worker guest environment", &self.guest_env)?; - validate_environment("worker target environment", &self.target_env)?; - Ok(()) - } - - pub fn tool(&self) -> &WorkerTool { - &self.tool - } - - pub fn trusted_executable(&self) -> &WorkerExecutablePath { - &self.trusted_executable - } - - pub fn trusted_executable_path(&self) -> &str { - self.trusted_executable.as_str() - } - - pub fn argv(&self) -> &[String] { - &self.argv - } - - pub fn cwd(&self) -> &WorkerRelativePath { - &self.cwd - } - - pub fn guest_env(&self) -> &[(String, String)] { - &self.guest_env - } - - pub fn target_env(&self) -> &[(String, String)] { - &self.target_env - } - - pub fn upload_token(&self) -> WorkerUploadId { - self.upload_token - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum WorkerMessage { - Hello { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - version: u16, - response: bool, - }, - UploadBegin { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - upload_id: WorkerUploadId, - entries: Vec, - }, - UploadEntry { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - upload_id: WorkerUploadId, - entry_index: u32, - entry: WorkerUploadEntry, - }, - UploadFileChunk { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - upload_id: WorkerUploadId, - path: WorkerRelativePath, - offset: u64, - data: Vec, - }, - UploadComplete { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - upload_id: WorkerUploadId, - }, - Build { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - build: WorkerBuild, - }, - Cleanup { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - upload_token: WorkerUploadId, - }, - SyncProgress { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - upload_id: WorkerUploadId, - completed_bytes: u64, - total_bytes: Option, - }, - Stdout { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - data: Vec, - }, - Stderr { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - data: Vec, - }, - Completed { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - operation: WorkerOperation, - exit_code: i32, - }, - Error { - request_id: WorkerRequestId, - session_id: WorkerSessionId, - operation: WorkerOperation, - kind: WorkerErrorKind, - message: String, - }, -} - -impl WorkerMessage { - pub fn hello(request_id: WorkerRequestId, session_id: WorkerSessionId, response: bool) -> Self { - Self::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response } - } - - pub fn build(request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild) -> Self { - Self::Build { request_id, session_id, build } - } - - pub fn stdout(request_id: WorkerRequestId, session_id: WorkerSessionId, data: Vec) -> Self { - Self::Stdout { request_id, session_id, data } - } - - pub fn stderr(request_id: WorkerRequestId, session_id: WorkerSessionId, data: Vec) -> Self { - Self::Stderr { request_id, session_id, data } - } - - pub fn completed(request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, exit_code: i32) -> Self { - Self::Completed { request_id, session_id, operation, exit_code } - } - - pub fn error( - request_id: WorkerRequestId, session_id: WorkerSessionId, operation: WorkerOperation, kind: WorkerErrorKind, message: impl Into, - ) -> Self { - Self::Error { request_id, session_id, operation, kind, message: message.into() } - } - - pub const fn kind(&self) -> WorkerFrameKind { - match self { - Self::Hello { .. } => WorkerFrameKind::Hello, - Self::UploadBegin { .. } => WorkerFrameKind::UploadBegin, - Self::UploadEntry { .. } => WorkerFrameKind::UploadEntry, - Self::UploadFileChunk { .. } => WorkerFrameKind::UploadFileChunk, - Self::UploadComplete { .. } => WorkerFrameKind::UploadComplete, - Self::Build { .. } => WorkerFrameKind::Build, - Self::Cleanup { .. } => WorkerFrameKind::Cleanup, - Self::SyncProgress { .. } => WorkerFrameKind::SyncProgress, - Self::Stdout { .. } => WorkerFrameKind::Stdout, - Self::Stderr { .. } => WorkerFrameKind::Stderr, - Self::Completed { .. } => WorkerFrameKind::Completed, - Self::Error { .. } => WorkerFrameKind::Error, - } - } - - pub fn request_id(&self) -> WorkerRequestId { - match self { - Self::Hello { request_id, .. } - | Self::UploadBegin { request_id, .. } - | Self::UploadEntry { request_id, .. } - | Self::UploadFileChunk { request_id, .. } - | Self::UploadComplete { request_id, .. } - | Self::Build { request_id, .. } - | Self::Cleanup { request_id, .. } - | Self::SyncProgress { request_id, .. } - | Self::Stdout { request_id, .. } - | Self::Stderr { request_id, .. } - | Self::Completed { request_id, .. } - | Self::Error { request_id, .. } => *request_id, - } - } - - pub fn session_id(&self) -> WorkerSessionId { - match self { - Self::Hello { session_id, .. } - | Self::UploadBegin { session_id, .. } - | Self::UploadEntry { session_id, .. } - | Self::UploadFileChunk { session_id, .. } - | Self::UploadComplete { session_id, .. } - | Self::Build { session_id, .. } - | Self::Cleanup { session_id, .. } - | Self::SyncProgress { session_id, .. } - | Self::Stdout { session_id, .. } - | Self::Stderr { session_id, .. } - | Self::Completed { session_id, .. } - | Self::Error { session_id, .. } => *session_id, - } - } - - pub fn upload_id(&self) -> Option { - match self { - Self::UploadBegin { upload_id, .. } - | Self::UploadEntry { upload_id, .. } - | Self::UploadFileChunk { upload_id, .. } - | Self::UploadComplete { upload_id, .. } - | Self::SyncProgress { upload_id, .. } => Some(*upload_id), - Self::Build { build, .. } => Some(build.upload_token), - Self::Cleanup { upload_token, .. } => Some(*upload_token), - Self::Hello { .. } | Self::Stdout { .. } | Self::Stderr { .. } | Self::Completed { .. } | Self::Error { .. } => None, - } - } - - pub fn validate(&self) -> WorkerResult<()> { - match self { - Self::Hello { version, .. } => { - if *version != WORKER_PROTOCOL_VERSION { - return Err(invalid(format!("unsupported worker hello version: {version}"))); - } - } - Self::UploadBegin { entries, .. } => { - validate_upload_manifest(entries)?; - } - Self::UploadEntry { entry_index, entry, .. } => { - validate_entry_index(*entry_index)?; - entry.validate()?; - } - Self::UploadFileChunk { path, offset, data, .. } => validate_chunk(path, *offset, data)?, - Self::UploadComplete { .. } => {} - Self::Build { build, .. } => build.validate()?, - Self::Cleanup { .. } => {} - Self::SyncProgress { completed_bytes, total_bytes, .. } => validate_progress(*completed_bytes, *total_bytes)?, - Self::Stdout { data, .. } | Self::Stderr { data, .. } => { - if data.len() > MAX_WORKER_OUTPUT_BYTES { - return Err(invalid(format!("worker output exceeds maximum length {MAX_WORKER_OUTPUT_BYTES}"))); - } - } - Self::Completed { operation, .. } => { - if *operation == WorkerOperation::Protocol { - return Err(invalid("worker completion cannot use protocol operation")); - } - } - Self::Error { message, .. } => validate_error_message(message)?, - } - Ok(()) - } - - pub fn encode(&self) -> WorkerResult> { - self.validate()?; - let mut payload = WireWriter::new(); - encode_payload(self, &mut payload)?; - let payload = payload.finish()?; - - let declared = u32::try_from(payload.len()).map_err(|_| invalid("worker payload length does not fit in u32"))?; - let mut frame = Vec::with_capacity(WORKER_FRAME_HEADER_LEN + payload.len()); - frame.extend_from_slice(&WORKER_PROTOCOL_MAGIC); - frame.extend_from_slice(&WORKER_PROTOCOL_VERSION.to_le_bytes()); - frame.push(self.kind().as_u8()); - frame.extend_from_slice(&declared.to_le_bytes()); - frame.extend_from_slice(&payload); - Ok(frame) - } - - pub fn decode(frame: &[u8]) -> WorkerResult { - let (kind, payload) = split_frame(frame)?; - decode_payload(kind, payload) - } - - pub async fn read_async(reader: &mut R) -> WorkerResult { - let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; - reader.read_exact(&mut header).await.map_err(|error| match error.kind() { - io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), - _ => WorkerProtocolError::Io(error.to_string()), - })?; - let (kind, payload_len) = decode_header(&header)?; - let mut payload = vec![0u8; payload_len]; - reader.read_exact(&mut payload).await.map_err(|error| match error.kind() { - io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), - _ => WorkerProtocolError::Io(error.to_string()), - })?; - decode_payload(kind, &payload) - } - - pub async fn write_async(&self, writer: &mut W) -> WorkerResult<()> { - let frame = self.encode()?; - writer.write_all(&frame).await.map_err(WorkerProtocolError::from)?; - writer.flush().await.map_err(WorkerProtocolError::from) - } -} - -pub async fn read_worker_message(reader: &mut R) -> WorkerResult { - WorkerMessage::read_async(reader).await -} - -pub async fn write_worker_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { - message.write_async(writer).await -} - -pub async fn read_message(reader: &mut R) -> WorkerResult { - read_worker_message(reader).await -} - -pub async fn write_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { - write_worker_message(writer, message).await -} - -pub fn encode_worker_message(message: &WorkerMessage) -> WorkerResult> { - message.encode() -} - -pub fn decode_worker_message(frame: &[u8]) -> WorkerResult { - WorkerMessage::decode(frame) -} - -pub fn validate_worker_relative_path(value: &str) -> WorkerResult<()> { - validate_relative_path("worker relative path", value, false) -} - -pub fn validate_worker_tool(value: &str) -> WorkerResult<()> { - validate_tool(value) -} - -pub fn validate_worker_executable_path(value: &str) -> WorkerResult<()> { - validate_executable_path(value) -} - -pub fn validate_upload_manifest(entries: &[WorkerUploadEntry]) -> WorkerResult { - validate_count(entries.len(), MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; - let mut previous: Option<&str> = None; - let mut total_bytes = 0u64; - let mut manifest_bytes = 4usize; - for entry in entries { - entry.validate()?; - if let Some(previous) = previous { - if previous >= entry.path.as_str() { - return Err(invalid("worker upload manifest must be strictly sorted by path")); - } - } - previous = Some(entry.path.as_str()); - let encoded_entry_bytes = 4usize - .checked_add(entry.path.as_str().len()) - .and_then(|bytes| bytes.checked_add(1 + 4 + 8 + 1)) - .and_then(|bytes| bytes.checked_add(if entry.digest.is_some() { WORKER_DIGEST_LEN } else { 0 })) - .ok_or_else(|| invalid("worker manifest length overflow"))?; - manifest_bytes = manifest_bytes.checked_add(encoded_entry_bytes).ok_or_else(|| invalid("worker manifest length overflow"))?; - if manifest_bytes > MAX_WORKER_MANIFEST_BYTES { - return Err(invalid(format!("worker manifest exceeds maximum length {MAX_WORKER_MANIFEST_BYTES}"))); - } - total_bytes = total_bytes.checked_add(entry.size).ok_or_else(|| invalid("worker upload byte count overflow"))?; - if total_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { - return Err(invalid(format!("worker upload exceeds maximum size {MAX_WORKER_TOTAL_UPLOAD_BYTES}"))); - } - } - Ok(total_bytes) -} - -fn encode_payload(message: &WorkerMessage, writer: &mut WireWriter) -> WorkerResult<()> { - match message { - WorkerMessage::Hello { request_id, session_id, version, response } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.u16(*version)?; - writer.boolean(*response)?; - } - WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.id(upload_id.0)?; - writer.count(entries.len(), MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; - for entry in entries { - encode_entry(writer, entry)?; - } - } - WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.id(upload_id.0)?; - writer.u32(*entry_index)?; - encode_entry(writer, entry)?; - } - WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.id(upload_id.0)?; - writer.string(path.as_str(), MAX_WORKER_PATH_BYTES, "worker chunk path")?; - writer.u64(*offset)?; - writer.blob(data, MAX_WORKER_CHUNK_BYTES, "worker file chunk")?; - } - WorkerMessage::UploadComplete { request_id, session_id, upload_id } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.id(upload_id.0)?; - } - WorkerMessage::Build { request_id, session_id, build } => { - encode_correlation(writer, *request_id, *session_id)?; - encode_build(writer, build)?; - } - WorkerMessage::Cleanup { request_id, session_id, upload_token } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.id(upload_token.0)?; - } - WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes, total_bytes } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.id(upload_id.0)?; - writer.u64(*completed_bytes)?; - match total_bytes { - Some(total_bytes) => { - writer.boolean(true)?; - writer.u64(*total_bytes)?; - } - None => writer.boolean(false)?, - } - } - WorkerMessage::Stdout { request_id, session_id, data } | WorkerMessage::Stderr { request_id, session_id, data } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.blob(data, MAX_WORKER_OUTPUT_BYTES, "worker output")?; - } - WorkerMessage::Completed { request_id, session_id, operation, exit_code } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.u8(operation.as_u8())?; - writer.i32(*exit_code)?; - } - WorkerMessage::Error { request_id, session_id, operation, kind, message } => { - encode_correlation(writer, *request_id, *session_id)?; - writer.u8(operation.as_u8())?; - writer.u8(kind.as_u8())?; - writer.string(message, MAX_WORKER_ERROR_BYTES, "worker error")?; - } - } - Ok(()) -} - -fn decode_payload(kind: WorkerFrameKind, payload: &[u8]) -> WorkerResult { - let mut reader = WireReader::new(payload); - let message = match kind { - WorkerFrameKind::Hello => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let version = reader.u16()?; - let response = reader.boolean("worker hello response")?; - WorkerMessage::Hello { request_id, session_id, version, response } - } - WorkerFrameKind::UploadBegin => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let upload_id = WorkerUploadId(reader.array16()?); - let count = reader.count(MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; - let mut entries = Vec::with_capacity(count); - for _ in 0..count { - entries.push(decode_entry(&mut reader)?); - } - WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } - } - WorkerFrameKind::UploadEntry => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let upload_id = WorkerUploadId(reader.array16()?); - let entry_index = reader.u32()?; - let entry = decode_entry(&mut reader)?; - WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } - } - WorkerFrameKind::UploadFileChunk => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let upload_id = WorkerUploadId(reader.array16()?); - let path = WorkerRelativePath::for_entry(reader.string(MAX_WORKER_PATH_BYTES, "worker chunk path")?)?; - let offset = reader.u64()?; - let data = reader.blob(MAX_WORKER_CHUNK_BYTES, "worker file chunk")?; - WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } - } - WorkerFrameKind::UploadComplete => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let upload_id = WorkerUploadId(reader.array16()?); - WorkerMessage::UploadComplete { request_id, session_id, upload_id } - } - WorkerFrameKind::Build => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - WorkerMessage::Build { request_id, session_id, build: decode_build(&mut reader)? } - } - WorkerFrameKind::Cleanup => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let upload_token = WorkerUploadId(reader.array16()?); - WorkerMessage::Cleanup { request_id, session_id, upload_token } - } - WorkerFrameKind::SyncProgress => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let upload_id = WorkerUploadId(reader.array16()?); - let completed_bytes = reader.u64()?; - let total_bytes = if reader.boolean("worker progress total flag")? { Some(reader.u64()?) } else { None }; - WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes, total_bytes } - } - WorkerFrameKind::Stdout => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - WorkerMessage::Stdout { request_id, session_id, data: reader.blob(MAX_WORKER_OUTPUT_BYTES, "worker stdout")? } - } - WorkerFrameKind::Stderr => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - WorkerMessage::Stderr { request_id, session_id, data: reader.blob(MAX_WORKER_OUTPUT_BYTES, "worker stderr")? } - } - WorkerFrameKind::Completed => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let operation = WorkerOperation::from_u8(reader.u8()?)?; - let exit_code = reader.i32()?; - WorkerMessage::Completed { request_id, session_id, operation, exit_code } - } - WorkerFrameKind::Error => { - let (request_id, session_id) = decode_correlation(&mut reader)?; - let operation = WorkerOperation::from_u8(reader.u8()?)?; - let kind = WorkerErrorKind::from_u8(reader.u8()?)?; - let message = reader.string(MAX_WORKER_ERROR_BYTES, "worker error")?; - WorkerMessage::Error { request_id, session_id, operation, kind, message } - } - }; - reader.finish()?; - message.validate()?; - Ok(message) -} - -fn encode_build(writer: &mut WireWriter, build: &WorkerBuild) -> WorkerResult<()> { - build.validate()?; - writer.string(build.tool.as_str(), MAX_WORKER_TOOL_BYTES, "worker tool")?; - writer.string(build.trusted_executable.as_str(), MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?; - writer.count(build.argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; - for argument in &build.argv { - writer.string(argument, MAX_WORKER_ARG_BYTES, "worker argument")?; - } - writer.string(build.cwd.as_str(), MAX_WORKER_CWD_BYTES, "worker cwd")?; - encode_environment(writer, &build.guest_env, "worker guest environment")?; - encode_environment(writer, &build.target_env, "worker target environment")?; - writer.id(build.upload_token.0)?; - Ok(()) -} - -fn decode_build(reader: &mut WireReader<'_>) -> WorkerResult { - let tool = WorkerTool::new(reader.string(MAX_WORKER_TOOL_BYTES, "worker tool")?)?; - let trusted_executable = WorkerExecutablePath::new(reader.string(MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?)?; - let argument_count = reader.count(MAX_WORKER_ARG_COUNT, "worker argv")?; - let mut argv = Vec::with_capacity(argument_count); - for _ in 0..argument_count { - argv.push(reader.string(MAX_WORKER_ARG_BYTES, "worker argument")?); - } - let cwd = WorkerRelativePath::new(reader.string(MAX_WORKER_CWD_BYTES, "worker cwd")?)?; - let guest_env = decode_environment(reader, "worker guest environment")?; - let target_env = decode_environment(reader, "worker target environment")?; - let upload_token = WorkerUploadId(reader.array16()?); - let build = WorkerBuild { tool, trusted_executable, argv, cwd, guest_env, target_env, upload_token }; - build.validate()?; - Ok(build) -} - -fn encode_environment(writer: &mut WireWriter, environment: &[(String, String)], field: &str) -> WorkerResult<()> { - validate_environment(field, environment)?; - writer.count(environment.len(), MAX_WORKER_ENV_COUNT, field)?; - for (key, value) in environment { - writer.string(key, MAX_WORKER_ENV_KEY_BYTES, "worker environment key")?; - writer.string(value, MAX_WORKER_ENV_VALUE_BYTES, "worker environment value")?; - } - Ok(()) -} - -fn decode_environment(reader: &mut WireReader<'_>, field: &str) -> WorkerResult> { - let count = reader.count(MAX_WORKER_ENV_COUNT, field)?; - let mut environment = Vec::with_capacity(count); - for _ in 0..count { - let key = reader.string(MAX_WORKER_ENV_KEY_BYTES, "worker environment key")?; - let value = reader.string(MAX_WORKER_ENV_VALUE_BYTES, "worker environment value")?; - environment.push((key, value)); - } - validate_environment(field, &environment)?; - Ok(environment) -} - -fn encode_entry(writer: &mut WireWriter, entry: &WorkerUploadEntry) -> WorkerResult<()> { - entry.validate()?; - writer.string(entry.path.as_str(), MAX_WORKER_PATH_BYTES, "worker entry path")?; - writer.u8(entry.kind.as_u8())?; - writer.u32(entry.mode)?; - writer.u64(entry.size)?; - match entry.digest { - Some(digest) => { - writer.boolean(true)?; - writer.bytes(&digest)?; - } - None => writer.boolean(false)?, - } - Ok(()) -} - -fn decode_entry(reader: &mut WireReader<'_>) -> WorkerResult { - let path = reader.string(MAX_WORKER_PATH_BYTES, "worker entry path")?; - let kind = WorkerEntryKind::from_u8(reader.u8()?)?; - let mode = reader.u32()?; - let size = reader.u64()?; - let digest = match reader.boolean("worker entry digest flag")? { - true => Some(reader.array32()?), - false => None, - }; - WorkerUploadEntry::new(path, kind, mode, size, digest) -} - -fn encode_correlation(writer: &mut WireWriter, request_id: WorkerRequestId, session_id: WorkerSessionId) -> WorkerResult<()> { - writer.id(request_id.0)?; - writer.id(session_id.0) -} - -fn decode_correlation(reader: &mut WireReader<'_>) -> WorkerResult<(WorkerRequestId, WorkerSessionId)> { - Ok((WorkerRequestId(reader.array16()?), WorkerSessionId(reader.array16()?))) -} - -fn validate_entry_index(index: u32) -> WorkerResult<()> { - if usize::try_from(index).map_or(true, |index| index >= MAX_WORKER_UPLOAD_ENTRIES) { - return Err(invalid(format!("worker upload entry index exceeds maximum {MAX_WORKER_UPLOAD_ENTRIES}"))); - } - Ok(()) -} - -fn validate_chunk(path: &WorkerRelativePath, offset: u64, data: &[u8]) -> WorkerResult<()> { - validate_relative_path("worker chunk path", path.as_str(), false)?; - if data.is_empty() { - return Err(invalid("worker file chunk must not be empty")); - } - if data.len() > MAX_WORKER_CHUNK_BYTES { - return Err(invalid(format!("worker file chunk exceeds maximum length {MAX_WORKER_CHUNK_BYTES}"))); - } - let end = offset - .checked_add(u64::try_from(data.len()).map_err(|_| invalid("worker file chunk length does not fit in u64"))?) - .ok_or_else(|| invalid("worker file chunk offset overflow"))?; - if end > MAX_WORKER_FILE_BYTES { - return Err(invalid(format!("worker file chunk exceeds maximum file size {MAX_WORKER_FILE_BYTES}"))); - } - Ok(()) -} - -fn validate_progress(completed_bytes: u64, total_bytes: Option) -> WorkerResult<()> { - if completed_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { - return Err(invalid("worker progress exceeds the maximum upload size")); - } - if let Some(total_bytes) = total_bytes { - if total_bytes > MAX_WORKER_TOTAL_UPLOAD_BYTES { - return Err(invalid("worker progress total exceeds the maximum upload size")); - } - if completed_bytes > total_bytes { - return Err(invalid("worker progress exceeds its total")); - } - } - Ok(()) -} - -fn validate_argv(argv: &[String]) -> WorkerResult<()> { - validate_count(argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; - let mut total_bytes = 0usize; - for argument in argv { - validate_text("worker argument", argument, MAX_WORKER_ARG_BYTES)?; - total_bytes = total_bytes.checked_add(argument.len()).ok_or_else(|| invalid("worker argv length overflow"))?; - if total_bytes > MAX_WORKER_ARG_TOTAL_BYTES { - return Err(invalid(format!("worker argv exceeds maximum length {MAX_WORKER_ARG_TOTAL_BYTES}"))); - } - } - Ok(()) -} - -fn validate_environment(field: &str, environment: &[(String, String)]) -> WorkerResult<()> { - validate_count(environment.len(), MAX_WORKER_ENV_COUNT, field)?; - let mut names = BTreeSet::new(); - let mut total_bytes = 0usize; - for (key, value) in environment { - validate_environment_key(key)?; - validate_environment_value(value)?; - total_bytes = total_bytes - .checked_add(key.len()) - .and_then(|bytes| bytes.checked_add(value.len())) - .ok_or_else(|| invalid(format!("{field} length overflow")))?; - if total_bytes > MAX_WORKER_ENV_TOTAL_BYTES { - return Err(invalid(format!("{field} exceeds maximum length {MAX_WORKER_ENV_TOTAL_BYTES}"))); - } - if !names.insert(key.as_str()) { - return Err(invalid(format!("duplicate worker environment key: {key}"))); - } - } - Ok(()) -} - -fn validate_environment_key(key: &str) -> WorkerResult<()> { - validate_text("worker environment key", key, MAX_WORKER_ENV_KEY_BYTES)?; - let mut bytes = key.bytes(); - let Some(first) = bytes.next() else { - return Err(invalid("worker environment key is empty")); - }; - if !(first == b'_' || first.is_ascii_alphabetic()) || !bytes.all(|byte| byte == b'_' || byte.is_ascii_alphanumeric()) { - return Err(invalid("worker environment key is not a valid variable name")); - } - Ok(()) -} - -fn validate_environment_value(value: &str) -> WorkerResult<()> { - validate_text("worker environment value", value, MAX_WORKER_ENV_VALUE_BYTES)?; - if value.chars().any(char::is_control) { - return Err(invalid("worker environment value contains control data")); - } - Ok(()) -} - -fn validate_error_message(message: &str) -> WorkerResult<()> { - validate_text("worker error", message, MAX_WORKER_ERROR_BYTES) -} - -fn validate_tool(value: &str) -> WorkerResult<()> { - validate_text("worker tool", value, MAX_WORKER_TOOL_BYTES)?; - if value.is_empty() - || value == "." - || value == ".." - || !value.bytes().all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.' | b'+')) - { - return Err(invalid("worker tool must be a single safe executable identity")); - } - Ok(()) -} - -fn validate_executable_path(value: &str) -> WorkerResult<()> { - validate_text("worker executable path", value, MAX_WORKER_EXECUTABLE_PATH_BYTES)?; - if !value.starts_with('/') || value == "/" || value.starts_with("//") { - return Err(invalid("worker executable path must be a normalized absolute path")); - } - if value.chars().any(|character| character.is_whitespace() || character.is_control()) { - return Err(invalid("worker executable path must not contain whitespace or control data")); - } - if value.bytes().any(|byte| { - byte < 0x20 - || byte == 0x7f - || matches!(byte, b'\\' | b';' | b'|' | b'&' | b'$' | b'`' | b'<' | b'>' | b'\'' | b'"' | b'(' | b')' | b'[' | b']' | b'{' | b'}') - }) { - return Err(invalid("worker executable path contains unsafe identity data")); - } - for (depth, component) in value.split('/').skip(1).enumerate() { - if depth >= MAX_WORKER_PATH_DEPTH - || component.len() > MAX_WORKER_PATH_COMPONENT_BYTES - || component.is_empty() - || component == "." - || component == ".." - || component.contains(':') - { - return Err(invalid("worker executable path is not normalized")); - } - } - Ok(()) -} - -fn validate_relative_path(field: &str, value: &str, allow_empty: bool) -> WorkerResult<()> { - validate_text(field, value, MAX_WORKER_PATH_BYTES)?; - if value.is_empty() { - if allow_empty { - return Ok(()); - } - return Err(invalid(format!("{field} must not be empty"))); - } - if value.starts_with('/') || value.starts_with('\\') || value.contains('\\') || value.bytes().any(|byte| byte == b':') { - return Err(invalid(format!("{field} must be a normalized relative path"))); - } - for (depth, component) in value.split('/').enumerate() { - if depth >= MAX_WORKER_PATH_DEPTH - || component.len() > MAX_WORKER_PATH_COMPONENT_BYTES - || component.is_empty() - || component == "." - || component == ".." - || component.chars().any(char::is_control) - { - return Err(invalid(format!("{field} must be a normalized relative path"))); - } - } - Ok(()) -} - -fn validate_text(field: &str, value: &str, maximum: usize) -> WorkerResult<()> { - if value.len() > maximum { - return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); - } - if value.as_bytes().contains(&0) { - return Err(invalid(format!("{field} contains a NUL byte"))); - } - Ok(()) -} - -fn validate_count(count: usize, maximum: usize, field: &str) -> WorkerResult<()> { - if count > maximum { - return Err(invalid(format!("{field} exceeds maximum count {maximum}"))); - } - Ok(()) -} - -fn invalid(message: impl Into) -> WorkerProtocolError { - WorkerProtocolError::Invalid(message.into()) -} - -fn split_frame(frame: &[u8]) -> WorkerResult<(WorkerFrameKind, &[u8])> { - if frame.len() < WORKER_FRAME_HEADER_LEN { - return Err(invalid("truncated worker frame header")); - } - let (kind, payload_len) = decode_header(&frame[..WORKER_FRAME_HEADER_LEN])?; - let expected = WORKER_FRAME_HEADER_LEN.checked_add(payload_len).ok_or_else(|| invalid("worker frame length overflow"))?; - if frame.len() < expected { - return Err(invalid("truncated worker frame payload")); - } - if frame.len() > expected { - return Err(invalid("extra bytes after worker frame")); - } - Ok((kind, &frame[WORKER_FRAME_HEADER_LEN..expected])) -} - -fn decode_header(header: &[u8]) -> WorkerResult<(WorkerFrameKind, usize)> { - if header.len() != WORKER_FRAME_HEADER_LEN { - return Err(invalid("invalid worker frame header length")); - } - if header[..4] != WORKER_PROTOCOL_MAGIC { - return Err(invalid("invalid worker frame magic")); - } - let version = u16::from_le_bytes([header[4], header[5]]); - if version != WORKER_PROTOCOL_VERSION { - return Err(invalid(format!("unsupported worker protocol version: {version}"))); - } - let kind = WorkerFrameKind::try_from(header[6])?; - let payload_len = u32::from_le_bytes([header[7], header[8], header[9], header[10]]) as usize; - if payload_len > MAX_WORKER_FRAME_PAYLOAD { - return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); - } - Ok((kind, payload_len)) -} - -struct WireWriter { - bytes: Vec, -} - -impl WireWriter { - fn new() -> Self { - Self { bytes: Vec::new() } - } - - fn finish(self) -> WorkerResult> { - if self.bytes.len() > MAX_WORKER_FRAME_PAYLOAD { - return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); - } - Ok(self.bytes) - } - - fn bytes(&mut self, value: &[u8]) -> WorkerResult<()> { - let new_length = self.bytes.len().checked_add(value.len()).ok_or_else(|| invalid("worker payload length overflow"))?; - if new_length > MAX_WORKER_FRAME_PAYLOAD { - return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); - } - self.bytes.extend_from_slice(value); - Ok(()) - } - - fn id(&mut self, value: [u8; WORKER_ID_LEN]) -> WorkerResult<()> { - self.bytes(&value) - } - - fn u8(&mut self, value: u8) -> WorkerResult<()> { - self.bytes(&[value]) - } - - fn boolean(&mut self, value: bool) -> WorkerResult<()> { - self.u8(u8::from(value)) - } - - fn u16(&mut self, value: u16) -> WorkerResult<()> { - self.bytes(&value.to_le_bytes()) - } - - fn u32(&mut self, value: u32) -> WorkerResult<()> { - self.bytes(&value.to_le_bytes()) - } - - fn u64(&mut self, value: u64) -> WorkerResult<()> { - self.bytes(&value.to_le_bytes()) - } - - fn i32(&mut self, value: i32) -> WorkerResult<()> { - self.bytes(&value.to_le_bytes()) - } - - fn count(&mut self, count: usize, maximum: usize, field: &str) -> WorkerResult<()> { - validate_count(count, maximum, field)?; - self.u32(u32::try_from(count).map_err(|_| invalid(format!("{field} count does not fit in u32")))?) - } - - fn string(&mut self, value: &str, maximum: usize, field: &str) -> WorkerResult<()> { - validate_text(field, value, maximum)?; - let length = u32::try_from(value.len()).map_err(|_| invalid(format!("{field} length does not fit in u32")))?; - self.u32(length)?; - self.bytes(value.as_bytes()) - } - - fn blob(&mut self, value: &[u8], maximum: usize, field: &str) -> WorkerResult<()> { - if value.len() > maximum { - return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); - } - let length = u32::try_from(value.len()).map_err(|_| invalid(format!("{field} length does not fit in u32")))?; - self.u32(length)?; - self.bytes(value) - } -} - -struct WireReader<'a> { - bytes: &'a [u8], - offset: usize, -} - -impl<'a> WireReader<'a> { - fn new(bytes: &'a [u8]) -> Self { - Self { bytes, offset: 0 } - } - - fn take(&mut self, length: usize) -> WorkerResult<&'a [u8]> { - let end = self.offset.checked_add(length).ok_or_else(|| invalid("worker payload length overflow"))?; - if end > self.bytes.len() { - return Err(invalid("truncated worker payload")); - } - let value = &self.bytes[self.offset..end]; - self.offset = end; - Ok(value) - } - - fn u8(&mut self) -> WorkerResult { - Ok(self.take(1)?[0]) - } - - fn boolean(&mut self, field: &str) -> WorkerResult { - match self.u8()? { - 0 => Ok(false), - 1 => Ok(true), - value => Err(invalid(format!("invalid {field} flag: {value}"))), - } - } - - fn u16(&mut self) -> WorkerResult { - let value = self.take(2)?; - Ok(u16::from_le_bytes([value[0], value[1]])) - } - - fn u32(&mut self) -> WorkerResult { - let value = self.take(4)?; - Ok(u32::from_le_bytes([value[0], value[1], value[2], value[3]])) - } - - fn u64(&mut self) -> WorkerResult { - let value = self.take(8)?; - Ok(u64::from_le_bytes(value.try_into().map_err(|_| invalid("invalid worker integer"))?)) - } - - fn i32(&mut self) -> WorkerResult { - let value = self.take(4)?; - Ok(i32::from_le_bytes(value.try_into().map_err(|_| invalid("invalid worker integer"))?)) - } - - fn array16(&mut self) -> WorkerResult<[u8; WORKER_ID_LEN]> { - self.take(WORKER_ID_LEN)?.try_into().map_err(|_| invalid("invalid worker identifier")) - } - - fn array32(&mut self) -> WorkerResult<[u8; WORKER_DIGEST_LEN]> { - self.take(WORKER_DIGEST_LEN)?.try_into().map_err(|_| invalid("invalid worker digest")) - } - - fn count(&mut self, maximum: usize, field: &str) -> WorkerResult { - let count = self.u32()? as usize; - validate_count(count, maximum, field)?; - Ok(count) - } - - fn string(&mut self, maximum: usize, field: &str) -> WorkerResult { - let length = self.u32()? as usize; - if length > maximum { - return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); - } - let value = std::str::from_utf8(self.take(length)?).map_err(|_| invalid(format!("{field} is not valid UTF-8")))?; - validate_text(field, value, maximum)?; - Ok(value.to_owned()) - } - - fn blob(&mut self, maximum: usize, field: &str) -> WorkerResult> { - let length = self.u32()? as usize; - if length > maximum { - return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); - } - Ok(self.take(length)?.to_vec()) - } - - fn finish(self) -> WorkerResult<()> { - if self.offset != self.bytes.len() { - return Err(invalid("extra bytes in worker payload")); - } - Ok(()) - } -} +pub use bunkerbox_worker_protocol::*; #[cfg(test)] #[path = "worker_protocol_ut.rs"] diff --git a/src/worker_protocol_ut.rs b/src/worker_protocol_ut.rs index dfe4c23..461d231 100644 --- a/src/worker_protocol_ut.rs +++ b/src/worker_protocol_ut.rs @@ -1,264 +1,68 @@ use super::*; -use tokio::io::AsyncWriteExt; fn ids() -> (WorkerRequestId, WorkerSessionId, WorkerUploadId) { - (WorkerRequestId([1; 16]), WorkerSessionId([2; 16]), WorkerUploadId([3; 16])) -} - -fn file_entry(path: &str) -> WorkerUploadEntry { - WorkerUploadEntry::file(path, 0o644, 4, [9; WORKER_DIGEST_LEN]).unwrap() -} - -fn build() -> WorkerBuild { - WorkerBuild::new( - "make", - "/usr/bin/make", - vec!["release mode".into(), "$(literal); still one arg".into()], - "src", - vec![("CC".into(), "gcc".into())], - vec![("PATH".into(), "/usr/bin".into())], - WorkerUploadId([3; 16]), - ) - .unwrap() -} - -fn messages() -> Vec { - let (request_id, session_id, upload_id) = ids(); - vec![ - WorkerMessage::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response: false }, - WorkerMessage::UploadBegin { - request_id, - session_id, - upload_id, - entries: vec![WorkerUploadEntry::directory("src", 0o755).unwrap(), file_entry("src/main.rs")], - }, - WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index: 1, entry: file_entry("src/main.rs") }, - WorkerMessage::UploadFileChunk { - request_id, - session_id, - upload_id, - path: WorkerRelativePath::new("src/main.rs").unwrap(), - offset: 0, - data: b"data".to_vec(), - }, - WorkerMessage::UploadComplete { request_id, session_id, upload_id }, - WorkerMessage::Build { request_id, session_id, build: build() }, - WorkerMessage::Cleanup { request_id, session_id, upload_token: upload_id }, - WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 4, total_bytes: Some(8) }, - WorkerMessage::Stdout { request_id, session_id, data: b"out".to_vec() }, - WorkerMessage::Stderr { request_id, session_id, data: b"err".to_vec() }, - WorkerMessage::Completed { request_id, session_id, operation: WorkerOperation::Build, exit_code: -17 }, - WorkerMessage::Error { - request_id, - session_id, - operation: WorkerOperation::Cleanup, - kind: WorkerErrorKind::Cleanup, - message: "cleanup failed".into(), - }, - ] -} - -#[test] -fn all_message_variants_round_trip() { - for message in messages() { - let frame = message.encode().unwrap(); - assert_eq!(&frame[..4], &WORKER_PROTOCOL_MAGIC); - assert_eq!(frame[4..6], WORKER_PROTOCOL_VERSION.to_le_bytes()); - assert_eq!(frame[6], message.kind().as_u8()); - assert_eq!(WorkerMessage::decode(&frame).unwrap(), message); - } -} - -#[tokio::test] -async fn async_helpers_handle_fragmented_duplex_frames() { - let message = WorkerMessage::Build { request_id: ids().0, session_id: ids().1, build: build() }; - let frame = message.encode().unwrap(); - let (mut reader, mut writer) = tokio::io::duplex(3); - let sender = tokio::spawn(async move { - for part in frame.chunks(2) { - writer.write_all(part).await.unwrap(); - tokio::task::yield_now().await; - } - }); - let decoded = WorkerMessage::read_async(&mut reader).await.unwrap(); - sender.await.unwrap(); - assert_eq!(decoded, message); -} - -#[tokio::test] -async fn async_write_helper_emits_a_decodable_frame() { - let message = WorkerMessage::Stdout { request_id: ids().0, session_id: ids().1, data: b"exact".to_vec() }; - let (mut reader, mut writer) = tokio::io::duplex(128); - let expected = message.clone(); - let sender = tokio::spawn(async move { message.write_async(&mut writer).await.unwrap() }); - assert_eq!(WorkerMessage::read_async(&mut reader).await.unwrap(), expected); - sender.await.unwrap(); + (WorkerRequestId([0x11; 16]), WorkerSessionId([0x22; 16]), WorkerUploadId([0x33; 16])) } #[test] -fn malformed_headers_and_lengths_are_rejected() { - let frame = - WorkerMessage::Hello { request_id: ids().0, session_id: ids().1, version: WORKER_PROTOCOL_VERSION, response: false }.encode().unwrap(); - - let mut bad_magic = frame.clone(); - bad_magic[0] ^= 1; - assert!(WorkerMessage::decode(&bad_magic).is_err()); - - let mut bad_version = frame.clone(); - bad_version[4..6].copy_from_slice(&(WORKER_PROTOCOL_VERSION + 1).to_le_bytes()); - assert!(WorkerMessage::decode(&bad_version).is_err()); - - let mut bad_kind = frame.clone(); - bad_kind[6] = 255; - assert!(WorkerMessage::decode(&bad_kind).is_err()); - - let mut oversized = frame[..WORKER_FRAME_HEADER_LEN].to_vec(); - oversized[7..11].copy_from_slice(&((MAX_WORKER_FRAME_PAYLOAD as u32) + 1).to_le_bytes()); - assert!(WorkerMessage::decode(&oversized).is_err()); - - assert!(WorkerMessage::decode(&frame[..frame.len() - 1]).is_err()); - let mut extra = frame.clone(); - extra.push(0); - assert!(WorkerMessage::decode(&extra).is_err()); -} - -#[test] -fn invalid_utf8_and_trailing_payload_are_rejected() { +fn ticket_11_v1_frame_header_and_payload_bytes_remain_stable() { let (request_id, session_id, _) = ids(); - let mut frame = - WorkerMessage::Error { request_id, session_id, operation: WorkerOperation::Build, kind: WorkerErrorKind::Build, message: "x".into() } - .encode() - .unwrap(); - let message_start = WORKER_FRAME_HEADER_LEN + 32 + 1 + 1 + 4; - frame[message_start] = 0xff; - assert!(WorkerMessage::decode(&frame).is_err()); - - let mut hello = WorkerMessage::hello(request_id, session_id, false).encode().unwrap(); - hello.push(0); - assert!(WorkerMessage::decode(&hello).is_err()); -} - -#[test] -fn paths_tools_and_executables_are_strictly_validated() { - for path in ["/absolute", "../parent", "a/../b", "./name", "a//b", "a\\b", ""] { - assert!(validate_worker_relative_path(path).is_err(), "accepted path {path:?}"); - } - assert!(validate_worker_relative_path("src/main.rs").is_ok()); - assert!(validate_worker_relative_path(&"a".repeat(MAX_WORKER_PATH_COMPONENT_BYTES + 1)).is_err()); - assert!(WorkerTool::new("make;rm").is_err()); - assert!(WorkerTool::new("/usr/bin/make").is_err()); - assert!(WorkerExecutablePath::new("make").is_err()); - assert!(WorkerExecutablePath::new("/usr/bin/make -f").is_err()); - assert!(WorkerExecutablePath::new("/usr/bin/../bin/make").is_err()); - assert!(WorkerExecutablePath::new("/usr/bin/make").is_ok()); + let message = WorkerMessage::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response: false }; + let expected = [ + b'B', + b'B', + b'W', + b'K', + 1, + 0, + WorkerFrameKind::Hello as u8, + 0x23, + 0, + 0, + 0, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x11, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 0x22, + 1, + 0, + 0, + ]; + assert_eq!(message.encode().unwrap(), expected); } #[test] -fn manifests_require_sorted_paths_and_valid_metadata() { - assert!(validate_upload_manifest(&[file_entry("z"), file_entry("a")]).is_err()); - assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, None).is_err()); - assert!(WorkerUploadEntry::new("dir", WorkerEntryKind::Directory, 0o755, 1, None).is_err()); - assert!(WorkerUploadEntry::new("dir", WorkerEntryKind::Directory, 0o755, 0, Some([1; 32])).is_err()); - assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o100000, 1, Some([1; 32])).is_err()); - assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, Some([1; 32])).is_ok()); -} - -#[test] -fn duplicate_environment_keys_and_bad_argv_are_rejected() { - assert!(WorkerBuild::new( - "make", - "/usr/bin/make", - vec!["x".into()], - "", - vec![("CC".into(), "one".into()), ("CC".into(), "two".into())], - Vec::new(), - ids().2, - ) - .is_err()); - - assert!(WorkerBuild::new("make", "/usr/bin/make", vec!["x\0y".into()], "", Vec::new(), Vec::new(), ids().2,).is_err()); - - assert!(WorkerBuild::new("make", "/usr/bin/make", Vec::new(), "", vec![("BAD-NAME".into(), "value".into())], Vec::new(), ids().2,).is_err()); - - assert!(WorkerBuild::new("make", "/usr/bin/make", vec!["arg".into(); MAX_WORKER_ARG_COUNT + 1], "", Vec::new(), Vec::new(), ids().2,).is_err()); - - assert!(WorkerBuild::new( - "make", - "/usr/bin/make", - Vec::new(), - "", - vec![("A".into(), "value".into()); MAX_WORKER_ENV_COUNT + 1], - Vec::new(), - ids().2, - ) - .is_err()); -} - -#[test] -fn bounded_chunks_output_errors_and_progress_are_rejected() { - let (request_id, session_id, upload_id) = ids(); - let path = WorkerRelativePath::new("file").unwrap(); - assert!(WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path: path.clone(), offset: 0, data: Vec::new() }.encode().is_err()); - assert!(WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset: 0, data: vec![0; MAX_WORKER_CHUNK_BYTES + 1] } - .encode() - .is_err()); - assert!(WorkerMessage::Stdout { request_id, session_id, data: vec![0; MAX_WORKER_OUTPUT_BYTES + 1] }.encode().is_err()); - assert!(WorkerMessage::Error { - request_id, - session_id, - operation: WorkerOperation::Build, - kind: WorkerErrorKind::Build, - message: "x".repeat(MAX_WORKER_ERROR_BYTES + 1), - } - .encode() - .is_err()); - assert!(WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 2, total_bytes: Some(1) }.encode().is_err()); -} - -#[test] -fn stdout_stderr_and_exact_completion_preserve_bytes_and_exit_code() { +fn host_reexport_decodes_ticket_11_v1_bytes() { let (request_id, session_id, _) = ids(); - let stdout = WorkerMessage::Stdout { request_id, session_id, data: b"out\0with bytes".to_vec() }; - let stderr = WorkerMessage::Stderr { request_id, session_id, data: b"err".to_vec() }; - let completed = WorkerMessage::Completed { request_id, session_id, operation: WorkerOperation::Build, exit_code: i32::MIN }; - assert_eq!(WorkerMessage::decode(&stdout.encode().unwrap()).unwrap(), stdout); - assert_eq!(WorkerMessage::decode(&stderr.encode().unwrap()).unwrap(), stderr); - assert_eq!(WorkerMessage::decode(&completed.encode().unwrap()).unwrap(), completed); -} - -#[test] -fn build_has_no_shell_packing_and_keeps_literal_arguments() { - let message = WorkerMessage::Build { request_id: ids().0, session_id: ids().1, build: build() }; - let decoded = WorkerMessage::decode(&message.encode().unwrap()).unwrap(); - let WorkerMessage::Build { build, .. } = decoded else { panic!("expected build") }; - assert_eq!(build.argv, vec!["release mode", "$(literal); still one arg"]); - assert_eq!(build.argv.len(), 2); - assert_eq!(build.tool.as_str(), "make"); - assert_eq!(build.trusted_executable.as_str(), "/usr/bin/make"); -} - -#[test] -fn error_kinds_preserve_protocol_build_and_cleanup_failures() { - let (request_id, session_id, _) = ids(); - for (operation, kind) in [ - (WorkerOperation::Protocol, WorkerErrorKind::WorkerProtocol), - (WorkerOperation::Build, WorkerErrorKind::Build), - (WorkerOperation::Cleanup, WorkerErrorKind::Cleanup), - ] { - let message = WorkerMessage::Error { request_id, session_id, operation, kind, message: "failure".into() }; - assert_eq!(WorkerMessage::decode(&message.encode().unwrap()).unwrap(), message); - } -} - -#[test] -fn unknown_nested_kinds_and_flags_are_rejected() { - let message = WorkerMessage::UploadBegin { request_id: ids().0, session_id: ids().1, upload_id: ids().2, entries: Vec::new() }; - let mut frame = message.encode().unwrap(); - frame[WORKER_FRAME_HEADER_LEN + 48..WORKER_FRAME_HEADER_LEN + 52].copy_from_slice(&u32::MAX.to_le_bytes()); - assert!(WorkerMessage::decode(&frame).is_err()); - - let hello = WorkerMessage::hello(ids().0, ids().1, false); - let mut frame = hello.encode().unwrap(); - frame[WORKER_FRAME_HEADER_LEN + 34] = 9; - assert!(WorkerMessage::decode(&frame).is_err()); + let message = WorkerMessage::hello(request_id, session_id, false); + let frame = message.encode().unwrap(); + assert_eq!(WorkerMessage::decode(&frame).unwrap(), message); } From 8848005bb952bab3f8ea27c500fe6ec92242078b Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 23:17:03 +0200 Subject: [PATCH 31/52] Implement artefact manifests and controlled retrieval --- crates/bunkerbox-worker-protocol/src/lib.rs | 514 +++++++++++++++-- crates/bunkerbox-worker/src/process.rs | 2 +- crates/bunkerbox-worker/src/storage.rs | 208 ++++++- crates/bunkerbox-worker/src/worker.rs | 212 +++++-- src/artifact.rs | 603 ++++++++++++++++++++ src/daemon.rs | 22 +- src/lib.rs | 1 + src/loopback.rs | 76 ++- src/main.rs | 5 +- src/remote.rs | 4 + src/remote_target.rs | 74 ++- src/ssh.rs | 231 +++++++- 12 files changed, 1847 insertions(+), 105 deletions(-) create mode 100644 src/artifact.rs diff --git a/crates/bunkerbox-worker-protocol/src/lib.rs b/crates/bunkerbox-worker-protocol/src/lib.rs index 39aa017..89c1b5a 100644 --- a/crates/bunkerbox-worker-protocol/src/lib.rs +++ b/crates/bunkerbox-worker-protocol/src/lib.rs @@ -21,6 +21,7 @@ pub const WORKER_PROTOCOL_MAGIC: [u8; 4] = *b"BBWK"; pub const WORKER_MAGIC: [u8; 4] = WORKER_PROTOCOL_MAGIC; pub const WORKER_PROTOCOL_VERSION: u16 = 1; pub const WORKER_VERSION: u16 = WORKER_PROTOCOL_VERSION; +pub const WORKER_ARTIFACT_PROTOCOL_VERSION: u16 = 2; pub const WORKER_FRAME_HEADER_LEN: usize = 4 + 2 + 1 + 4; pub const WORKER_FRAME_HEADER_SIZE: usize = WORKER_FRAME_HEADER_LEN; pub const WORKER_ID_LEN: usize = 16; @@ -59,6 +60,10 @@ pub const MAX_WORKER_MANIFEST_BYTES: usize = 512 * 1024; pub const MAX_WORKER_CHUNK_BYTES: usize = 64 * 1024; pub const MAX_WORKER_OUTPUT_BYTES: usize = 64 * 1024; pub const MAX_WORKER_ERROR_BYTES: usize = 4 * 1024; +pub const MAX_WORKER_ARTIFACT_ENTRIES: usize = 256; +pub const MAX_WORKER_ARTIFACT_PATH_BYTES: usize = MAX_WORKER_PATH_BYTES; +pub const MAX_WORKER_ARTIFACT_FILE_BYTES: u64 = 256 * 1024 * 1024; +pub const MAX_WORKER_ARTIFACT_TOTAL_BYTES: u64 = 512 * 1024 * 1024; #[derive(Debug, Clone, PartialEq, Eq)] pub enum WorkerProtocolError { @@ -127,6 +132,7 @@ macro_rules! worker_id { worker_id!(WorkerRequestId); worker_id!(WorkerSessionId); worker_id!(WorkerUploadId); +worker_id!(WorkerArtifactSetId); pub type RequestId = WorkerRequestId; pub type SessionId = WorkerSessionId; @@ -150,6 +156,10 @@ pub enum WorkerFrameKind { Stderr = 10, Completed = 11, Error = 12, + ArtifactManifest = 13, + FetchArtifact = 14, + ArtifactChunk = 15, + ArtifactComplete = 16, } impl WorkerFrameKind { @@ -171,6 +181,10 @@ impl WorkerFrameKind { 10 => Some(Self::Stderr), 11 => Some(Self::Completed), 12 => Some(Self::Error), + 13 => Some(Self::ArtifactManifest), + 14 => Some(Self::FetchArtifact), + 15 => Some(Self::ArtifactChunk), + 16 => Some(Self::ArtifactComplete), _ => None, } } @@ -192,6 +206,7 @@ pub enum WorkerOperation { Build = 2, Cleanup = 3, Sync = 4, + Artifact = 5, } impl WorkerOperation { @@ -206,6 +221,7 @@ impl WorkerOperation { 2 => Ok(Self::Build), 3 => Ok(Self::Cleanup), 4 => Ok(Self::Sync), + 5 => Ok(Self::Artifact), _ => Err(invalid(format!("unknown worker operation: {value}"))), } } @@ -219,6 +235,7 @@ pub enum WorkerErrorKind { Build = 3, Cleanup = 4, Sync = 5, + Artifact = 6, } pub type WorkerErrorClass = WorkerErrorKind; @@ -238,6 +255,7 @@ impl WorkerErrorKind { 3 => Ok(Self::Build), 4 => Ok(Self::Cleanup), 5 => Ok(Self::Sync), + 6 => Ok(Self::Artifact), _ => Err(invalid(format!("unknown worker error kind: {value}"))), } } @@ -340,6 +358,76 @@ impl TryFrom for WorkerExecutablePath { } } +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct WorkerArtifactPath(String); + +impl WorkerArtifactPath { + pub fn new(value: impl Into) -> WorkerResult { + let value = value.into(); + validate_relative_path("worker artifact path", &value, false)?; + Ok(Self(value)) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl AsRef for WorkerArtifactPath { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl TryFrom for WorkerArtifactPath { + type Error = WorkerProtocolError; + + fn try_from(value: String) -> WorkerResult { + Self::new(value) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct WorkerArtifactEntry { + pub path: WorkerArtifactPath, + pub mode: u32, + pub size: u64, + pub digest: WorkerDigest, +} + +impl WorkerArtifactEntry { + pub fn new(path: impl Into, mode: u32, size: u64, digest: WorkerDigest) -> WorkerResult { + let entry = Self { path: WorkerArtifactPath::new(path)?, mode, size, digest }; + entry.validate() + } + + pub fn validate(&self) -> WorkerResult { + if self.mode & !0o777 != 0 { + return Err(invalid(format!("worker artifact mode has unsupported bits: {:o}", self.mode))); + } + if self.size > MAX_WORKER_ARTIFACT_FILE_BYTES { + return Err(invalid(format!("worker artifact exceeds maximum file size {MAX_WORKER_ARTIFACT_FILE_BYTES}"))); + } + Ok(self.clone()) + } + + pub fn path(&self) -> &WorkerArtifactPath { + &self.path + } + + pub fn mode(&self) -> u32 { + self.mode + } + + pub fn size(&self) -> u64 { + self.size + } + + pub fn digest(&self) -> &WorkerDigest { + &self.digest + } +} + #[repr(u8)] #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum WorkerEntryKind { @@ -442,6 +530,9 @@ pub struct WorkerBuild { pub guest_env: Vec<(String, String)>, pub target_env: Vec<(String, String)>, pub upload_token: WorkerUploadId, + pub artifact_paths: Vec, + pub artifact_max_file_bytes: u64, + pub artifact_max_total_bytes: u64, } impl WorkerBuild { @@ -457,6 +548,9 @@ impl WorkerBuild { guest_env, target_env, upload_token, + artifact_paths: Vec::new(), + artifact_max_file_bytes: MAX_WORKER_ARTIFACT_FILE_BYTES, + artifact_max_total_bytes: MAX_WORKER_ARTIFACT_TOTAL_BYTES, }; build.validate()?; Ok(build) @@ -469,6 +563,17 @@ impl WorkerBuild { validate_argv(&self.argv)?; validate_environment("worker guest environment", &self.guest_env)?; validate_environment("worker target environment", &self.target_env)?; + validate_count(self.artifact_paths.len(), MAX_WORKER_ARTIFACT_ENTRIES, "worker artifact paths")?; + for path in &self.artifact_paths { + validate_relative_path("worker artifact path", path.as_str(), false)?; + } + validate_artifact_limits(self.artifact_max_file_bytes, self.artifact_max_total_bytes)?; + let mut paths = BTreeSet::new(); + for path in &self.artifact_paths { + if !paths.insert(path.as_str()) { + return Err(invalid(format!("duplicate worker artifact path: {}", path.as_str()))); + } + } Ok(()) } @@ -503,6 +608,26 @@ impl WorkerBuild { pub fn upload_token(&self) -> WorkerUploadId { self.upload_token } + + pub fn artifact_paths(&self) -> &[WorkerArtifactPath] { + &self.artifact_paths + } + + pub fn artifact_max_file_bytes(&self) -> u64 { + self.artifact_max_file_bytes + } + + pub fn artifact_max_total_bytes(&self) -> u64 { + self.artifact_max_total_bytes + } + + pub fn with_artifacts(mut self, paths: Vec, max_file_bytes: u64, max_total_bytes: u64) -> WorkerResult { + self.artifact_paths = paths; + self.artifact_max_file_bytes = max_file_bytes; + self.artifact_max_total_bytes = max_total_bytes; + self.validate()?; + Ok(self) + } } #[derive(Debug, Clone, PartialEq, Eq)] @@ -579,11 +704,42 @@ pub enum WorkerMessage { kind: WorkerErrorKind, message: String, }, + ArtifactManifest { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + artifact_set_id: WorkerArtifactSetId, + entries: Vec, + total_bytes: u64, + }, + FetchArtifact { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + artifact_set_id: WorkerArtifactSetId, + entry_index: u32, + }, + ArtifactChunk { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + artifact_set_id: WorkerArtifactSetId, + entry_index: u32, + offset: u64, + data: Vec, + }, + ArtifactComplete { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + artifact_set_id: WorkerArtifactSetId, + entry_index: u32, + }, } impl WorkerMessage { pub fn hello(request_id: WorkerRequestId, session_id: WorkerSessionId, response: bool) -> Self { - Self::Hello { request_id, session_id, version: WORKER_PROTOCOL_VERSION, response } + Self::hello_for_version(request_id, session_id, response, WORKER_PROTOCOL_VERSION) + } + + pub fn hello_for_version(request_id: WorkerRequestId, session_id: WorkerSessionId, response: bool, version: u16) -> Self { + Self::Hello { request_id, session_id, version, response } } pub fn build(request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild) -> Self { @@ -622,6 +778,10 @@ impl WorkerMessage { Self::Stderr { .. } => WorkerFrameKind::Stderr, Self::Completed { .. } => WorkerFrameKind::Completed, Self::Error { .. } => WorkerFrameKind::Error, + Self::ArtifactManifest { .. } => WorkerFrameKind::ArtifactManifest, + Self::FetchArtifact { .. } => WorkerFrameKind::FetchArtifact, + Self::ArtifactChunk { .. } => WorkerFrameKind::ArtifactChunk, + Self::ArtifactComplete { .. } => WorkerFrameKind::ArtifactComplete, } } @@ -638,7 +798,11 @@ impl WorkerMessage { | Self::Stdout { request_id, .. } | Self::Stderr { request_id, .. } | Self::Completed { request_id, .. } - | Self::Error { request_id, .. } => *request_id, + | Self::Error { request_id, .. } + | Self::ArtifactManifest { request_id, .. } + | Self::FetchArtifact { request_id, .. } + | Self::ArtifactChunk { request_id, .. } + | Self::ArtifactComplete { request_id, .. } => *request_id, } } @@ -655,7 +819,11 @@ impl WorkerMessage { | Self::Stdout { session_id, .. } | Self::Stderr { session_id, .. } | Self::Completed { session_id, .. } - | Self::Error { session_id, .. } => *session_id, + | Self::Error { session_id, .. } + | Self::ArtifactManifest { session_id, .. } + | Self::FetchArtifact { session_id, .. } + | Self::ArtifactChunk { session_id, .. } + | Self::ArtifactComplete { session_id, .. } => *session_id, } } @@ -668,14 +836,22 @@ impl WorkerMessage { | Self::SyncProgress { upload_id, .. } => Some(*upload_id), Self::Build { build, .. } => Some(build.upload_token), Self::Cleanup { upload_token, .. } => Some(*upload_token), - Self::Hello { .. } | Self::Stdout { .. } | Self::Stderr { .. } | Self::Completed { .. } | Self::Error { .. } => None, + Self::Hello { .. } + | Self::Stdout { .. } + | Self::Stderr { .. } + | Self::Completed { .. } + | Self::Error { .. } + | Self::ArtifactManifest { .. } + | Self::FetchArtifact { .. } + | Self::ArtifactChunk { .. } + | Self::ArtifactComplete { .. } => None, } } pub fn validate(&self) -> WorkerResult<()> { match self { Self::Hello { version, .. } => { - if *version != WORKER_PROTOCOL_VERSION { + if !is_supported_version(*version) { return Err(invalid(format!("unsupported worker hello version: {version}"))); } } @@ -702,20 +878,53 @@ impl WorkerMessage { } } Self::Error { message, .. } => validate_error_message(message)?, + Self::ArtifactManifest { artifact_set_id, entries, total_bytes, .. } => { + validate_nonzero_id(artifact_set_id.0, "worker artifact set")?; + validate_artifact_manifest(entries, *total_bytes)?; + } + Self::FetchArtifact { artifact_set_id, entry_index, .. } => { + validate_nonzero_id(artifact_set_id.0, "worker artifact set")?; + validate_artifact_index(*entry_index)?; + } + Self::ArtifactChunk { artifact_set_id, entry_index, offset, data, .. } => { + validate_nonzero_id(artifact_set_id.0, "worker artifact set")?; + validate_artifact_index(*entry_index)?; + validate_chunk_offset(*offset, data)?; + } + Self::ArtifactComplete { artifact_set_id, entry_index, .. } => { + validate_nonzero_id(artifact_set_id.0, "worker artifact set")?; + validate_artifact_index(*entry_index)?; + } } Ok(()) } + fn requires_artifact_version(&self) -> bool { + match self { + Self::Build { build, .. } => !build.artifact_paths.is_empty(), + Self::ArtifactManifest { .. } | Self::FetchArtifact { .. } | Self::ArtifactChunk { .. } | Self::ArtifactComplete { .. } => true, + _ => false, + } + } + pub fn encode(&self) -> WorkerResult> { + self.encode_version(WORKER_PROTOCOL_VERSION) + } + + pub fn encode_version(&self, version: u16) -> WorkerResult> { + validate_version(version)?; + if version == WORKER_PROTOCOL_VERSION && self.requires_artifact_version() { + return Err(invalid("worker artifact message requires artifact-capable protocol version")); + } self.validate()?; let mut payload = WireWriter::new(); - encode_payload(self, &mut payload)?; + encode_payload(self, &mut payload, version)?; let payload = payload.finish()?; let declared = u32::try_from(payload.len()).map_err(|_| invalid("worker payload length does not fit in u32"))?; let mut frame = Vec::with_capacity(WORKER_FRAME_HEADER_LEN + payload.len()); frame.extend_from_slice(&WORKER_PROTOCOL_MAGIC); - frame.extend_from_slice(&WORKER_PROTOCOL_VERSION.to_le_bytes()); + frame.extend_from_slice(&version.to_le_bytes()); frame.push(self.kind().as_u8()); frame.extend_from_slice(&declared.to_le_bytes()); frame.extend_from_slice(&payload); @@ -723,42 +932,74 @@ impl WorkerMessage { } pub fn decode(frame: &[u8]) -> WorkerResult { - let (kind, payload) = split_frame(frame)?; - decode_payload(kind, payload) + let (version, message) = Self::decode_versioned(frame)?; + if version != WORKER_PROTOCOL_VERSION { + return Err(invalid(format!("unsupported worker protocol version: {version}"))); + } + Ok(message) + } + + pub fn decode_versioned(frame: &[u8]) -> WorkerResult<(u16, Self)> { + let (version, kind, payload) = split_frame_versioned(frame)?; + let message = decode_payload(kind, payload, version)?; + Ok((version, message)) } #[cfg(feature = "async")] pub async fn read_async(reader: &mut R) -> WorkerResult { + let (version, message) = Self::read_async_versioned(reader).await?; + require_default_version(version)?; + Ok(message) + } + + #[cfg(feature = "async")] + pub async fn read_async_versioned(reader: &mut R) -> WorkerResult<(u16, Self)> { let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; reader.read_exact(&mut header).await.map_err(|error| match error.kind() { io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), _ => WorkerProtocolError::Io(error.to_string()), })?; - let (kind, payload_len) = decode_header(&header)?; + let (version, kind, payload_len) = decode_header(&header)?; let mut payload = vec![0u8; payload_len]; reader.read_exact(&mut payload).await.map_err(|error| match error.kind() { io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), _ => WorkerProtocolError::Io(error.to_string()), })?; - decode_payload(kind, &payload) + Ok((version, decode_payload(kind, &payload, version)?)) } pub fn read_blocking(reader: &mut R) -> WorkerResult { + let (version, message) = Self::read_blocking_versioned(reader)?; + require_default_version(version)?; + Ok(message) + } + + pub fn read_blocking_versioned(reader: &mut R) -> WorkerResult<(u16, Self)> { let mut header = [0u8; WORKER_FRAME_HEADER_LEN]; reader.read_exact(&mut header).map_err(|error| match error.kind() { io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), _ => WorkerProtocolError::Io(error.to_string()), })?; - let (kind, payload_len) = decode_header(&header)?; + let (version, kind, payload_len) = decode_header(&header)?; let mut payload = vec![0u8; payload_len]; reader.read_exact(&mut payload).map_err(|error| match error.kind() { io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), _ => WorkerProtocolError::Io(error.to_string()), })?; - decode_payload(kind, &payload) + Ok((version, decode_payload(kind, &payload, version)?)) } pub fn read_blocking_optional(reader: &mut R) -> WorkerResult> { + let message = Self::read_blocking_optional_versioned(reader)?; + message + .map(|(version, message)| { + require_default_version(version)?; + Ok(message) + }) + .transpose() + } + + pub fn read_blocking_optional_versioned(reader: &mut R) -> WorkerResult> { let mut first = [0u8; 1]; match reader.read_exact(&mut first) { Ok(()) => {} @@ -771,24 +1012,33 @@ impl WorkerMessage { io::ErrorKind::UnexpectedEof => invalid("truncated worker frame header"), _ => WorkerProtocolError::Io(error.to_string()), })?; - let (kind, payload_len) = decode_header(&header)?; + let (version, kind, payload_len) = decode_header(&header)?; let mut payload = vec![0u8; payload_len]; reader.read_exact(&mut payload).map_err(|error| match error.kind() { io::ErrorKind::UnexpectedEof => invalid("truncated worker frame payload"), _ => WorkerProtocolError::Io(error.to_string()), })?; - decode_payload(kind, &payload).map(Some) + decode_payload(kind, &payload, version).map(|message| Some((version, message))) } #[cfg(feature = "async")] pub async fn write_async(&self, writer: &mut W) -> WorkerResult<()> { - let frame = self.encode()?; + self.write_async_version(writer, WORKER_PROTOCOL_VERSION).await + } + + #[cfg(feature = "async")] + pub async fn write_async_version(&self, writer: &mut W, version: u16) -> WorkerResult<()> { + let frame = self.encode_version(version)?; writer.write_all(&frame).await.map_err(WorkerProtocolError::from)?; writer.flush().await.map_err(WorkerProtocolError::from) } pub fn write_blocking(&self, writer: &mut W) -> WorkerResult<()> { - let frame = self.encode()?; + self.write_blocking_version(writer, WORKER_PROTOCOL_VERSION) + } + + pub fn write_blocking_version(&self, writer: &mut W, version: u16) -> WorkerResult<()> { + let frame = self.encode_version(version)?; writer.write_all(&frame).map_err(WorkerProtocolError::from)?; writer.flush().map_err(WorkerProtocolError::from) } @@ -809,11 +1059,21 @@ pub async fn read_message(reader: &mut R) -> WorkerResult< read_worker_message(reader).await } +#[cfg(feature = "async")] +pub async fn read_message_versioned(reader: &mut R) -> WorkerResult<(u16, WorkerMessage)> { + WorkerMessage::read_async_versioned(reader).await +} + #[cfg(feature = "async")] pub async fn write_message(writer: &mut W, message: &WorkerMessage) -> WorkerResult<()> { write_worker_message(writer, message).await } +#[cfg(feature = "async")] +pub async fn write_message_versioned(writer: &mut W, message: &WorkerMessage, version: u16) -> WorkerResult<()> { + message.write_async_version(writer, version).await +} + pub fn encode_worker_message(message: &WorkerMessage) -> WorkerResult> { message.encode() } @@ -864,7 +1124,7 @@ pub fn validate_upload_manifest(entries: &[WorkerUploadEntry]) -> WorkerResult WorkerResult<()> { +fn encode_payload(message: &WorkerMessage, writer: &mut WireWriter, version: u16) -> WorkerResult<()> { match message { WorkerMessage::Hello { request_id, session_id, version, response } => { encode_correlation(writer, *request_id, *session_id)?; @@ -898,7 +1158,7 @@ fn encode_payload(message: &WorkerMessage, writer: &mut WireWriter) -> WorkerRes } WorkerMessage::Build { request_id, session_id, build } => { encode_correlation(writer, *request_id, *session_id)?; - encode_build(writer, build)?; + encode_build(writer, build, version)?; } WorkerMessage::Cleanup { request_id, session_id, upload_token } => { encode_correlation(writer, *request_id, *session_id)?; @@ -931,11 +1191,37 @@ fn encode_payload(message: &WorkerMessage, writer: &mut WireWriter) -> WorkerRes writer.u8(kind.as_u8())?; writer.string(message, MAX_WORKER_ERROR_BYTES, "worker error")?; } + WorkerMessage::ArtifactManifest { request_id, session_id, artifact_set_id, entries, total_bytes } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(artifact_set_id.0)?; + writer.count(entries.len(), MAX_WORKER_ARTIFACT_ENTRIES, "worker artifact entries")?; + for entry in entries { + encode_artifact_entry(writer, entry)?; + } + writer.u64(*total_bytes)?; + } + WorkerMessage::FetchArtifact { request_id, session_id, artifact_set_id, entry_index } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(artifact_set_id.0)?; + writer.u32(*entry_index)?; + } + WorkerMessage::ArtifactChunk { request_id, session_id, artifact_set_id, entry_index, offset, data } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(artifact_set_id.0)?; + writer.u32(*entry_index)?; + writer.u64(*offset)?; + writer.blob(data, MAX_WORKER_CHUNK_BYTES, "worker artifact chunk")?; + } + WorkerMessage::ArtifactComplete { request_id, session_id, artifact_set_id, entry_index } => { + encode_correlation(writer, *request_id, *session_id)?; + writer.id(artifact_set_id.0)?; + writer.u32(*entry_index)?; + } } Ok(()) } -fn decode_payload(kind: WorkerFrameKind, payload: &[u8]) -> WorkerResult { +fn decode_payload(kind: WorkerFrameKind, payload: &[u8], version: u16) -> WorkerResult { let mut reader = WireReader::new(payload); let message = match kind { WorkerFrameKind::Hello => { @@ -976,7 +1262,7 @@ fn decode_payload(kind: WorkerFrameKind, payload: &[u8]) -> WorkerResult { let (request_id, session_id) = decode_correlation(&mut reader)?; - WorkerMessage::Build { request_id, session_id, build: decode_build(&mut reader)? } + WorkerMessage::Build { request_id, session_id, build: decode_build(&mut reader, version)? } } WorkerFrameKind::Cleanup => { let (request_id, session_id) = decode_correlation(&mut reader)?; @@ -1011,13 +1297,65 @@ fn decode_payload(kind: WorkerFrameKind, payload: &[u8]) -> WorkerResult { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let artifact_set_id = WorkerArtifactSetId(reader.array16()?); + let count = reader.count(MAX_WORKER_ARTIFACT_ENTRIES, "worker artifact entries")?; + let mut entries = Vec::with_capacity(count); + for _ in 0..count { + entries.push(decode_artifact_entry(&mut reader)?); + } + let total_bytes = reader.u64()?; + WorkerMessage::ArtifactManifest { request_id, session_id, artifact_set_id, entries, total_bytes } + } + WorkerFrameKind::FetchArtifact => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let artifact_set_id = WorkerArtifactSetId(reader.array16()?); + let entry_index = reader.u32()?; + WorkerMessage::FetchArtifact { request_id, session_id, artifact_set_id, entry_index } + } + WorkerFrameKind::ArtifactChunk => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let artifact_set_id = WorkerArtifactSetId(reader.array16()?); + let entry_index = reader.u32()?; + let offset = reader.u64()?; + let data = reader.blob(MAX_WORKER_CHUNK_BYTES, "worker artifact chunk")?; + WorkerMessage::ArtifactChunk { request_id, session_id, artifact_set_id, entry_index, offset, data } + } + WorkerFrameKind::ArtifactComplete => { + let (request_id, session_id) = decode_correlation(&mut reader)?; + let artifact_set_id = WorkerArtifactSetId(reader.array16()?); + let entry_index = reader.u32()?; + WorkerMessage::ArtifactComplete { request_id, session_id, artifact_set_id, entry_index } + } }; reader.finish()?; message.validate()?; + if version == WORKER_PROTOCOL_VERSION && message.requires_artifact_version() { + return Err(invalid("worker artifact message requires artifact-capable protocol version")); + } + if let WorkerMessage::Hello { version: hello_version, .. } = &message { + if *hello_version != version { + return Err(invalid("worker Hello version does not match frame version")); + } + } Ok(message) } -fn encode_build(writer: &mut WireWriter, build: &WorkerBuild) -> WorkerResult<()> { +fn encode_artifact_entry(writer: &mut WireWriter, entry: &WorkerArtifactEntry) -> WorkerResult<()> { + entry.validate()?; + writer.string(entry.path.as_str(), MAX_WORKER_ARTIFACT_PATH_BYTES, "worker artifact path")?; + writer.u32(entry.mode)?; + writer.u64(entry.size)?; + writer.bytes(&entry.digest)?; + Ok(()) +} + +fn decode_artifact_entry(reader: &mut WireReader<'_>) -> WorkerResult { + WorkerArtifactEntry::new(reader.string(MAX_WORKER_ARTIFACT_PATH_BYTES, "worker artifact path")?, reader.u32()?, reader.u64()?, reader.array32()?) +} + +fn encode_build(writer: &mut WireWriter, build: &WorkerBuild, version: u16) -> WorkerResult<()> { build.validate()?; writer.string(build.tool.as_str(), MAX_WORKER_TOOL_BYTES, "worker tool")?; writer.string(build.trusted_executable.as_str(), MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?; @@ -1029,10 +1367,18 @@ fn encode_build(writer: &mut WireWriter, build: &WorkerBuild) -> WorkerResult<() encode_environment(writer, &build.guest_env, "worker guest environment")?; encode_environment(writer, &build.target_env, "worker target environment")?; writer.id(build.upload_token.0)?; + if version >= WORKER_ARTIFACT_PROTOCOL_VERSION { + writer.count(build.artifact_paths.len(), MAX_WORKER_ARTIFACT_ENTRIES, "worker artifact paths")?; + for path in &build.artifact_paths { + writer.string(path.as_str(), MAX_WORKER_ARTIFACT_PATH_BYTES, "worker artifact path")?; + } + writer.u64(build.artifact_max_file_bytes)?; + writer.u64(build.artifact_max_total_bytes)?; + } Ok(()) } -fn decode_build(reader: &mut WireReader<'_>) -> WorkerResult { +fn decode_build(reader: &mut WireReader<'_>, version: u16) -> WorkerResult { let tool = WorkerTool::new(reader.string(MAX_WORKER_TOOL_BYTES, "worker tool")?)?; let trusted_executable = WorkerExecutablePath::new(reader.string(MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?)?; let argument_count = reader.count(MAX_WORKER_ARG_COUNT, "worker argv")?; @@ -1044,7 +1390,28 @@ fn decode_build(reader: &mut WireReader<'_>) -> WorkerResult { let guest_env = decode_environment(reader, "worker guest environment")?; let target_env = decode_environment(reader, "worker target environment")?; let upload_token = WorkerUploadId(reader.array16()?); - let build = WorkerBuild { tool, trusted_executable, argv, cwd, guest_env, target_env, upload_token }; + let (artifact_paths, artifact_max_file_bytes, artifact_max_total_bytes) = if version >= WORKER_ARTIFACT_PROTOCOL_VERSION { + let count = reader.count(MAX_WORKER_ARTIFACT_ENTRIES, "worker artifact paths")?; + let mut paths = Vec::with_capacity(count); + for _ in 0..count { + paths.push(WorkerArtifactPath::new(reader.string(MAX_WORKER_ARTIFACT_PATH_BYTES, "worker artifact path")?)?); + } + (paths, reader.u64()?, reader.u64()?) + } else { + (Vec::new(), MAX_WORKER_ARTIFACT_FILE_BYTES, MAX_WORKER_ARTIFACT_TOTAL_BYTES) + }; + let build = WorkerBuild { + tool, + trusted_executable, + argv, + cwd, + guest_env, + target_env, + upload_token, + artifact_paths, + artifact_max_file_bytes, + artifact_max_total_bytes, + }; build.validate()?; Ok(build) } @@ -1292,11 +1659,96 @@ fn invalid(message: impl Into) -> WorkerProtocolError { WorkerProtocolError::Invalid(message.into()) } -fn split_frame(frame: &[u8]) -> WorkerResult<(WorkerFrameKind, &[u8])> { +fn is_supported_version(version: u16) -> bool { + matches!(version, WORKER_PROTOCOL_VERSION | WORKER_ARTIFACT_PROTOCOL_VERSION) +} + +fn validate_version(version: u16) -> WorkerResult<()> { + if is_supported_version(version) { + Ok(()) + } else { + Err(invalid(format!("unsupported worker protocol version: {version}"))) + } +} + +fn require_default_version(version: u16) -> WorkerResult<()> { + if version == WORKER_PROTOCOL_VERSION { + Ok(()) + } else { + Err(invalid(format!("unsupported worker protocol version: {version}"))) + } +} + +fn validate_nonzero_id(bytes: [u8; WORKER_ID_LEN], label: &str) -> WorkerResult<()> { + if bytes == [0; WORKER_ID_LEN] { + return Err(invalid(format!("{label} ID must be nonzero"))); + } + Ok(()) +} + +fn validate_artifact_limits(max_file_bytes: u64, max_total_bytes: u64) -> WorkerResult<()> { + if max_file_bytes == 0 || max_file_bytes > MAX_WORKER_ARTIFACT_FILE_BYTES { + return Err(invalid("worker artifact per-file limit is invalid")); + } + if max_total_bytes == 0 || max_total_bytes > MAX_WORKER_ARTIFACT_TOTAL_BYTES { + return Err(invalid("worker artifact total limit is invalid")); + } + Ok(()) +} + +fn validate_artifact_index(index: u32) -> WorkerResult<()> { + if usize::try_from(index).map_or(true, |index| index >= MAX_WORKER_ARTIFACT_ENTRIES) { + return Err(invalid(format!("worker artifact index exceeds maximum {MAX_WORKER_ARTIFACT_ENTRIES}"))); + } + Ok(()) +} + +fn validate_chunk_offset(offset: u64, data: &[u8]) -> WorkerResult<()> { + if data.is_empty() { + return Err(invalid("worker artifact chunk must not be empty")); + } + if data.len() > MAX_WORKER_CHUNK_BYTES { + return Err(invalid(format!("worker artifact chunk exceeds maximum length {MAX_WORKER_CHUNK_BYTES}"))); + } + let end = offset.checked_add(data.len() as u64).ok_or_else(|| invalid("worker artifact chunk offset overflow"))?; + if end > MAX_WORKER_ARTIFACT_FILE_BYTES { + return Err(invalid("worker artifact chunk exceeds maximum file size")); + } + Ok(()) +} + +fn validate_artifact_manifest(entries: &[WorkerArtifactEntry], total_bytes: u64) -> WorkerResult<()> { + validate_count(entries.len(), MAX_WORKER_ARTIFACT_ENTRIES, "worker artifact entries")?; + let mut paths = BTreeSet::new(); + let mut total = 0u64; + let mut manifest_bytes = 4usize; + for entry in entries { + entry.validate()?; + if !paths.insert(entry.path.as_str()) { + return Err(invalid(format!("duplicate worker artifact path: {}", entry.path.as_str()))); + } + total = total.checked_add(entry.size).ok_or_else(|| invalid("worker artifact total size overflow"))?; + if total > MAX_WORKER_ARTIFACT_TOTAL_BYTES || total > total_bytes { + return Err(invalid("worker artifact total exceeds its limit")); + } + manifest_bytes = manifest_bytes + .checked_add(4 + entry.path.as_str().len() + 4 + 8 + WORKER_DIGEST_LEN) + .ok_or_else(|| invalid("worker artifact manifest length overflow"))?; + if manifest_bytes > MAX_WORKER_MANIFEST_BYTES { + return Err(invalid("worker artifact manifest exceeds maximum length")); + } + } + if total != total_bytes { + return Err(invalid("worker artifact manifest total does not match entries")); + } + Ok(()) +} + +fn split_frame_versioned(frame: &[u8]) -> WorkerResult<(u16, WorkerFrameKind, &[u8])> { if frame.len() < WORKER_FRAME_HEADER_LEN { return Err(invalid("truncated worker frame header")); } - let (kind, payload_len) = decode_header(&frame[..WORKER_FRAME_HEADER_LEN])?; + let (version, kind, payload_len) = decode_header(&frame[..WORKER_FRAME_HEADER_LEN])?; let expected = WORKER_FRAME_HEADER_LEN.checked_add(payload_len).ok_or_else(|| invalid("worker frame length overflow"))?; if frame.len() < expected { return Err(invalid("truncated worker frame payload")); @@ -1304,10 +1756,10 @@ fn split_frame(frame: &[u8]) -> WorkerResult<(WorkerFrameKind, &[u8])> { if frame.len() > expected { return Err(invalid("extra bytes after worker frame")); } - Ok((kind, &frame[WORKER_FRAME_HEADER_LEN..expected])) + Ok((version, kind, &frame[WORKER_FRAME_HEADER_LEN..expected])) } -fn decode_header(header: &[u8]) -> WorkerResult<(WorkerFrameKind, usize)> { +fn decode_header(header: &[u8]) -> WorkerResult<(u16, WorkerFrameKind, usize)> { if header.len() != WORKER_FRAME_HEADER_LEN { return Err(invalid("invalid worker frame header length")); } @@ -1315,15 +1767,13 @@ fn decode_header(header: &[u8]) -> WorkerResult<(WorkerFrameKind, usize)> { return Err(invalid("invalid worker frame magic")); } let version = u16::from_le_bytes([header[4], header[5]]); - if version != WORKER_PROTOCOL_VERSION { - return Err(invalid(format!("unsupported worker protocol version: {version}"))); - } + validate_version(version)?; let kind = WorkerFrameKind::try_from(header[6])?; let payload_len = u32::from_le_bytes([header[7], header[8], header[9], header[10]]) as usize; if payload_len > MAX_WORKER_FRAME_PAYLOAD { return Err(invalid(format!("worker payload exceeds maximum length {MAX_WORKER_FRAME_PAYLOAD}"))); } - Ok((kind, payload_len)) + Ok((version, kind, payload_len)) } struct WireWriter { diff --git a/crates/bunkerbox-worker/src/process.rs b/crates/bunkerbox-worker/src/process.rs index 3841c60..e08bd87 100644 --- a/crates/bunkerbox-worker/src/process.rs +++ b/crates/bunkerbox-worker/src/process.rs @@ -94,7 +94,7 @@ impl Drop for JobWorkspace { pub fn cleanup_stale_jobs(parent: &File) -> io::Result<()> { for lock_name in platform::list_names(parent)?.into_iter().take(MAX_STALE_JOBS) { let Some(name) = lock_name.strip_suffix(".lock") else { continue }; - if !name.starts_with("job-") { + if !name.starts_with("job-") && !name.starts_with("artifact-") { continue; } let Ok(lock) = platform::open_lock_at(parent, &lock_name) else { continue }; diff --git a/crates/bunkerbox-worker/src/storage.rs b/crates/bunkerbox-worker/src/storage.rs index e6f9c63..f6774bd 100644 --- a/crates/bunkerbox-worker/src/storage.rs +++ b/crates/bunkerbox-worker/src/storage.rs @@ -1,7 +1,8 @@ use crate::platform; use bunkerbox_worker_protocol::{ - validate_upload_manifest, WorkerDigest, WorkerEntryKind, WorkerProtocolError, WorkerRelativePath, WorkerSessionId, WorkerUploadEntry, - WorkerUploadId, MAX_WORKER_MANIFEST_BYTES, + validate_upload_manifest, WorkerArtifactEntry, WorkerArtifactPath, WorkerArtifactSetId, WorkerDigest, WorkerEntryKind, WorkerProtocolError, + WorkerRelativePath, WorkerSessionId, WorkerUploadEntry, WorkerUploadId, MAX_WORKER_ARTIFACT_FILE_BYTES, MAX_WORKER_ARTIFACT_TOTAL_BYTES, + MAX_WORKER_MANIFEST_BYTES, }; use sha2::{Digest, Sha256}; use std::collections::BTreeMap; @@ -292,6 +293,118 @@ impl StoredUpload { } } +pub struct ArtifactSpool { + parent: File, + name: String, + root: File, + lock: File, + files: File, + artifact_set_id: WorkerArtifactSetId, + entries: Vec, + total_bytes: u64, +} + +impl ArtifactSpool { + pub fn capture(parent: &File, job_root: &File, paths: &[WorkerArtifactPath], max_file_bytes: u64, max_total_bytes: u64) -> Result { + if max_file_bytes == 0 || max_file_bytes > MAX_WORKER_ARTIFACT_FILE_BYTES { + return Err("worker artifact per-file limit is invalid".to_string()); + } + if max_total_bytes == 0 || max_total_bytes > MAX_WORKER_ARTIFACT_TOTAL_BYTES { + return Err("worker artifact total limit is invalid".to_string()); + } + let (name, lock_name, lock, root) = reserve_artifact_root(parent)?; + let files = match private_directory(&root, FILES_DIRECTORY) { + Ok(files) => files, + Err(error) => { + let _ = platform::remove_tree_at(parent, &name); + let _ = platform::unlink_at(parent, &lock_name, 0); + return Err(error); + } + }; + + let result = (|| { + let mut entries = Vec::with_capacity(paths.len()); + let mut total_bytes = 0u64; + for path in paths { + let source = open_relative_file(job_root, path.as_str())?; + let before = platform::stat_fd(source.as_raw_fd()).map_err(|error| format!("stat worker artifact {}: {error}", path.as_str()))?; + validate_artifact_source(&before, path.as_str())?; + let size = u64::try_from(before.st_size).map_err(|_| format!("worker artifact size is invalid: {}", path.as_str()))?; + if size > max_file_bytes { + return Err(format!("worker artifact exceeds per-file limit: {}", path.as_str())); + } + total_bytes = total_bytes.checked_add(size).ok_or_else(|| "worker artifact total size overflow".to_string())?; + if total_bytes > max_total_bytes { + return Err("worker artifacts exceed total size limit".to_string()); + } + let mode = (before.st_mode as u32) & 0o777; + if let Some((parents, _)) = path.as_str().rsplit_once('/') { + ensure_directory(&files, parents, 0o700)?; + } + let destination = create_relative_file(&files, path.as_str(), mode)?; + let digest = copy_artifact(&source, &destination, size, path.as_str())?; + let after = platform::stat_fd(source.as_raw_fd()).map_err(|error| format!("restat worker artifact {}: {error}", path.as_str()))?; + if after.st_dev != before.st_dev + || after.st_ino != before.st_ino + || after.st_size != before.st_size + || after.st_mode & 0o777 != before.st_mode & 0o777 + || after.st_nlink != before.st_nlink + { + return Err(format!("worker artifact changed during capture: {}", path.as_str())); + } + entries.push(WorkerArtifactEntry::new(path.as_str().to_string(), mode, size, digest).map_err(protocol_error)?); + } + let artifact_set_id = artifact_set_id(&name, &entries, total_bytes); + Ok((artifact_set_id, entries, total_bytes)) + })(); + + match result { + Ok((artifact_set_id, entries, total_bytes)) => Ok(Self { + parent: parent.try_clone().map_err(|error| format!("clone worker jobs directory: {error}"))?, + name, + root, + lock, + files, + artifact_set_id, + entries, + total_bytes, + }), + Err(error) => { + let _ = platform::remove_tree_at(parent, &name); + let _ = platform::unlink_at(parent, &lock_name, 0); + Err(error) + } + } + } + + pub fn artifact_set_id(&self) -> WorkerArtifactSetId { + self.artifact_set_id + } + + pub fn entries(&self) -> &[WorkerArtifactEntry] { + &self.entries + } + + pub fn total_bytes(&self) -> u64 { + self.total_bytes + } + + pub fn open_entry(&self, index: usize) -> Result { + let entry = self.entries.get(index).ok_or_else(|| "worker artifact index is out of range".to_string())?; + let file = open_relative_file(&self.files, entry.path().as_str())?; + let metadata = platform::stat_fd(file.as_raw_fd()).map_err(|error| format!("stat worker artifact spool: {error}"))?; + validate_regular_file(&metadata, entry.size(), entry.mode(), entry.path().as_str())?; + Ok(file) + } +} + +impl Drop for ArtifactSpool { + fn drop(&mut self) { + let _ = (&self.root, &self.lock, &self.files); + let _ = platform::remove_tree_at(&self.parent, &self.name); + } +} + struct PendingFile { file: File, size: u64, @@ -441,6 +554,97 @@ fn private_directory(parent: &File, name: &str) -> Result { Ok(directory) } +fn reserve_artifact_root(parent: &File) -> Result<(String, String, File, File), String> { + for _ in 0..32 { + let name = format!("artifact-{}-{}", unsafe { libc::getpid() }, next_job_id()); + let lock_name = format!("{name}.lock"); + let lock = match platform::create_file_at(parent, &lock_name, 0o600) { + Ok(lock) => lock, + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(format!("create worker artifact lock: {error}")), + }; + if !platform::lock_exclusive(&lock).map_err(|error| format!("lock worker artifact spool: {error}"))? { + let _ = platform::unlink_at(parent, &lock_name, 0); + continue; + } + if let Err(error) = platform::create_dir_at(parent, &name, 0o700) { + let _ = platform::unlink_at(parent, &lock_name, 0); + if error.kind() == io::ErrorKind::AlreadyExists { + continue; + } + return Err(format!("create worker artifact spool: {error}")); + } + let root = match platform::open_dir_at(parent, &name) { + Ok(root) => root, + Err(error) => { + let _ = platform::remove_tree_at(parent, &name); + let _ = platform::unlink_at(parent, &lock_name, 0); + return Err(format!("open worker artifact spool: {error}")); + } + }; + if let Err(error) = platform::validate_private_directory(&root, "worker artifact spool") { + let _ = platform::remove_tree_at(parent, &name); + let _ = platform::unlink_at(parent, &lock_name, 0); + return Err(error); + } + return Ok((name, lock_name, lock, root)); + } + Err("could not reserve a worker artifact spool".to_string()) +} + +fn validate_artifact_source(metadata: &libc::stat, path: &str) -> Result<(), String> { + if metadata.st_mode & libc::S_IFMT != libc::S_IFREG || metadata.st_nlink != 1 { + return Err(format!("worker artifact is not a private regular file: {path}")); + } + if metadata.st_size < 0 { + return Err(format!("worker artifact has an invalid size: {path}")); + } + Ok(()) +} + +fn copy_artifact(source: &File, destination: &File, expected_size: u64, path: &str) -> Result { + let mut source = source.try_clone().map_err(|error| format!("clone worker artifact {path}: {error}"))?; + let mut destination = destination.try_clone().map_err(|error| format!("clone worker artifact spool {path}: {error}"))?; + let mut hasher = Sha256::new(); + let mut copied = 0u64; + let mut buffer = vec![0u8; COPY_BUFFER_BYTES]; + loop { + let count = source.read(&mut buffer).map_err(|error| format!("read worker artifact {path}: {error}"))?; + if count == 0 { + break; + } + copied = copied.checked_add(count as u64).ok_or_else(|| "worker artifact size overflow".to_string())?; + if copied > expected_size { + return Err(format!("worker artifact grew during capture: {path}")); + } + hasher.update(&buffer[..count]); + destination.write_all(&buffer[..count]).map_err(|error| format!("write worker artifact spool {path}: {error}"))?; + } + if copied != expected_size { + return Err(format!("worker artifact size changed during capture: {path}")); + } + platform::sync_fd(&destination).map_err(|error| format!("flush worker artifact spool {path}: {error}"))?; + Ok(hasher.finalize().into()) +} + +fn artifact_set_id(name: &str, entries: &[WorkerArtifactEntry], total_bytes: u64) -> WorkerArtifactSetId { + let mut hasher = Sha256::new(); + hasher.update(name.as_bytes()); + hasher.update(total_bytes.to_le_bytes()); + for entry in entries { + hasher.update(entry.path().as_str().as_bytes()); + hasher.update(entry.size().to_le_bytes()); + hasher.update(entry.digest()); + } + let digest = hasher.finalize(); + let mut id = [0u8; 16]; + id.copy_from_slice(&digest[..16]); + if id == [0; 16] { + id[0] = 1; + } + WorkerArtifactSetId(id) +} + fn open_existing_private_directory(parent: &File, name: &str, label: &str) -> Result { let directory = platform::open_dir_at(parent, name).map_err(|error| format!("open {label}: {error}"))?; platform::validate_private_directory(&directory, label)?; diff --git a/crates/bunkerbox-worker/src/worker.rs b/crates/bunkerbox-worker/src/worker.rs index cd0d6c2..495d132 100644 --- a/crates/bunkerbox-worker/src/worker.rs +++ b/crates/bunkerbox-worker/src/worker.rs @@ -1,25 +1,31 @@ use crate::process::{self, JobWorkspace, OutputSink}; -use crate::storage::{UploadStore, UploadTransaction}; +use crate::storage::{ArtifactSpool, UploadStore, UploadTransaction}; use bunkerbox_worker_protocol::{ - WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadId, MAX_WORKER_ERROR_BYTES, + WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadId, MAX_WORKER_CHUNK_BYTES, + MAX_WORKER_ERROR_BYTES, WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, }; use std::io::{self, Read, Write}; -use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU16, AtomicUsize, Ordering}; use std::sync::mpsc::{self, Receiver}; use std::sync::{Arc, Mutex}; pub struct FrameWriter { writer: Mutex, + version: AtomicU16, } impl FrameWriter { pub fn new(writer: W) -> Self { - Self { writer: Mutex::new(writer) } + Self { writer: Mutex::new(writer), version: AtomicU16::new(WORKER_PROTOCOL_VERSION) } + } + + pub fn set_version(&self, version: u16) { + self.version.store(version, Ordering::Release); } pub fn send_message(&self, message: &WorkerMessage) -> Result<(), String> { let mut writer = self.writer.lock().map_err(|_| "worker protocol writer lock poisoned".to_string())?; - message.write_blocking(&mut *writer).map_err(|error| error.to_string()) + message.write_blocking_version(&mut *writer, self.version.load(Ordering::Acquire)).map_err(|error| error.to_string()) } #[cfg(test)] @@ -48,7 +54,7 @@ impl WorkerService { pub fn run(&self, input: R, writer: &FrameWriter) -> Result<(), String> { let input = InputChannel::spawn(input); - let Some(hello) = input.next()? else { + let Some((hello_version, hello)) = input.next()? else { return Ok(()); }; let (request_id, session_id, response, version) = match hello { @@ -58,19 +64,34 @@ impl WorkerService { if response { return Err("worker received a Hello response instead of a request".to_string()); } - if version != bunkerbox_worker_protocol::WORKER_PROTOCOL_VERSION { + if version != hello_version { + return Err("worker Hello version does not match frame version".to_string()); + } + if version != WORKER_PROTOCOL_VERSION && version != WORKER_ARTIFACT_PROTOCOL_VERSION { return Err(format!("unsupported worker protocol version: {version}")); } if session_id.0 == [0; 16] { return Err("worker session ID must be nonzero".to_string()); } - writer.send_message(&WorkerMessage::hello(request_id, session_id, true))?; + writer.set_version(version); + writer.send_message(&WorkerMessage::hello_for_version(request_id, session_id, true, version))?; let mut active_upload: Option = None; loop { - let Some(message) = input.next()? else { + let Some((message_version, message)) = input.next()? else { return Ok(()); }; + if message_version != version { + send_error( + writer, + message.request_id(), + session_id, + WorkerOperation::Protocol, + WorkerErrorKind::WorkerProtocol, + "worker frame version changed during connection", + )?; + return Ok(()); + } if message.session_id() != session_id { send_error( writer, @@ -145,7 +166,7 @@ impl WorkerService { } }, WorkerMessage::Build { request_id, session_id, build } => { - self.handle_build(&input, writer, request_id, session_id, build)?; + self.handle_build(&input, writer, request_id, session_id, build, version)?; return Ok(()); } _ => { @@ -165,7 +186,7 @@ impl WorkerService { fn handle_build( &self, input: &InputChannel, writer: &FrameWriter, request_id: WorkerRequestId, session_id: WorkerSessionId, - build: bunkerbox_worker_protocol::WorkerBuild, + build: bunkerbox_worker_protocol::WorkerBuild, protocol_version: u16, ) -> Result<(), String> { let upload = match self.store.open_completed(session_id, build.upload_token()) { Ok(upload) => upload, @@ -193,34 +214,161 @@ impl WorkerService { return Ok(()); } }; - writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, exit_code))?; - - let Some(cleanup) = input.next()? else { - return Ok(()); - }; - match cleanup { - WorkerMessage::Cleanup { request_id: cleanup_request, session_id: cleanup_session, upload_token } - if cleanup_request == request_id && cleanup_session == session_id && upload_token == build.upload_token() => + let artifact_spool = if exit_code == 0 && !build.artifact_paths().is_empty() { + match ArtifactSpool::capture(&jobs, job.root(), build.artifact_paths(), build.artifact_max_file_bytes(), build.artifact_max_total_bytes()) { - match self.store.cleanup(session_id, upload_token) { - Ok(()) => writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0)), - Err(error) => send_error(writer, request_id, session_id, WorkerOperation::Cleanup, WorkerErrorKind::Cleanup, &error), + Ok(spool) => Some(spool), + Err(error) => { + send_error(writer, request_id, session_id, WorkerOperation::Artifact, WorkerErrorKind::Artifact, &error)?; + return self.wait_for_cleanup( + input, + writer, + BuildCleanup { request_id, session_id, upload_token: build.upload_token(), protocol_version }, + ); } } - other => send_error( - writer, - other.request_id(), + } else { + None + }; + + writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, exit_code))?; + if let Some(spool) = artifact_spool { + writer.send_message(&WorkerMessage::ArtifactManifest { + request_id, session_id, - WorkerOperation::Cleanup, - WorkerErrorKind::WorkerProtocol, - "unexpected worker cleanup message", - ), + artifact_set_id: spool.artifact_set_id(), + entries: spool.entries().to_vec(), + total_bytes: spool.total_bytes(), + })?; + return self.wait_for_artifacts( + input, + writer, + BuildCleanup { request_id, session_id, upload_token: build.upload_token(), protocol_version }, + spool, + ); + } + self.wait_for_cleanup(input, writer, BuildCleanup { request_id, session_id, upload_token: build.upload_token(), protocol_version }) + } + + fn wait_for_cleanup(&self, input: &InputChannel, writer: &FrameWriter, cleanup: BuildCleanup) -> Result<(), String> { + let BuildCleanup { request_id, session_id, upload_token, protocol_version } = cleanup; + loop { + let Some((version, cleanup)) = input.next()? else { + return Ok(()); + }; + if version != protocol_version { + return Err("worker cleanup frame version changed".to_string()); + } + match cleanup { + WorkerMessage::Cleanup { request_id: cleanup_request, session_id: cleanup_session, upload_token: received_upload } + if cleanup_request == request_id && cleanup_session == session_id && received_upload == upload_token => + { + return match self.store.cleanup(session_id, received_upload) { + Ok(()) => writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0)), + Err(error) => send_error(writer, request_id, session_id, WorkerOperation::Cleanup, WorkerErrorKind::Cleanup, &error), + }; + } + other => { + send_error( + writer, + other.request_id(), + session_id, + WorkerOperation::Cleanup, + WorkerErrorKind::WorkerProtocol, + "unexpected worker cleanup message", + )?; + } + } + } + } + + fn wait_for_artifacts( + &self, input: &InputChannel, writer: &FrameWriter, cleanup: BuildCleanup, spool: ArtifactSpool, + ) -> Result<(), String> { + let BuildCleanup { request_id, session_id, upload_token, protocol_version } = cleanup; + let mut fetched = vec![false; spool.entries().len()]; + loop { + let Some((version, message)) = input.next()? else { + return Ok(()); + }; + if version != protocol_version { + return Err("worker artifact frame version changed".to_string()); + } + match message { + WorkerMessage::FetchArtifact { request_id: fetch_request, session_id: fetch_session, artifact_set_id, entry_index } + if fetch_session == session_id && artifact_set_id == spool.artifact_set_id() => + { + let index = usize::try_from(entry_index).map_err(|_| "worker artifact index is invalid".to_string())?; + if index >= fetched.len() || fetched[index] { + send_error( + writer, + fetch_request, + session_id, + WorkerOperation::Artifact, + WorkerErrorKind::Artifact, + "worker artifact was fetched more than once or is out of range", + )?; + continue; + } + let mut file = match spool.open_entry(index) { + Ok(file) => file, + Err(error) => { + send_error(writer, fetch_request, session_id, WorkerOperation::Artifact, WorkerErrorKind::Artifact, &error)?; + continue; + } + }; + let mut offset = 0u64; + let mut buffer = [0u8; MAX_WORKER_CHUNK_BYTES]; + loop { + let count = file.read(&mut buffer).map_err(|error| format!("read worker artifact: {error}"))?; + if count == 0 { + break; + } + writer.send_message(&WorkerMessage::ArtifactChunk { + request_id: fetch_request, + session_id, + artifact_set_id, + entry_index, + offset, + data: buffer[..count].to_vec(), + })?; + offset = offset.checked_add(count as u64).ok_or_else(|| "worker artifact offset overflow".to_string())?; + } + writer.send_message(&WorkerMessage::ArtifactComplete { request_id: fetch_request, session_id, artifact_set_id, entry_index })?; + fetched[index] = true; + } + WorkerMessage::Cleanup { request_id: cleanup_request, session_id: cleanup_session, upload_token: received_upload } + if cleanup_request == request_id && cleanup_session == session_id && received_upload == upload_token => + { + return match self.store.cleanup(session_id, received_upload) { + Ok(()) => writer.send_message(&WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0)), + Err(error) => send_error(writer, request_id, session_id, WorkerOperation::Cleanup, WorkerErrorKind::Cleanup, &error), + }; + } + other => { + send_error( + writer, + other.request_id(), + session_id, + WorkerOperation::Artifact, + WorkerErrorKind::WorkerProtocol, + "unexpected worker artifact message", + )?; + } + } } } } +struct BuildCleanup { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + upload_token: WorkerUploadId, + protocol_version: u16, +} + struct InputChannel { - receiver: Receiver, String>>, + receiver: Receiver, String>>, eof_seen: Arc, pending: Arc, } @@ -233,7 +381,7 @@ impl InputChannel { let pending = Arc::new(AtomicUsize::new(0)); let pending_for_thread = pending.clone(); std::thread::spawn(move || loop { - match bunkerbox_worker_protocol::WorkerMessage::read_blocking_optional(&mut input) { + match bunkerbox_worker_protocol::WorkerMessage::read_blocking_optional_versioned(&mut input) { Ok(Some(message)) => { pending_for_thread.fetch_add(1, Ordering::Release); if sender.send(Ok(Some(message))).is_err() { @@ -255,7 +403,7 @@ impl InputChannel { Self { receiver, eof_seen, pending } } - fn next(&self) -> Result, String> { + fn next(&self) -> Result, String> { let result = self.receiver.recv().map_err(|_| "worker input reader stopped".to_string())?; if matches!(&result, Ok(Some(_))) { self.pending.fetch_sub(1, Ordering::AcqRel); diff --git a/src/artifact.rs b/src/artifact.rs new file mode 100644 index 0000000..591e0df --- /dev/null +++ b/src/artifact.rs @@ -0,0 +1,603 @@ +use sha2::{Digest, Sha256}; +use std::collections::BTreeSet; +use std::ffi::CString; +use std::fs::{self, File, OpenOptions}; +use std::io::{self, Read, Write}; +use std::os::fd::{AsRawFd, FromRawFd}; +use std::os::unix::fs::{MetadataExt, OpenOptionsExt, PermissionsExt}; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::Duration; + +pub const MAX_ARTIFACT_ENTRIES: usize = 256; +pub const MAX_ARTIFACT_PATH_BYTES: usize = 4 * 1024; +pub const MAX_ARTIFACT_FILE_BYTES: u64 = 256 * 1024 * 1024; +pub const MAX_ARTIFACT_TOTAL_BYTES: u64 = 512 * 1024 * 1024; +pub const DEFAULT_ARTIFACT_TIMEOUT: Duration = Duration::from_secs(30); +pub const DEFAULT_MAX_ARTIFACT_FILE_BYTES: u64 = MAX_ARTIFACT_FILE_BYTES; +pub const DEFAULT_MAX_ARTIFACT_TOTAL_BYTES: u64 = MAX_ARTIFACT_TOTAL_BYTES; +pub const DEFAULT_MAX_ARTIFACT_ENTRIES: usize = MAX_ARTIFACT_ENTRIES; +const COPY_BUFFER_BYTES: usize = 64 * 1024; +const ARTIFACT_ROOT: &str = ".bunkerbox"; +const ARTIFACT_DIRECTORY: &str = "artifacts"; +const STAGING_DIRECTORY: &str = ".staging"; +static NEXT_SPOOL_ID: AtomicU64 = AtomicU64::new(1); + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ArtifactPolicy { + paths: Vec, +} + +impl ArtifactPolicy { + pub fn new(paths: Vec) -> Result { + if paths.len() > MAX_ARTIFACT_ENTRIES { + return Err(format!("artifact policy exceeds maximum entry count {MAX_ARTIFACT_ENTRIES}")); + } + let mut seen = BTreeSet::new(); + for path in &paths { + validate_artifact_path(path)?; + if !seen.insert(path.clone()) { + return Err(format!("duplicate artifact path: {path}")); + } + } + Ok(Self { paths }) + } + + pub fn paths(&self) -> &[String] { + &self.paths + } + + pub fn is_enabled(&self) -> bool { + !self.paths.is_empty() + } + + pub fn validate_limits(&self, limits: ArtifactLimits) -> Result<(), String> { + if self.paths.len() > limits.max_entries { + return Err(format!("artifact policy exceeds configured entry count {}", limits.max_entries)); + } + Ok(()) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ArtifactLimits { + pub timeout: Duration, + pub max_entries: usize, + pub max_file_bytes: u64, + pub max_total_bytes: u64, +} + +impl Default for ArtifactLimits { + fn default() -> Self { + Self { + timeout: DEFAULT_ARTIFACT_TIMEOUT, + max_entries: DEFAULT_MAX_ARTIFACT_ENTRIES, + max_file_bytes: DEFAULT_MAX_ARTIFACT_FILE_BYTES, + max_total_bytes: DEFAULT_MAX_ARTIFACT_TOTAL_BYTES, + } + } +} + +impl ArtifactLimits { + pub fn new(timeout: Duration, max_entries: usize, max_file_bytes: u64, max_total_bytes: u64) -> Result { + if timeout.is_zero() { + return Err("artifact timeout must be positive".to_string()); + } + if max_entries == 0 || max_entries > MAX_ARTIFACT_ENTRIES { + return Err(format!("artifact entry limit must be between 1 and {MAX_ARTIFACT_ENTRIES}")); + } + if max_file_bytes == 0 || max_file_bytes > MAX_ARTIFACT_FILE_BYTES { + return Err(format!("artifact file limit must be between 1 and {MAX_ARTIFACT_FILE_BYTES}")); + } + if max_total_bytes == 0 || max_total_bytes > MAX_ARTIFACT_TOTAL_BYTES { + return Err(format!("artifact total limit must be between 1 and {MAX_ARTIFACT_TOTAL_BYTES}")); + } + Ok(Self { timeout, max_entries, max_file_bytes, max_total_bytes }) + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ArtifactEntry { + path: String, + mode: u32, + size: u64, + digest: [u8; 32], +} + +impl ArtifactEntry { + pub fn new(path: impl Into, mode: u32, size: u64, digest: [u8; 32]) -> Result { + let path = path.into(); + validate_artifact_path(&path)?; + if mode & !0o777 != 0 { + return Err(format!("artifact mode has unsupported bits: {mode:o}")); + } + Ok(Self { path, mode, size, digest }) + } + + pub fn path(&self) -> &str { + &self.path + } + + pub fn mode(&self) -> u32 { + self.mode + } + + pub fn size(&self) -> u64 { + self.size + } + + pub fn digest(&self) -> &[u8; 32] { + &self.digest + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ArtifactManifest { + entries: Vec, + total_bytes: u64, +} + +impl ArtifactManifest { + pub fn new(entries: Vec, total_bytes: u64, policy: &ArtifactPolicy, limits: ArtifactLimits) -> Result { + policy.validate_limits(limits)?; + if entries.len() != policy.paths.len() { + return Err(format!("artifact manifest has {} entries but {} are required", entries.len(), policy.paths.len())); + } + if entries.len() > limits.max_entries { + return Err("artifact manifest exceeds configured entry count".to_string()); + } + let mut total = 0u64; + let mut seen = BTreeSet::new(); + for (entry, expected_path) in entries.iter().zip(policy.paths.iter()) { + if entry.path != *expected_path { + return Err(format!("artifact manifest path does not match policy: {}", entry.path)); + } + if !seen.insert(entry.path.as_str()) { + return Err(format!("artifact manifest contains duplicate path: {}", entry.path)); + } + if entry.size > limits.max_file_bytes { + return Err(format!("artifact exceeds configured per-file limit: {}", entry.path)); + } + total = total.checked_add(entry.size).ok_or_else(|| "artifact manifest total size overflow".to_string())?; + if total > limits.max_total_bytes { + return Err("artifact manifest exceeds configured total size".to_string()); + } + } + if total != total_bytes { + return Err("artifact manifest total size does not match entries".to_string()); + } + Ok(Self { entries, total_bytes }) + } + + pub fn entries(&self) -> &[ArtifactEntry] { + &self.entries + } + + pub fn total_bytes(&self) -> u64 { + self.total_bytes + } + + pub fn entry(&self, index: usize) -> Option<&ArtifactEntry> { + self.entries.get(index) + } +} + +pub struct LocalArtifactSpool { + path: PathBuf, + root: File, + manifest: ArtifactManifest, +} + +impl LocalArtifactSpool { + pub fn capture(job_root: &Path, parent: &Path, policy: &ArtifactPolicy, limits: ArtifactLimits) -> Result { + policy.validate_limits(limits)?; + let path = create_unique_directory(parent, "artifact-spool")?; + let root = match open_directory(&path) { + Ok(root) => root, + Err(error) => { + let _ = fs::remove_dir_all(&path); + return Err(format!("open local artifact spool: {error}")); + } + }; + + let result = (|| { + let mut entries = Vec::with_capacity(policy.paths.len()); + let mut total = 0u64; + for relative in policy.paths() { + let source = open_regular_file(job_root, relative)?; + let metadata = source.metadata().map_err(|error| format!("stat artifact {relative}: {error}"))?; + if metadata.nlink() != 1 { + return Err(format!("artifact is a hard-link alias: {relative}")); + } + let size = metadata.len(); + if size > limits.max_file_bytes { + return Err(format!("artifact exceeds configured per-file limit: {relative}")); + } + total = total.checked_add(size).ok_or_else(|| "artifact total size overflow".to_string())?; + if total > limits.max_total_bytes { + return Err("artifacts exceed configured total size".to_string()); + } + let mode = metadata.mode() & 0o777; + let destination = create_relative_file(&root, relative, mode)?; + let digest = copy_and_hash(&source, &destination, size, relative)?; + let after = source.metadata().map_err(|error| format!("restat artifact {relative}: {error}"))?; + if after.dev() != metadata.dev() + || after.ino() != metadata.ino() + || after.len() != metadata.len() + || after.mode() & 0o777 != metadata.mode() & 0o777 + || after.nlink() != metadata.nlink() + { + return Err(format!("artifact changed during capture: {relative}")); + } + entries.push(ArtifactEntry::new(relative.clone(), mode, size, digest)?); + } + ArtifactManifest::new(entries, total, policy, limits) + })(); + + match result { + Ok(manifest) => Ok(Self { path, root, manifest }), + Err(error) => { + let _ = fs::remove_dir_all(&path); + Err(error) + } + } + } + + pub fn manifest(&self) -> &ArtifactManifest { + &self.manifest + } + + pub fn open_entry(&self, index: usize) -> Result { + let entry = self.manifest.entry(index).ok_or_else(|| "artifact index is out of range".to_string())?; + let file = open_regular_file_from_fd(&self.root, &entry.path)?; + let metadata = file.metadata().map_err(|error| format!("stat spooled artifact {}: {error}", entry.path))?; + if metadata.nlink() != 1 || metadata.len() != entry.size || metadata.mode() & 0o777 != entry.mode { + return Err(format!("spooled artifact metadata changed: {}", entry.path)); + } + Ok(file) + } +} + +impl Drop for LocalArtifactSpool { + fn drop(&mut self) { + let _ = &self.root; + let _ = fs::remove_dir_all(&self.path); + } +} + +pub struct ArtifactPublication { + staging: PathBuf, + final_path: PathBuf, + manifest: ArtifactManifest, + written: BTreeSet, + published: bool, +} + +impl ArtifactPublication { + pub fn new(workspace_root: &Path, request_id: [u8; 16], manifest: ArtifactManifest) -> Result { + let artifacts = ensure_directory_path(workspace_root, &[ARTIFACT_ROOT, ARTIFACT_DIRECTORY])?; + let staging_root = ensure_directory_path(&artifacts, &[STAGING_DIRECTORY])?; + let request = hex_id(request_id); + let staging = staging_root.join(&request); + let final_path = artifacts.join(&request); + if fs::symlink_metadata(&staging).is_ok() || fs::symlink_metadata(&final_path).is_ok() { + return Err(format!("artifact destination already exists for request {request}")); + } + fs::create_dir(&staging).map_err(|error| format!("create artifact staging directory: {error}"))?; + set_private_mode(&staging)?; + Ok(Self { staging, final_path, manifest, written: BTreeSet::new(), published: false }) + } + + pub fn begin(&self, index: usize) -> Result { + if self.written.contains(&index) { + return Err(format!("artifact entry was already written: {index}")); + } + let entry = self.manifest.entry(index).ok_or_else(|| "artifact index is out of range".to_string())?.clone(); + let components = artifact_components(&entry.path)?; + let (name, parents) = components.split_last().ok_or_else(|| "artifact path is empty".to_string())?; + let parent = ensure_directory_components(&self.staging, parents)?; + let path = parent.join(name); + let file = OpenOptions::new() + .write(true) + .create_new(true) + .mode(0o600) + .custom_flags(libc::O_NOFOLLOW | libc::O_CLOEXEC) + .open(&path) + .map_err(|error| format!("create artifact staging file {}: {error}", entry.path))?; + Ok(ArtifactWriter { index, entry, file, written: 0, hasher: Sha256::new() }) + } + + pub fn complete(&mut self, writer: ArtifactWriter) -> Result<(), String> { + let index = writer.index; + writer.finish()?; + if !self.written.insert(index) { + return Err(format!("artifact entry was completed twice: {index}")); + } + Ok(()) + } + + pub fn copy_from_reader(&mut self, index: usize, reader: &mut R) -> Result<(), String> { + let mut writer = self.begin(index)?; + let mut buffer = [0u8; COPY_BUFFER_BYTES]; + loop { + let count = reader.read(&mut buffer).map_err(|error| format!("read artifact spool: {error}"))?; + if count == 0 { + break; + } + writer.write_chunk(&buffer[..count])?; + } + self.complete(writer) + } + + pub fn publish(mut self) -> Result<(), String> { + if self.written.len() != self.manifest.entries.len() { + return Err("artifact publication is missing entries".to_string()); + } + if fs::symlink_metadata(&self.final_path).is_ok() { + return Err("artifact destination appeared before publication".to_string()); + } + fs::rename(&self.staging, &self.final_path).map_err(|error| format!("publish artifacts: {error}"))?; + let parent = self.final_path.parent().ok_or_else(|| "artifact destination has no parent".to_string())?; + File::open(parent).and_then(|file| file.sync_all()).map_err(|error| format!("flush artifact destination: {error}"))?; + self.published = true; + Ok(()) + } +} + +impl Drop for ArtifactPublication { + fn drop(&mut self) { + if !self.published { + let _ = fs::remove_dir_all(&self.staging); + } + } +} + +pub struct ArtifactWriter { + index: usize, + entry: ArtifactEntry, + file: File, + written: u64, + hasher: Sha256, +} + +impl ArtifactWriter { + pub fn expected_offset(&self) -> Result { + self.file.metadata().map(|metadata| metadata.len()).map_err(|error| format!("stat artifact staging file: {error}")) + } + + pub fn write_chunk(&mut self, bytes: &[u8]) -> Result<(), String> { + let end = self.written.checked_add(bytes.len() as u64).ok_or_else(|| "artifact size overflow".to_string())?; + if end > self.entry.size { + return Err(format!("artifact exceeds manifest size: {}", self.entry.path)); + } + self.file.write_all(bytes).map_err(|error| format!("write artifact staging file: {error}"))?; + self.hasher.update(bytes); + self.written = end; + Ok(()) + } + + fn finish(mut self) -> Result<(), String> { + self.file.flush().map_err(|error| format!("flush artifact staging file: {error}"))?; + self.file.sync_all().map_err(|error| format!("sync artifact staging file: {error}"))?; + let metadata = self.file.metadata().map_err(|error| format!("stat artifact staging file: {error}"))?; + if self.written != self.entry.size || metadata.len() != self.entry.size { + return Err(format!("artifact size does not match manifest: {}", self.entry.path)); + } + if metadata.nlink() != 1 { + return Err(format!("artifact staging file has an invalid link count: {}", self.entry.path)); + } + if self.hasher.finalize().as_slice() != self.entry.digest { + return Err(format!("artifact digest does not match manifest: {}", self.entry.path)); + } + if unsafe { libc::fchmod(self.file.as_raw_fd(), self.entry.mode as libc::mode_t) } != 0 { + return Err(format!("set artifact mode {}: {}", self.entry.path, io::Error::last_os_error())); + } + Ok(()) + } +} + +pub fn validate_artifact_path(path: &str) -> Result<(), String> { + if path.is_empty() || path.len() > MAX_ARTIFACT_PATH_BYTES || path.as_bytes().contains(&0) { + return Err("artifact path is empty, too long, or contains NUL".to_string()); + } + if path.starts_with('/') || path.starts_with('\\') || path.contains('\\') { + return Err(format!("artifact path must be relative: {path}")); + } + if path.bytes().any(|byte| matches!(byte, b'*' | b'?' | b'[' | b']' | b'{' | b'}')) { + return Err(format!("artifact path must not contain glob syntax: {path}")); + } + if path.split('/').any(|component| component.is_empty() || component == "." || component == "..") { + return Err(format!("artifact path is not normalized: {path}")); + } + Ok(()) +} + +fn artifact_components(path: &str) -> Result, String> { + validate_artifact_path(path)?; + Ok(path.split('/').collect()) +} + +fn create_unique_directory(parent: &Path, label: &str) -> Result { + fs::create_dir_all(parent).map_err(|error| format!("create {label} parent: {error}"))?; + for _ in 0..32 { + let name = format!(".bunkerbox-{label}-{}-{}", std::process::id(), NEXT_SPOOL_ID.fetch_add(1, Ordering::Relaxed)); + let path = parent.join(name); + match fs::create_dir(&path) { + Ok(()) => { + set_private_mode(&path)?; + return Ok(path); + } + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(format!("create {label}: {error}")), + } + } + Err(format!("could not reserve a private {label}")) +} + +fn ensure_directory_path(root: &Path, components: &[&str]) -> Result { + let mut current = root.to_path_buf(); + for component in components { + if component.is_empty() || *component == "." || *component == ".." || component.contains('/') { + return Err("artifact destination contains an invalid component".to_string()); + } + current.push(component); + match fs::symlink_metadata(¤t) { + Ok(metadata) if metadata.file_type().is_dir() => {} + Ok(_) => return Err(format!("artifact destination component is not a directory: {}", current.display())), + Err(error) if error.kind() == io::ErrorKind::NotFound => { + fs::create_dir(¤t).map_err(|create_error| format!("create artifact destination directory: {create_error}"))?; + set_private_mode(¤t)?; + } + Err(error) => return Err(format!("inspect artifact destination: {error}")), + } + } + Ok(current) +} + +fn ensure_directory_components(root: &Path, components: &[&str]) -> Result { + let mut current = root.to_path_buf(); + for component in components { + if component.is_empty() || *component == "." || *component == ".." || component.contains('/') { + return Err("artifact path contains an invalid component".to_string()); + } + current.push(component); + match fs::symlink_metadata(¤t) { + Ok(metadata) if metadata.file_type().is_dir() => {} + Ok(_) => return Err(format!("artifact parent is not a directory: {}", current.display())), + Err(error) if error.kind() == io::ErrorKind::NotFound => { + fs::create_dir(¤t).map_err(|create_error| format!("create artifact parent: {create_error}"))?; + set_private_mode(¤t)?; + } + Err(error) => return Err(format!("inspect artifact parent: {error}")), + } + } + Ok(current) +} + +fn open_directory(path: &Path) -> io::Result { + OpenOptions::new().read(true).custom_flags(libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC).open(path) +} + +fn open_directory_at(parent: &File, name: &str) -> io::Result { + open_at(parent, name, libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) +} + +fn open_file_at(parent: &File, name: &str) -> io::Result { + open_at(parent, name, libc::O_RDONLY | libc::O_NOFOLLOW | libc::O_CLOEXEC) +} + +fn open_at(parent: &File, name: &str, flags: i32) -> io::Result { + let name = CString::new(name.as_bytes()).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "NUL in artifact path"))?; + let fd = unsafe { libc::openat(parent.as_raw_fd(), name.as_ptr(), flags, 0) }; + if fd < 0 { + return Err(io::Error::last_os_error()); + } + Ok(unsafe { File::from_raw_fd(fd) }) +} + +fn create_file_at(parent: &File, name: &str, mode: u32) -> io::Result { + let name = CString::new(name.as_bytes()).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "NUL in artifact path"))?; + let fd = unsafe { + libc::openat( + parent.as_raw_fd(), + name.as_ptr(), + libc::O_WRONLY | libc::O_CREAT | libc::O_EXCL | libc::O_NOFOLLOW | libc::O_CLOEXEC, + mode as libc::mode_t, + ) + }; + if fd < 0 { + return Err(io::Error::last_os_error()); + } + Ok(unsafe { File::from_raw_fd(fd) }) +} + +fn open_relative_directory(root: &Path, components: &[&str]) -> Result { + let mut current = open_directory(root).map_err(|error| format!("open artifact root: {error}"))?; + for component in components { + current = open_directory_at(¤t, component).map_err(|error| format!("open artifact directory {component}: {error}"))?; + } + Ok(current) +} + +fn open_regular_file(root: &Path, relative: &str) -> Result { + let components = artifact_components(relative)?; + let (name, parents) = components.split_last().ok_or_else(|| "artifact path is empty".to_string())?; + let parent = open_relative_directory(root, parents)?; + let file = open_file_at(&parent, name).map_err(|error| format!("open artifact {relative}: {error}"))?; + let metadata = file.metadata().map_err(|error| format!("stat artifact {relative}: {error}"))?; + if !metadata.file_type().is_file() { + return Err(format!("artifact is not a regular file: {relative}")); + } + Ok(file) +} + +fn open_regular_file_from_fd(root: &File, relative: &str) -> Result { + let components = artifact_components(relative)?; + let (name, parents) = components.split_last().ok_or_else(|| "artifact path is empty".to_string())?; + let mut parent = root.try_clone().map_err(|error| format!("clone artifact spool root: {error}"))?; + for component in parents { + parent = open_directory_at(&parent, component).map_err(|error| format!("open spooled artifact directory: {error}"))?; + } + let file = open_file_at(&parent, name).map_err(|error| format!("open spooled artifact: {error}"))?; + let metadata = file.metadata().map_err(|error| format!("stat spooled artifact: {error}"))?; + if !metadata.file_type().is_file() { + return Err(format!("spooled artifact is not a regular file: {relative}")); + } + Ok(file) +} + +fn create_relative_file(root: &File, relative: &str, mode: u32) -> Result { + let components = artifact_components(relative)?; + let (name, parents) = components.split_last().ok_or_else(|| "artifact path is empty".to_string())?; + let mut parent = root.try_clone().map_err(|error| format!("clone artifact spool root: {error}"))?; + for component in parents { + match unsafe { libc::mkdirat(parent.as_raw_fd(), CString::new(*component).unwrap().as_ptr(), 0o700) } { + 0 => {} + -1 if io::Error::last_os_error().kind() == io::ErrorKind::AlreadyExists => {} + _ => return Err(format!("create artifact spool directory: {}", io::Error::last_os_error())), + } + parent = open_directory_at(&parent, component).map_err(|error| format!("open artifact spool directory: {error}"))?; + } + let file = create_file_at(&parent, name, mode & 0o777).map_err(|error| format!("create artifact spool file {relative}: {error}"))?; + if unsafe { libc::fchmod(file.as_raw_fd(), (mode & 0o777) as libc::mode_t) } != 0 { + return Err(format!("set artifact spool file mode: {}", io::Error::last_os_error())); + } + Ok(file) +} + +fn copy_and_hash(source: &File, destination: &File, expected_size: u64, path: &str) -> Result<[u8; 32], String> { + let mut source = source.try_clone().map_err(|error| format!("clone artifact source {path}: {error}"))?; + let mut destination = destination.try_clone().map_err(|error| format!("clone artifact spool file {path}: {error}"))?; + let mut hasher = Sha256::new(); + let mut copied = 0u64; + let mut buffer = [0u8; COPY_BUFFER_BYTES]; + loop { + let count = source.read(&mut buffer).map_err(|error| format!("read artifact {path}: {error}"))?; + if count == 0 { + break; + } + copied = copied.checked_add(count as u64).ok_or_else(|| "artifact size overflow".to_string())?; + if copied > expected_size { + return Err(format!("artifact grew during capture: {path}")); + } + hasher.update(&buffer[..count]); + destination.write_all(&buffer[..count]).map_err(|error| format!("write artifact spool {path}: {error}"))?; + } + if copied != expected_size { + return Err(format!("artifact size changed during capture: {path}")); + } + destination.sync_all().map_err(|error| format!("sync artifact spool {path}: {error}"))?; + Ok(hasher.finalize().into()) +} + +fn set_private_mode(path: &Path) -> Result<(), String> { + fs::set_permissions(path, fs::Permissions::from_mode(0o700)).map_err(|error| format!("set private artifact mode: {error}")) +} + +fn hex_id(bytes: [u8; 16]) -> String { + bytes.iter().map(|byte| format!("{byte:02x}")).collect() +} + +#[cfg(test)] +#[path = "artifact_ut.rs"] +mod tests; diff --git a/src/daemon.rs b/src/daemon.rs index 5645990..65932b8 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -1,3 +1,4 @@ +use crate::artifact::{ArtifactLimits, ArtifactPolicy}; use crate::cfg::EnvMode; use crate::logging; use crate::loopback::{LoopbackBackend, RunRemoteSession}; @@ -102,6 +103,8 @@ pub struct RemoteDaemonConfig { environment: Option, backend: RemoteBackendSelection, resources: RemoteResourcePolicy, + artifact_policy: ArtifactPolicy, + artifact_limits: ArtifactLimits, } enum RemoteBackendSelection { @@ -118,6 +121,8 @@ impl RemoteDaemonConfig { environment: None, backend: RemoteBackendSelection::Loopback { tools, target_environment: std::collections::BTreeMap::new() }, resources: RemoteResourcePolicy::default(), + artifact_policy: ArtifactPolicy::default(), + artifact_limits: ArtifactLimits::default(), } } @@ -130,6 +135,8 @@ impl RemoteDaemonConfig { environment: None, backend: RemoteBackendSelection::Ssh { target: Box::new(target) }, resources: RemoteResourcePolicy::default(), + artifact_policy: ArtifactPolicy::default(), + artifact_limits: ArtifactLimits::default(), }) } @@ -150,6 +157,12 @@ impl RemoteDaemonConfig { self.resources = resources; self } + + pub fn with_artifacts(mut self, policy: ArtifactPolicy, limits: ArtifactLimits) -> Self { + self.artifact_policy = policy; + self.artifact_limits = limits; + self + } } struct RemoteComponents { @@ -163,7 +176,7 @@ impl VsockDaemon { passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, remote: RemoteDaemonConfig, ) -> Result { - let RemoteDaemonConfig { session, allowed_tools, tool_policies, environment, backend, resources } = remote; + let RemoteDaemonConfig { session, allowed_tools, tool_policies, environment, backend, resources, artifact_policy, artifact_limits } = remote; let remote_policy = match (tool_policies, environment) { (Some(tool_policies), Some(environment)) => { RemoteAuthorizationPolicy::from_policies(session.target(), session.session_id(), tool_policies, environment)? @@ -178,9 +191,12 @@ impl VsockDaemon { LoopbackBackend::new(session, tools) .with_target_environment(target_environment) .with_timeout(resources.build_timeout) - .with_output_limit(resources.max_output_bytes), + .with_output_limit(resources.max_output_bytes) + .with_artifacts(artifact_policy.clone(), artifact_limits), ), - RemoteBackendSelection::Ssh { target } => Arc::new(SshBackend::new(session, *target)?), + RemoteBackendSelection::Ssh { target } => { + Arc::new(SshBackend::new(session, *target)?.with_artifacts(artifact_policy.clone(), artifact_limits)) + } }; let remote_components = RemoteComponents { context: remote_context, policy: remote_policy, backend }; Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, remote_components) diff --git a/src/lib.rs b/src/lib.rs index bb36354..710def6 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,4 @@ +pub mod artifact; pub mod cfg; pub mod cfgsetup; pub mod clidef; diff --git a/src/loopback.rs b/src/loopback.rs index 77c8730..c41ecb5 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -1,3 +1,4 @@ +use crate::artifact::{ArtifactLimits, ArtifactPolicy, ArtifactPublication, LocalArtifactSpool}; use crate::remote::{ AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteResourcePolicy, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, WorkspaceSessionId, @@ -76,6 +77,10 @@ impl RunRemoteSession { &self.workspace_root } + pub(crate) fn jobs_root(&self) -> &Path { + &self.jobs_root + } + pub fn snapshot_store(&self) -> SnapshotStore { self.snapshot_store.clone() } @@ -309,6 +314,8 @@ pub struct LoopbackBackend { tools: Arc>, target_environment: Arc>, resources: RemoteResourcePolicy, + artifact_policy: ArtifactPolicy, + artifact_limits: ArtifactLimits, } impl LoopbackBackend { @@ -318,6 +325,8 @@ impl LoopbackBackend { tools: Arc::new(tools), target_environment: Arc::new(trusted_target_environment()), resources: RemoteResourcePolicy::default(), + artifact_policy: ArtifactPolicy::default(), + artifact_limits: ArtifactLimits::default(), } } @@ -337,6 +346,12 @@ impl LoopbackBackend { self.target_environment = Arc::new(trusted); self } + + pub fn with_artifacts(mut self, policy: ArtifactPolicy, limits: ArtifactLimits) -> Self { + self.artifact_policy = policy; + self.artifact_limits = limits; + self + } } impl RemoteBackend for LoopbackBackend { @@ -347,15 +362,27 @@ impl RemoteBackend for LoopbackBackend { let tools = self.tools.clone(); let target_environment = self.target_environment.clone(); let resources = self.resources; + let artifact_policy = self.artifact_policy.clone(); + let artifact_limits = self.artifact_limits; + let request_id = request.request_id().0; + let options = LoopbackBuildOptions { resources, artifact_policy, artifact_limits, request_id }; Box::pin(async move { match request.request().operation() { RemoteOperation::Sync(sync) => execute_sync(session, sync.retain_capability(), events).await, - RemoteOperation::Build(build) => execute_build(session, tools, target_environment, resources, build, events).await, + RemoteOperation::Build(build) => execute_build(session, tools, target_environment, options, build, events).await, } }) } } +#[derive(Clone)] +struct LoopbackBuildOptions { + resources: RemoteResourcePolicy, + artifact_policy: ArtifactPolicy, + artifact_limits: ArtifactLimits, + request_id: [u8; 16], +} + async fn execute_sync( session: Arc, retain_capability: bool, events: mpsc::Sender, ) -> Result<(), RemoteBackendError> { @@ -371,8 +398,9 @@ async fn execute_sync( async fn execute_build( session: Arc, tools: Arc>, target_environment: Arc>, - resources: RemoteResourcePolicy, build: &crate::remote::RemoteBuild, events: mpsc::Sender, + options: LoopbackBuildOptions, build: &crate::remote::RemoteBuild, events: mpsc::Sender, ) -> Result<(), RemoteBackendError> { + let LoopbackBuildOptions { resources, artifact_policy, artifact_limits, request_id } = options; let snapshot = session.claim_snapshot(build.snapshot_id()).map_err(RemoteBackendError::Failed)?; let executable = tools .get(build.tool().as_str()) @@ -523,7 +551,49 @@ async fn execute_build( return Err(error); } let status = child_status.unwrap()?; - send_event(&events, RemoteBackendEvent::Completed { exit_code: status.code().unwrap_or(-1) }).await + let exit_code = status.code().unwrap_or(-1); + if exit_code == 0 && artifact_policy.is_enabled() { + let job_root = job_path.clone(); + let workspace_root = session.workspace_root().to_path_buf(); + let spool_parent = session.jobs_root().to_path_buf(); + let policy = artifact_policy; + let limits = artifact_limits; + let retrieval = tokio::time::timeout( + limits.timeout, + tokio::task::spawn_blocking(move || { + let spool = LocalArtifactSpool::capture(&job_root, &spool_parent, &policy, limits) + .map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactManifest, message: error })?; + let manifest = spool.manifest().clone(); + let mut publication = ArtifactPublication::new(&workspace_root, request_id, manifest) + .map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactTransfer, message: error })?; + for index in 0..spool.manifest().entries().len() { + let mut source = spool.open_entry(index).map_err(|error| RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::ArtifactTransfer, + message: error, + })?; + publication.copy_from_reader(index, &mut source).map_err(|error| RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::ArtifactTransfer, + message: error, + })?; + } + publication + .publish() + .map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactTransfer, message: error }) + }), + ) + .await; + match retrieval { + Ok(Ok(result)) => result?, + Ok(Err(error)) => return Err(RemoteBackendError::Failed(format!("artifact worker failed: {error}"))), + Err(_) => { + return Err(RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::ArtifactTransfer, + message: "artifact retrieval timed out".to_string(), + }) + } + } + } + send_event(&events, RemoteBackendEvent::Completed { exit_code }).await } #[derive(Clone, Copy)] diff --git a/src/main.rs b/src/main.rs index 6bcabce..a436c3a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -191,6 +191,9 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option "worker version", Self::WorkerProtocol => "worker protocol", Self::SnapshotTransfer => "snapshot transfer", + Self::ArtifactManifest => "artifact manifest", + Self::ArtifactTransfer => "artifact transfer", Self::Disconnect => "disconnect", Self::Cleanup => "cleanup", } diff --git a/src/remote_target.rs b/src/remote_target.rs index 316a620..5cc0ef7 100644 --- a/src/remote_target.rs +++ b/src/remote_target.rs @@ -1,3 +1,4 @@ +use crate::artifact::{ArtifactLimits, ArtifactPolicy}; use serde::de::{self, MapAccess, Visitor}; use serde::Deserialize; use std::collections::BTreeMap; @@ -40,6 +41,7 @@ pub struct ResourceLimits { pub sync_timeout: Duration, pub build_timeout: Duration, pub max_output_bytes: u64, + pub artifact: ArtifactLimits, } impl ResourceLimits { @@ -62,6 +64,10 @@ impl ResourceLimits { pub fn max_output_bytes(&self) -> u64 { self.max_output_bytes } + + pub fn artifact_limits(&self) -> ArtifactLimits { + self.artifact + } } /// An SSH target after all configuration and local-file checks have passed. @@ -155,6 +161,7 @@ impl SshTarget { pub struct ProjectBinding { pub backend: BackendMode, pub target: Option, + pub artifacts: ArtifactPolicy, } impl ProjectBinding { @@ -165,6 +172,10 @@ impl ProjectBinding { pub fn target(&self) -> Option<&str> { self.target.as_deref() } + + pub fn artifacts(&self) -> &ArtifactPolicy { + &self.artifacts + } } /// The result of resolving a canonical project path. @@ -176,6 +187,7 @@ pub struct ResolvedBackend { pub project_root: PathBuf, pub backend: BackendMode, pub target: Option, + pub artifacts: ArtifactPolicy, } impl fmt::Debug for ResolvedBackend { @@ -205,6 +217,10 @@ impl ResolvedBackend { pub fn target(&self) -> Option<&SshTarget> { self.target.as_ref() } + + pub fn artifacts(&self) -> &ArtifactPolicy { + &self.artifacts + } } /// Configuration loaded from the host's remote-targets file. @@ -283,11 +299,18 @@ impl RemoteTargetConfig { let binding = self.projects.get(&canonical).ok_or_else(|| format!("project has no remote backend binding: {}", canonical.display()))?; match binding.backend { - BackendMode::Loopback => Ok(ResolvedBackend { project_root: canonical, backend: BackendMode::Loopback, target: None }), + BackendMode::Loopback => { + Ok(ResolvedBackend { project_root: canonical, backend: BackendMode::Loopback, target: None, artifacts: binding.artifacts.clone() }) + } BackendMode::Ssh => { let target_name = binding.target.as_deref().ok_or_else(|| "SSH project binding is missing a target".to_string())?; let target = self.targets.get(target_name).ok_or_else(|| format!("unknown SSH target '{target_name}'"))?; - Ok(ResolvedBackend { project_root: canonical, backend: BackendMode::Ssh, target: Some(target.clone()) }) + Ok(ResolvedBackend { + project_root: canonical, + backend: BackendMode::Ssh, + target: Some(target.clone()), + artifacts: binding.artifacts.clone(), + }) } } } @@ -309,16 +332,17 @@ impl RemoteTargetConfig { let mut projects = BTreeMap::new(); for (project, binding) in raw.projects.0 { let canonical = validate_project_binding_path(&project)?; - validate_project_binding(&binding)?; + let artifacts = validate_project_binding(&binding)?; if binding.backend == BackendMode::Ssh { let Some(target_name) = binding.target.as_deref() else { return Err("SSH project binding requires a target".to_string()); }; - if !targets.contains_key(target_name) { - return Err(format!("unknown SSH target '{target_name}'")); - } + let target = targets.get(target_name).ok_or_else(|| format!("unknown SSH target '{target_name}'"))?; + artifacts.validate_limits(target.resources().artifact_limits())?; + } else { + artifacts.validate_limits(ArtifactLimits::default())?; } - if projects.insert(canonical, binding.into_public()).is_some() { + if projects.insert(canonical, binding.into_public(artifacts)).is_some() { return Err("duplicate project binding after canonicalization".to_string()); } } @@ -468,14 +492,23 @@ struct RawProjectBinding { backend: BackendMode, #[serde(default)] target: Option, + #[serde(default)] + artifacts: Option, } impl RawProjectBinding { - fn into_public(self) -> ProjectBinding { - ProjectBinding { backend: self.backend, target: self.target } + fn into_public(self, artifacts: ArtifactPolicy) -> ProjectBinding { + ProjectBinding { backend: self.backend, target: self.target, artifacts } } } +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct RawArtifacts { + #[serde(default)] + paths: Vec, +} + #[derive(Deserialize)] #[serde(deny_unknown_fields)] struct RawResources { @@ -493,6 +526,14 @@ struct RawResources { build_timeout: RawQuantity, #[serde(rename = "max-output-bytes", alias = "max-output", alias = "max_output_bytes", alias = "max_output")] max_output: RawQuantity, + #[serde(default, rename = "artifact-timeout-seconds", alias = "artifact-timeout", alias = "artifact_timeout")] + artifact_timeout: Option, + #[serde(default, rename = "max-artifact-bytes", alias = "max-artifact-bytes-per-file", alias = "max_artifact_bytes")] + max_artifact_bytes: Option, + #[serde(default, rename = "max-artifact-total-bytes", alias = "max_artifact_total_bytes")] + max_artifact_total_bytes: Option, + #[serde(default, rename = "max-artifact-entries", alias = "max_artifact_entries")] + max_artifact_entries: Option, } #[derive(Deserialize)] @@ -557,14 +598,14 @@ fn validate_target(name: String, raw: RawTarget) -> Result { }) } -fn validate_project_binding(binding: &RawProjectBinding) -> Result<(), String> { +fn validate_project_binding(binding: &RawProjectBinding) -> Result { if let Some(target) = &binding.target { validate_name("project target name", target)?; } if binding.backend == BackendMode::Ssh && binding.target.is_none() { return Err("SSH project binding requires a target".to_string()); } - Ok(()) + binding.artifacts.as_ref().map_or_else(|| Ok(ArtifactPolicy::default()), |artifacts| ArtifactPolicy::new(artifacts.paths.clone())) } fn validate_resources(raw: RawResources) -> Result { @@ -572,7 +613,16 @@ fn validate_resources(raw: RawResources) -> Result { let sync_timeout = parse_duration("sync-timeout", raw.sync_timeout)?; let build_timeout = parse_duration("build-timeout", raw.build_timeout)?; let max_output_bytes = parse_size("max-output", raw.max_output)?; - Ok(ResourceLimits { connect_timeout, sync_timeout, build_timeout, max_output_bytes }) + let defaults = ArtifactLimits::default(); + let artifact_timeout = raw.artifact_timeout.map_or(Ok(defaults.timeout), |value| parse_duration("artifact-timeout", value))?; + let max_artifact_bytes = raw.max_artifact_bytes.map_or(Ok(defaults.max_file_bytes), |value| parse_size("max-artifact-bytes", value))?; + let max_artifact_total_bytes = + raw.max_artifact_total_bytes.map_or(Ok(defaults.max_total_bytes), |value| parse_size("max-artifact-total-bytes", value))?; + let max_artifact_entries = raw + .max_artifact_entries + .map_or(Ok(defaults.max_entries), |value| usize::try_from(value).map_err(|_| "max-artifact-entries is too large".to_string()))?; + let artifact = ArtifactLimits::new(artifact_timeout, max_artifact_entries, max_artifact_bytes, max_artifact_total_bytes)?; + Ok(ResourceLimits { connect_timeout, sync_timeout, build_timeout, max_output_bytes, artifact }) } fn parse_duration(field: &str, quantity: RawQuantity) -> Result { diff --git a/src/ssh.rs b/src/ssh.rs index 1bde16c..b3d1562 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -1,3 +1,4 @@ +use crate::artifact::{ArtifactLimits, ArtifactManifest, ArtifactPolicy, ArtifactPublication}; use crate::loopback::{RunRemoteSession, SnapshotExportClaim}; use crate::remote::{ AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFailureClass, RemoteFuture, RemoteOperation, @@ -6,8 +7,9 @@ use crate::remote::{ use crate::remote_target::{ResourceLimits, SshTarget}; use crate::snapshot::SnapshotEntryKind; use crate::worker_protocol::{ - self, WorkerBuild, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRelativePath, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, - WorkerUploadId, MAX_WORKER_CHUNK_BYTES, MAX_WORKER_FILE_BYTES, + self, WorkerArtifactEntry, WorkerArtifactPath, WorkerArtifactSetId, WorkerBuild, WorkerErrorKind, WorkerMessage, WorkerOperation, + WorkerRelativePath, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, WorkerUploadId, MAX_WORKER_CHUNK_BYTES, MAX_WORKER_FILE_BYTES, + WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, }; use rand::RngCore; use std::collections::BTreeMap; @@ -205,12 +207,21 @@ pub struct SshBackend { target: SshTarget, factory: Arc, uploads: Arc>>, + artifact_policy: ArtifactPolicy, + artifact_limits: ArtifactLimits, } impl SshBackend { pub fn new(session: Arc, target: SshTarget) -> Result { let _ = SshLaunchSpec::from_target(&target)?; - Ok(Self { session, target, factory: Arc::new(SystemSshProcessFactory), uploads: Arc::new(Mutex::new(BTreeMap::new())) }) + Ok(Self { + artifact_limits: target.resources().artifact_limits(), + session, + target, + factory: Arc::new(SystemSshProcessFactory), + uploads: Arc::new(Mutex::new(BTreeMap::new())), + artifact_policy: ArtifactPolicy::default(), + }) } pub fn with_process_factory(mut self, factory: Arc) -> Self { @@ -218,6 +229,12 @@ impl SshBackend { self } + pub fn with_artifacts(mut self, policy: ArtifactPolicy, limits: ArtifactLimits) -> Self { + self.artifact_policy = policy; + self.artifact_limits = limits; + self + } + pub fn target(&self) -> &SshTarget { &self.target } @@ -237,8 +254,10 @@ impl RemoteBackend for SshBackend { let target = self.target.clone(); let factory = self.factory.clone(); let uploads = self.uploads.clone(); + let artifact_policy = self.artifact_policy.clone(); + let artifact_limits = self.artifact_limits; Box::pin(async move { - let backend = SshExecution { session, target, factory, uploads }; + let backend = SshExecution { session, target, factory, uploads, artifact_policy, artifact_limits }; match operation { RemoteOperation::Sync(sync) => backend.execute_sync(request_id.0, sync.retain_capability(), events).await, RemoteOperation::Build(build) => backend.execute_build(request_id.0, &build, events).await, @@ -252,6 +271,8 @@ struct SshExecution { target: SshTarget, factory: Arc, uploads: Arc>>, + artifact_policy: ArtifactPolicy, + artifact_limits: ArtifactLimits, } impl SshExecution { @@ -363,11 +384,38 @@ impl SshExecution { let worker_build = WorkerBuild::new(build.tool().as_str(), executable, build.argv().to_vec(), build.cwd().as_str(), guest_env, target_env, upload_id) .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, message: error.to_string() })?; + let worker_build = if self.artifact_policy.is_enabled() { + let paths = self + .artifact_policy + .paths() + .iter() + .map(|path| WorkerArtifactPath::new(path.clone())) + .collect::, _>>() + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error.to_string() })?; + worker_build + .with_artifacts(paths, self.artifact_limits.max_file_bytes, self.artifact_limits.max_total_bytes) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error.to_string() })? + } else { + worker_build + }; let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; let session_id = WorkerSessionId(self.session.session_id().0); - let operation = - build_and_finish(&mut connection, WorkerRequestId(request_id), session_id, worker_build, &events, upload_id, self.target.resources()); + let operation = build_and_finish( + &mut connection, + BuildPlan { + request_id: WorkerRequestId(request_id), + session_id, + build: worker_build, + events: &events, + upload_id, + resources: self.target.resources(), + artifact_policy: self.artifact_policy.clone(), + artifact_limits: self.artifact_limits, + workspace_root: self.session.workspace_root().to_path_buf(), + protocol_version: if self.artifact_policy.is_enabled() { WORKER_ARTIFACT_PROTOCOL_VERSION } else { WORKER_PROTOCOL_VERSION }, + }, + ); let result = match timeout(self.target.resources().build_timeout(), operation).await { Ok(result) => result, Err(_) => { @@ -396,6 +444,7 @@ struct WorkerConnection { writer: Option, reader: WorkerReader, stderr_task: Option>>, + version: u16, } impl WorkerConnection { @@ -407,18 +456,25 @@ impl WorkerConnection { let reader = process.take_stdout().ok_or_else(|| unavailable("SSH transport has no stdout"))?; let stderr = process.take_stderr().ok_or_else(|| unavailable("SSH transport has no stderr"))?; let stderr_task = tokio::spawn(read_diagnostic(stderr)); - Ok(Self { process, writer: Some(writer), reader, stderr_task: Some(stderr_task) }) + Ok(Self { process, writer: Some(writer), reader, stderr_task: Some(stderr_task), version: WORKER_PROTOCOL_VERSION }) } - async fn handshake(&mut self, request_id: WorkerRequestId, session_id: WorkerSessionId) -> Result<(), RemoteBackendError> { - self.write(&WorkerMessage::hello(request_id, session_id, false)).await?; - let message = self.read().await?; + async fn handshake(&mut self, request_id: WorkerRequestId, session_id: WorkerSessionId, version: u16) -> Result<(), RemoteBackendError> { + self.version = version; + self.write(&WorkerMessage::hello_for_version(request_id, session_id, false, version)).await?; + let (received_version, message) = self.read_versioned().await?; + if received_version != version { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::WorkerVersion, + message: format!("worker selected protocol version {received_version}, requested {version}"), + }); + } match message { WorkerMessage::Hello { request_id: received_request, session_id: received_session, version, response } => { if received_request != request_id || received_session != session_id { return Err(worker_protocol("worker hello correlation mismatch")); } - if version != worker_protocol::WORKER_PROTOCOL_VERSION { + if version != self.version { return Err(RemoteBackendError::Transport { class: RemoteFailureClass::WorkerVersion, message: format!("unsupported worker version: {version}"), @@ -436,11 +492,22 @@ impl WorkerConnection { async fn write(&mut self, message: &WorkerMessage) -> Result<(), RemoteBackendError> { let writer = self.writer.as_mut().ok_or_else(|| disconnected("SSH worker stdin is closed"))?; - worker_protocol::write_message(writer, message).await.map_err(worker_io_error) + worker_protocol::write_message_versioned(writer, message, self.version).await.map_err(worker_io_error) } async fn read(&mut self) -> Result { - worker_protocol::read_message(&mut self.reader).await.map_err(worker_io_error) + let (version, message) = self.read_versioned().await?; + if version != self.version { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::WorkerVersion, + message: format!("worker frame version changed from {} to {version}", self.version), + }); + } + Ok(message) + } + + async fn read_versioned(&mut self) -> Result<(u16, WorkerMessage), RemoteBackendError> { + worker_protocol::read_message_versioned(&mut self.reader).await.map_err(worker_io_error) } async fn cleanup( @@ -512,9 +579,22 @@ struct UploadPlan<'a> { cleanup: bool, } +struct BuildPlan<'a> { + request_id: WorkerRequestId, + session_id: WorkerSessionId, + build: WorkerBuild, + events: &'a tokio::sync::mpsc::Sender, + upload_id: WorkerUploadId, + resources: ResourceLimits, + artifact_policy: ArtifactPolicy, + artifact_limits: ArtifactLimits, + workspace_root: PathBuf, + protocol_version: u16, +} + async fn upload_and_finish(connection: &mut WorkerConnection, plan: UploadPlan<'_>) -> Result<(), RemoteBackendError> { let UploadPlan { request_id, session_id, upload_id, entries, export, events, cleanup } = plan; - connection.handshake(request_id, session_id).await?; + connection.handshake(request_id, session_id, WORKER_PROTOCOL_VERSION).await?; let total_bytes = export.total_file_bytes(); connection.write(&WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries: entries.to_vec() }).await?; let mut completed_bytes = 0u64; @@ -572,11 +652,10 @@ async fn upload_and_finish(connection: &mut WorkerConnection, plan: UploadPlan<' connection.finish().await } -async fn build_and_finish( - connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild, - events: &tokio::sync::mpsc::Sender, upload_id: WorkerUploadId, resources: ResourceLimits, -) -> Result<(), RemoteBackendError> { - connection.handshake(request_id, session_id).await?; +async fn build_and_finish(connection: &mut WorkerConnection, plan: BuildPlan<'_>) -> Result<(), RemoteBackendError> { + let BuildPlan { request_id, session_id, build, events, upload_id, resources, artifact_policy, artifact_limits, workspace_root, protocol_version } = + plan; + connection.handshake(request_id, session_id, protocol_version).await?; connection.write(&WorkerMessage::build(request_id, session_id, build)).await?; let mut output_bytes = 0u64; let exit_code = loop { @@ -605,11 +684,124 @@ async fn build_and_finish( _ => return Err(worker_protocol("unexpected worker build response")), } }; + if exit_code == 0 && artifact_policy.is_enabled() { + let (artifact_set_id, manifest) = match connection.read().await? { + WorkerMessage::ArtifactManifest { request_id: received_request, session_id: received_session, artifact_set_id, entries, total_bytes } => { + check_correlation(received_request, received_session, request_id, session_id, "worker artifact manifest")?; + (artifact_set_id, artifact_manifest_from_worker(entries, total_bytes, &artifact_policy, artifact_limits)?) + } + WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Artifact, kind, message)), + _ => return Err(worker_protocol("unexpected worker artifact manifest response")), + }; + let retrieval = timeout( + artifact_limits.timeout, + fetch_and_publish_artifacts(connection, request_id, session_id, artifact_set_id, &manifest, workspace_root), + ) + .await; + match retrieval { + Ok(result) => result?, + Err(_) => { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::ArtifactTransfer, + message: "artifact retrieval timed out".to_string(), + }) + } + } + } connection.cleanup(request_id, session_id, upload_id).await?; connection.finish().await?; send_event(events, RemoteBackendEvent::Completed { exit_code }).await } +async fn fetch_and_publish_artifacts( + connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, artifact_set_id: WorkerArtifactSetId, + manifest: &ArtifactManifest, workspace_root: PathBuf, +) -> Result<(), RemoteBackendError> { + let mut publication = ArtifactPublication::new(&workspace_root, request_id.0, manifest.clone()) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; + for index in 0..manifest.entries().len() { + connection + .write(&WorkerMessage::FetchArtifact { + request_id, + session_id, + artifact_set_id, + entry_index: u32::try_from(index).map_err(|_| RemoteBackendError::Transport { + class: RemoteFailureClass::ArtifactTransfer, + message: "artifact index does not fit worker protocol".to_string(), + })?, + }) + .await?; + let mut writer = publication + .begin(index) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; + loop { + match connection.read().await? { + WorkerMessage::ArtifactChunk { + request_id: received_request, + session_id: received_session, + artifact_set_id: received_set, + entry_index, + offset, + data, + } => { + check_correlation(received_request, received_session, request_id, session_id, "worker artifact chunk")?; + if received_set != artifact_set_id || entry_index != index as u32 { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::ArtifactTransfer, + message: "worker artifact chunk identity mismatch".to_string(), + }); + } + let expected = writer + .expected_offset() + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; + if offset != expected { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::ArtifactTransfer, + message: "worker artifact chunk offset is out of order".to_string(), + }); + } + writer + .write_chunk(&data) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; + } + WorkerMessage::ArtifactComplete { + request_id: received_request, + session_id: received_session, + artifact_set_id: received_set, + entry_index, + } => { + check_correlation(received_request, received_session, request_id, session_id, "worker artifact completion")?; + if received_set != artifact_set_id || entry_index != index as u32 { + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::ArtifactTransfer, + message: "worker artifact completion identity mismatch".to_string(), + }); + } + publication + .complete(writer) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; + break; + } + WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Artifact, kind, message)), + _ => return Err(worker_protocol("unexpected worker artifact transfer response")), + } + } + } + publication.publish().map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error }) +} + +fn artifact_manifest_from_worker( + entries: Vec, total_bytes: u64, policy: &ArtifactPolicy, limits: ArtifactLimits, +) -> Result { + let entries = entries + .into_iter() + .map(|entry| crate::artifact::ArtifactEntry::new(entry.path().as_str().to_string(), entry.mode(), entry.size(), *entry.digest())) + .collect::, _>>() + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error })?; + ArtifactManifest::new(entries, total_bytes, policy, limits) + .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error }) +} + fn worker_entry(entry: &crate::snapshot::SnapshotEntry) -> Result { match entry.kind() { SnapshotEntryKind::Directory => WorkerUploadEntry::directory(entry.path().as_str(), entry.mode() as u32).map_err(worker_io_error), @@ -694,6 +886,7 @@ fn worker_error(operation: WorkerOperation, kind: WorkerErrorKind, message: Stri let class = match kind { WorkerErrorKind::WorkerProtocol => RemoteFailureClass::WorkerProtocol, WorkerErrorKind::Upload | WorkerErrorKind::Sync => RemoteFailureClass::SnapshotTransfer, + WorkerErrorKind::Artifact => RemoteFailureClass::ArtifactTransfer, WorkerErrorKind::Cleanup => RemoteFailureClass::Cleanup, WorkerErrorKind::Build => return RemoteBackendError::Failed(format!("remote worker build failure ({operation:?}): {message}")), }; From b7a7b9df312eff2bb405b6d83620069064e83dd6 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 23:17:17 +0200 Subject: [PATCH 32/52] Add artefact retrieval unit tests --- .../bunkerbox-worker-protocol/src/lib_ut.rs | 44 ++++++ crates/bunkerbox-worker/src/worker_ut.rs | 79 +++++++++- src/artifact_ut.rs | 92 +++++++++++ src/loopback_ut.rs | 43 ++++- src/remote_target_ut.rs | 52 ++++++ src/ssh_ut.rs | 148 ++++++++++++++++-- 6 files changed, 441 insertions(+), 17 deletions(-) create mode 100644 src/artifact_ut.rs diff --git a/crates/bunkerbox-worker-protocol/src/lib_ut.rs b/crates/bunkerbox-worker-protocol/src/lib_ut.rs index 00d1bcd..bff8c92 100644 --- a/crates/bunkerbox-worker-protocol/src/lib_ut.rs +++ b/crates/bunkerbox-worker-protocol/src/lib_ut.rs @@ -297,3 +297,47 @@ fn unknown_nested_kinds_and_flags_are_rejected() { frame[WORKER_FRAME_HEADER_LEN + 34] = 9; assert!(WorkerMessage::decode(&frame).is_err()); } + +#[test] +fn artifact_messages_round_trip_only_in_protocol_v2() { + let (request_id, session_id, _upload_id) = ids(); + let artifact_set_id = WorkerArtifactSetId([4; 16]); + let artifact = WorkerArtifactEntry::new("dist/result", 0o755, 4, [8; 32]).unwrap(); + let artifact_build = build().with_artifacts(vec![WorkerArtifactPath::new("dist/result").unwrap()], 1024, 2048).unwrap(); + let messages = [ + WorkerMessage::Build { request_id, session_id, build: artifact_build }, + WorkerMessage::ArtifactManifest { request_id, session_id, artifact_set_id, entries: vec![artifact], total_bytes: 4 }, + WorkerMessage::FetchArtifact { request_id, session_id, artifact_set_id, entry_index: 0 }, + WorkerMessage::ArtifactChunk { request_id, session_id, artifact_set_id, entry_index: 0, offset: 0, data: b"data".to_vec() }, + WorkerMessage::ArtifactComplete { request_id, session_id, artifact_set_id, entry_index: 0 }, + ]; + + for message in messages { + let frame = message.encode_version(WORKER_ARTIFACT_PROTOCOL_VERSION).unwrap(); + assert_eq!(frame[4..6], WORKER_ARTIFACT_PROTOCOL_VERSION.to_le_bytes()); + let (version, decoded) = WorkerMessage::decode_versioned(&frame).unwrap(); + assert_eq!(version, WORKER_ARTIFACT_PROTOCOL_VERSION); + assert_eq!(decoded, message); + assert!(message.encode().is_err()); + } + + let v1_build = WorkerMessage::Build { request_id, session_id, build: build() }; + let v1 = v1_build.encode().unwrap(); + assert_eq!(WorkerMessage::decode(&v1).unwrap(), v1_build); +} + +#[test] +fn artifact_protocol_rejects_bad_identity_offsets_and_manifest_totals() { + let (request_id, session_id, _) = ids(); + let artifact_set_id = WorkerArtifactSetId([4; 16]); + assert!(WorkerMessage::FetchArtifact { request_id, session_id, artifact_set_id, entry_index: MAX_WORKER_ARTIFACT_ENTRIES as u32 } + .encode_version(WORKER_ARTIFACT_PROTOCOL_VERSION) + .is_err()); + assert!(WorkerMessage::ArtifactChunk { request_id, session_id, artifact_set_id, entry_index: 0, offset: u64::MAX, data: vec![1] } + .encode_version(WORKER_ARTIFACT_PROTOCOL_VERSION) + .is_err()); + let entry = WorkerArtifactEntry::new("result", 0o644, 4, [1; 32]).unwrap(); + assert!(WorkerMessage::ArtifactManifest { request_id, session_id, artifact_set_id, entries: vec![entry], total_bytes: 3 } + .encode_version(WORKER_ARTIFACT_PROTOCOL_VERSION) + .is_err()); +} diff --git a/crates/bunkerbox-worker/src/worker_ut.rs b/crates/bunkerbox-worker/src/worker_ut.rs index 9798d3c..1ebd142 100644 --- a/crates/bunkerbox-worker/src/worker_ut.rs +++ b/crates/bunkerbox-worker/src/worker_ut.rs @@ -1,13 +1,16 @@ use super::*; use crate::platform; use bunkerbox_worker_protocol::{ - WorkerBuild, WorkerEntryKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, WorkerUploadId, + WorkerArtifactPath, WorkerBuild, WorkerEntryKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, + WorkerUploadId, WORKER_ARTIFACT_PROTOCOL_VERSION, }; use sha2::{Digest, Sha256}; use std::fs; use std::io::Cursor; use std::os::unix::fs::PermissionsExt; +use std::os::unix::net::UnixStream; use std::path::Path; +use std::thread; use tempfile::{tempdir, TempDir}; const REQUEST_ID: WorkerRequestId = WorkerRequestId([1; 16]); @@ -183,3 +186,77 @@ fn root_and_protocol_paths_are_confined() { assert!(bunkerbox_worker_protocol::WorkerUploadEntry::file("../escape", 0o644, 1, [0; 32]).is_err()); assert!(bunkerbox_worker_protocol::WorkerUploadEntry::new("src/a", WorkerEntryKind::Directory, 0o755, 0, None).is_ok()); } + +#[test] +fn artifact_capable_build_emits_manifest_fetches_from_spool_and_cleans() { + let fixture = fixture(); + let success_script = fixture._temp.path().join("success.sh"); + fs::write(&success_script, b"#!/bin/sh\nprintf 'data' > result\nexit 0\n").unwrap(); + fs::set_permissions(&success_script, fs::Permissions::from_mode(0o755)).unwrap(); + + let (mut host, worker) = UnixStream::pair().unwrap(); + let worker_input = worker.try_clone().unwrap(); + let service = WorkerService::new(&fixture.root).unwrap(); + let worker_thread = thread::spawn(move || { + let writer = FrameWriter::new(worker); + service.run(worker_input, &writer) + }); + + let request = WorkerRequestId([4; 16]); + write_v2(&mut host, &WorkerMessage::hello_for_version(request, SESSION_ID, false, WORKER_ARTIFACT_PROTOCOL_VERSION)); + write_v2( + &mut host, + &WorkerMessage::UploadBegin { + request_id: request, + session_id: SESSION_ID, + upload_id: UPLOAD_ID, + entries: vec![WorkerUploadEntry::directory("src", 0o755).unwrap(), file_entry("src/input", b"input")], + }, + ); + write_v2( + &mut host, + &WorkerMessage::UploadFileChunk { + request_id: request, + session_id: SESSION_ID, + upload_id: UPLOAD_ID, + path: bunkerbox_worker_protocol::WorkerRelativePath::new("src/input").unwrap(), + offset: 0, + data: b"input".to_vec(), + }, + ); + write_v2(&mut host, &WorkerMessage::UploadComplete { request_id: request, session_id: SESSION_ID, upload_id: UPLOAD_ID }); + let build = WorkerBuild::new("tool", success_script.to_string_lossy().into_owned(), Vec::new(), "src", Vec::new(), Vec::new(), UPLOAD_ID) + .unwrap() + .with_artifacts(vec![WorkerArtifactPath::new("src/result").unwrap()], 1024, 2048) + .unwrap(); + write_v2(&mut host, &WorkerMessage::Build { request_id: request, session_id: SESSION_ID, build }); + + let artifact_set_id = loop { + let (_, message) = WorkerMessage::read_blocking_versioned(&mut host).unwrap(); + match message { + WorkerMessage::ArtifactManifest { artifact_set_id: id, entries, total_bytes, .. } => { + assert_eq!(entries.len(), 1); + assert_eq!(entries[0].path().as_str(), "src/result"); + assert_eq!(entries[0].size(), 4); + assert_eq!(total_bytes, 4); + break id; + } + WorkerMessage::Hello { .. } | WorkerMessage::UploadComplete { .. } | WorkerMessage::Completed { .. } => {} + other => panic!("unexpected worker message: {other:?}"), + } + }; + write_v2(&mut host, &WorkerMessage::FetchArtifact { request_id: request, session_id: SESSION_ID, artifact_set_id, entry_index: 0 }); + let (_, chunk) = WorkerMessage::read_blocking_versioned(&mut host).unwrap(); + assert!(matches!(chunk, WorkerMessage::ArtifactChunk { offset: 0, data, .. } if data == b"data")); + let (_, complete) = WorkerMessage::read_blocking_versioned(&mut host).unwrap(); + assert!(matches!(complete, WorkerMessage::ArtifactComplete { entry_index: 0, .. })); + write_v2(&mut host, &WorkerMessage::Cleanup { request_id: request, session_id: SESSION_ID, upload_token: UPLOAD_ID }); + let (_, cleanup) = WorkerMessage::read_blocking_versioned(&mut host).unwrap(); + assert!(matches!(cleanup, WorkerMessage::Completed { operation: WorkerOperation::Cleanup, exit_code: 0, .. })); + drop(host); + assert_eq!(worker_thread.join().unwrap(), Ok(())); +} + +fn write_v2(stream: &mut UnixStream, message: &WorkerMessage) { + message.write_blocking_version(stream, WORKER_ARTIFACT_PROTOCOL_VERSION).unwrap(); +} diff --git a/src/artifact_ut.rs b/src/artifact_ut.rs new file mode 100644 index 0000000..33d93c7 --- /dev/null +++ b/src/artifact_ut.rs @@ -0,0 +1,92 @@ +use super::*; +use std::fs; +use std::io::Read; +use tempfile::tempdir; + +fn limits() -> ArtifactLimits { + ArtifactLimits::new(Duration::from_secs(1), 4, 1024, 2048).unwrap() +} + +#[test] +fn policy_rejects_unsafe_and_duplicate_paths() { + for paths in [ + vec!["/absolute/file".to_string()], + vec!["../escape".to_string()], + vec!["dir/./file".to_string()], + vec!["dir/*".to_string()], + vec!["same".to_string(), "same".to_string()], + ] { + assert!(ArtifactPolicy::new(paths).is_err()); + } +} + +#[test] +fn local_spool_is_manifested_and_published_without_buffering_the_file() { + let temp = tempdir().unwrap(); + let job = temp.path().join("job"); + let workspace = temp.path().join("workspace"); + let jobs = temp.path().join("jobs"); + fs::create_dir_all(job.join("out")).unwrap(); + fs::create_dir_all(&workspace).unwrap(); + fs::create_dir(&jobs).unwrap(); + fs::write(job.join("out/result.bin"), b"artifact bytes").unwrap(); + + let policy = ArtifactPolicy::new(vec!["out/result.bin".to_string()]).unwrap(); + let spool = LocalArtifactSpool::capture(&job, &jobs, &policy, limits()).unwrap(); + let manifest = spool.manifest().clone(); + assert_eq!(manifest.total_bytes(), 14); + assert_eq!(manifest.entries()[0].path(), "out/result.bin"); + + let mut publication = ArtifactPublication::new(&workspace, [1; 16], manifest).unwrap(); + let mut source = spool.open_entry(0).unwrap(); + publication.copy_from_reader(0, &mut source).unwrap(); + publication.publish().unwrap(); + + let mut actual = Vec::new(); + fs::File::open(workspace.join(".bunkerbox/artifacts/01010101010101010101010101010101/out/result.bin")).unwrap().read_to_end(&mut actual).unwrap(); + assert_eq!(actual, b"artifact bytes"); +} + +#[test] +fn publication_rejects_collisions_and_cleans_partial_staging() { + let temp = tempdir().unwrap(); + let workspace = temp.path().join("workspace"); + fs::create_dir(&workspace).unwrap(); + let policy = ArtifactPolicy::new(vec!["result".to_string()]).unwrap(); + let entry = ArtifactEntry::new("result", 0o644, 1, sha256(b"x")).unwrap(); + let manifest = ArtifactManifest::new(vec![entry], 1, &policy, limits()).unwrap(); + + let publication = ArtifactPublication::new(&workspace, [2; 16], manifest.clone()).unwrap(); + drop(publication); + assert!(!workspace.join(".bunkerbox/artifacts/.staging/02020202020202020202020202020202").exists()); + + let mut publication = ArtifactPublication::new(&workspace, [2; 16], manifest).unwrap(); + let mut writer = publication.begin(0).unwrap(); + writer.write_chunk(b"x").unwrap(); + publication.complete(writer).unwrap(); + publication.publish().unwrap(); + assert!(ArtifactPublication::new( + &workspace, + [2; 16], + ArtifactManifest::new(vec![ArtifactEntry::new("result", 0o644, 1, sha256(b"x")).unwrap()], 1, &policy, limits(),).unwrap(), + ) + .is_err()); +} + +#[cfg(unix)] +#[test] +fn local_capture_rejects_symlink_outputs() { + let temp = tempdir().unwrap(); + let job = temp.path().join("job"); + let jobs = temp.path().join("jobs"); + fs::create_dir(&job).unwrap(); + fs::create_dir(&jobs).unwrap(); + std::os::unix::fs::symlink("/etc/passwd", job.join("result")).unwrap(); + let policy = ArtifactPolicy::new(vec!["result".to_string()]).unwrap(); + assert!(LocalArtifactSpool::capture(&job, &jobs, &policy, limits()).is_err()); +} + +fn sha256(value: &[u8]) -> [u8; 32] { + use sha2::{Digest, Sha256}; + Sha256::digest(value).into() +} diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 60f8fc2..8f71b52 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -1,7 +1,8 @@ use super::*; +use crate::artifact::{ArtifactLimits, ArtifactPolicy}; use crate::remote::{ - RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTool, RequestId, - WorkspaceRelativePath, + RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteFailureClass, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, + RemoteTool, RequestId, WorkspaceRelativePath, }; use tempfile::TempDir; @@ -472,3 +473,41 @@ async fn output_limit_kills_a_flooding_direct_child() { let request = authorized_build(target, session_id, "printf", vec!["0123456789".into()], Vec::new(), snapshot_id); assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::OutputLimit { limit: 8 })); } + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn loopback_retrieves_only_trusted_declared_artifacts_after_successful_build() { + let (_temp, session, target, session_id) = fixture(); + fs::write(session.workspace_root().join("src/Makefile"), ".PHONY: artifact\nartifact:\n\t@printf 'data' > result\n").unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["make".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session.clone(), tools) + .with_artifacts(ArtifactPolicy::new(vec!["src/result".to_string()]).unwrap(), ArtifactLimits::default()); + let (events, receiver) = mpsc::channel(16); + let request = authorized_build(target, session_id, "make", vec!["artifact".into()], Vec::new(), snapshot_id); + assert_eq!(backend.execute(request, events).await, Ok(())); + let events = collect_events(receiver).await; + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); + assert_eq!(fs::read(session.workspace_root().join(".bunkerbox/artifacts/03030303030303030303030303030303/src/result")).unwrap(), b"data"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn loopback_missing_required_artifact_is_a_terminal_artifact_failure() { + let (_temp, session, target, session_id) = fixture(); + fs::write(session.workspace_root().join("src/Makefile"), ".PHONY: ok\nok:\n\t@true\n").unwrap(); + let snapshot_id = session.sync_snapshot().unwrap(); + let tools = resolve_fixed_tools(["make".to_string()]); + if tools.is_empty() { + return; + } + let backend = LoopbackBackend::new(session.clone(), tools) + .with_artifacts(ArtifactPolicy::new(vec!["src/missing".to_string()]).unwrap(), ArtifactLimits::default()); + let (events, receiver) = mpsc::channel(16); + let request = authorized_build(target, session_id, "make", vec!["ok".into()], Vec::new(), snapshot_id); + assert!(matches!(backend.execute(request, events).await, Err(RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, .. }))); + let events = collect_events(receiver).await; + assert!(!events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { .. }))); + assert!(fs::read_dir(&session.jobs_root).unwrap().next().is_none()); +} diff --git a/src/remote_target_ut.rs b/src/remote_target_ut.rs index 1753ffc..1a666c8 100644 --- a/src/remote_target_ut.rs +++ b/src/remote_target_ut.rs @@ -319,3 +319,55 @@ fn config_version_and_transport_are_strict() { write_config(&fixture, &transport); assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); } + +#[test] +fn project_artifact_policy_and_limits_are_loaded_from_trusted_binding() { + let fixture = Fixture::new(); + let marker = format!(" {}:\n backend: ssh\n target: ssh-one\n", scalar(fixture.project.to_str().unwrap())); + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")) + .replace( + " max-output-bytes: 67108864\n", + " max-output-bytes: 67108864\n artifact-timeout-seconds: 7\n max-artifact-bytes: 1M\n max-artifact-total-bytes: 2M\n max-artifact-entries: 3\n", + ) + .replace( + &marker, + &format!( + "{marker} artifacts:\n paths:\n - target/result\n - dist/package.tar.gz\n" + ), + ); + write_config(&fixture, &yaml); + + let config = RemoteTargetConfig::load_from(&fixture.config).unwrap(); + let resolved = config.resolve_for_project(&fixture.project).unwrap(); + assert_eq!(resolved.artifacts().paths(), ["target/result", "dist/package.tar.gz"]); + let limits = resolved.target().unwrap().resources().artifact_limits(); + assert_eq!(limits.timeout, Duration::from_secs(7)); + assert_eq!(limits.max_entries, 3); + assert_eq!(limits.max_file_bytes, 1024 * 1024); + assert_eq!(limits.max_total_bytes, 2 * 1024 * 1024); +} + +#[test] +fn project_artifact_policy_must_fit_target_limits() { + let fixture = Fixture::new(); + let marker = format!(" {}:\n backend: ssh\n target: ssh-one\n", scalar(fixture.project.to_str().unwrap())); + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")) + .replace(" max-output-bytes: 67108864\n", " max-output-bytes: 67108864\n max-artifact-entries: 1\n") + .replace(&marker, &format!("{marker} artifacts:\n paths:\n - target/result\n - dist/package.tar.gz\n")); + write_config(&fixture, &yaml); + + let error = RemoteTargetConfig::load_from(&fixture.config).unwrap_err(); + assert!(error.contains("artifact policy exceeds configured entry count 1")); +} + +#[test] +fn project_artifact_paths_reject_traversal_globs_and_duplicates() { + let fixture = Fixture::new(); + let marker = format!(" {}:\n backend: loopback\n", scalar(fixture.project.to_str().unwrap())); + for paths in [" - /absolute\n", " - ../escape\n", " - dist/*\n", " - result\n - result\n"] { + let yaml = + valid_yaml(&fixture, &fixture.project, "loopback", None).replace(&marker, &format!("{marker} artifacts:\n paths:\n{paths}")); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + } +} diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index e9bf160..8d4f3aa 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -1,11 +1,16 @@ use super::*; +use crate::artifact::{ArtifactLimits, ArtifactPolicy}; use crate::remote::{ RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteSnapshotId, RemoteTool, RequestId, WorkspaceRelativePath, WorkspaceSessionId, }; use crate::remote_target::RemoteTargetConfig; use crate::snapshot::{SnapshotBuilder, SnapshotExclusionPolicy, SnapshotLimits, SnapshotStore}; -use crate::worker_protocol::{self, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerSessionId, WorkerUploadEntry, WorkerUploadId}; +use crate::worker_protocol::{ + self, WorkerArtifactEntry, WorkerArtifactSetId, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerSessionId, WorkerUploadEntry, + WorkerUploadId, +}; +use sha2::{Digest, Sha256}; use std::fs; use std::path::Path; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; @@ -22,6 +27,7 @@ enum ScriptMode { Success, WrongVersion, ProtocolError, + ArtifactSuccess, } struct ScriptedFactory { @@ -106,7 +112,7 @@ where R: AsyncRead + Unpin, W: AsyncWrite + Unpin, { - let Ok(WorkerMessage::Hello { request_id, session_id, .. }) = worker_protocol::read_message(&mut reader).await else { + let Ok((hello_version, WorkerMessage::Hello { request_id, session_id, .. })) = worker_protocol::read_message_versioned(&mut reader).await else { return 71; }; if matches!(mode, ScriptMode::WrongVersion) { @@ -115,9 +121,15 @@ where writer.write_all(&frame).await.unwrap(); return 0; } - worker_protocol::write_message(&mut writer, &WorkerMessage::hello(request_id, session_id, true)).await.unwrap(); + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::hello_for_version(request_id, session_id, true, hello_version), + hello_version, + ) + .await + .unwrap(); - let Ok(first) = worker_protocol::read_message(&mut reader).await else { return 72 }; + let Ok((_, first)) = worker_protocol::read_message_versioned(&mut reader).await else { return 72 }; messages.lock().unwrap().push(first.clone()); if matches!(mode, ScriptMode::ProtocolError) { worker_protocol::write_message( @@ -131,7 +143,7 @@ where } match first { WorkerMessage::UploadBegin { request_id, session_id, upload_id, .. } => loop { - let Ok(message) = worker_protocol::read_message(&mut reader).await else { return 73 }; + let Ok((_, message)) = worker_protocol::read_message_versioned(&mut reader).await else { return 73 }; messages.lock().unwrap().push(message.clone()); match message { WorkerMessage::UploadFileChunk { .. } => {} @@ -139,13 +151,20 @@ where if received_request != request_id || received_session != session_id || received_upload != upload_id { return 74; } - worker_protocol::write_message( + worker_protocol::write_message_versioned( &mut writer, &WorkerMessage::SyncProgress { request_id, session_id, upload_id, completed_bytes: 1, total_bytes: Some(1) }, + hello_version, + ) + .await + .unwrap(); + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::UploadComplete { request_id, session_id, upload_id }, + hello_version, ) .await .unwrap(); - worker_protocol::write_message(&mut writer, &WorkerMessage::UploadComplete { request_id, session_id, upload_id }).await.unwrap(); return 0; } _ => return 75, @@ -153,18 +172,84 @@ where }, WorkerMessage::Build { request_id, session_id, build } => { messages.lock().unwrap().push(WorkerMessage::Build { request_id, session_id, build: build.clone() }); - worker_protocol::write_message(&mut writer, &WorkerMessage::stdout(request_id, session_id, b"remote stdout\n".to_vec())).await.unwrap(); - worker_protocol::write_message(&mut writer, &WorkerMessage::stderr(request_id, session_id, b"remote stderr\n".to_vec())).await.unwrap(); - worker_protocol::write_message(&mut writer, &WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, 7)).await.unwrap(); - let Ok(cleanup) = worker_protocol::read_message(&mut reader).await else { return 76 }; + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::stdout(request_id, session_id, b"remote stdout\n".to_vec()), + hello_version, + ) + .await + .unwrap(); + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::stderr(request_id, session_id, b"remote stderr\n".to_vec()), + hello_version, + ) + .await + .unwrap(); + let exit_code = if matches!(mode, ScriptMode::ArtifactSuccess) { 0 } else { 7 }; + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, exit_code), + hello_version, + ) + .await + .unwrap(); + if matches!(mode, ScriptMode::ArtifactSuccess) { + let digest: [u8; 32] = Sha256::digest(b"data").into(); + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::ArtifactManifest { + request_id, + session_id, + artifact_set_id: WorkerArtifactSetId([9; 16]), + entries: vec![WorkerArtifactEntry::new("result", 0o644, 4, digest).unwrap()], + total_bytes: 4, + }, + hello_version, + ) + .await + .unwrap(); + let Ok((_, fetch)) = worker_protocol::read_message_versioned(&mut reader).await else { return 76 }; + messages.lock().unwrap().push(fetch.clone()); + if !matches!(fetch, WorkerMessage::FetchArtifact { artifact_set_id, entry_index: 0, .. } if artifact_set_id == WorkerArtifactSetId([9; 16])) + { + return 77; + } + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::ArtifactChunk { + request_id, + session_id, + artifact_set_id: WorkerArtifactSetId([9; 16]), + entry_index: 0, + offset: 0, + data: b"data".to_vec(), + }, + hello_version, + ) + .await + .unwrap(); + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::ArtifactComplete { request_id, session_id, artifact_set_id: WorkerArtifactSetId([9; 16]), entry_index: 0 }, + hello_version, + ) + .await + .unwrap(); + } + let Ok((_, cleanup)) = worker_protocol::read_message_versioned(&mut reader).await else { return 76 }; messages.lock().unwrap().push(cleanup.clone()); if !matches!(cleanup, WorkerMessage::Cleanup { request_id: received_request, session_id: received_session, upload_token } if received_request == request_id && received_session == session_id && upload_token == build.upload_token()) { return 77; } - worker_protocol::write_message(&mut writer, &WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0)) - .await - .unwrap(); + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::completed(request_id, session_id, WorkerOperation::Cleanup, 0), + hello_version, + ) + .await + .unwrap(); 0 } _ => 78, @@ -343,3 +428,38 @@ async fn worker_protocol_failure_is_terminal_without_local_fallback() { assert!(matches!(error, RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, .. }), "{error:?}"); assert_eq!(fixture.session.snapshot_capability_count(), 0); } + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn artifact_capable_worker_manifest_is_fetched_verified_published_and_cleaned() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::Success, ScriptMode::ArtifactSuccess]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()) + .unwrap() + .with_process_factory(factory.clone()) + .with_artifacts(ArtifactPolicy::new(vec!["result".to_string()]).unwrap(), ArtifactLimits::default()); + + let (sync_tx, sync_rx) = tokio::sync::mpsc::channel(64); + backend.execute(authorize_sync(&fixture), sync_tx).await.unwrap(); + let snapshot_id = collect(sync_rx) + .await + .into_iter() + .find_map(|event| match event { + RemoteBackendEvent::SyncCompleted { snapshot_id } => Some(snapshot_id), + _ => None, + }) + .unwrap(); + + let (build_tx, build_rx) = tokio::sync::mpsc::channel(64); + backend.execute(authorize_build(&fixture, snapshot_id), build_tx).await.unwrap(); + let events = collect(build_rx).await; + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); + let published = fixture.session.workspace_root().join(".bunkerbox/artifacts/04040404040404040404040404040404/result"); + assert_eq!(fs::read(published).unwrap(), b"data"); + + let messages = factory.messages.lock().unwrap(); + assert!(messages + .iter() + .any(|message| matches!(message, WorkerMessage::Build { build, .. } if build.artifact_paths().iter().any(|path| path.as_str() == "result")))); + assert!(messages.iter().any(|message| matches!(message, WorkerMessage::FetchArtifact { entry_index: 0, .. }))); + assert!(messages.iter().any(|message| matches!(message, WorkerMessage::Cleanup { .. }))); +} From ef5aa82242af68915a9b36f85b6f15839a942683 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 23:48:26 +0200 Subject: [PATCH 33/52] Anchor ArtifactPublication workspace publication through opened directory FDs --- src/artifact.rs | 196 +++++++++++++++++++++++++++++++++--------------- src/ssh.rs | 34 +++++++-- 2 files changed, 164 insertions(+), 66 deletions(-) diff --git a/src/artifact.rs b/src/artifact.rs index 591e0df..5170cd3 100644 --- a/src/artifact.rs +++ b/src/artifact.rs @@ -266,8 +266,10 @@ impl Drop for LocalArtifactSpool { } pub struct ArtifactPublication { - staging: PathBuf, - final_path: PathBuf, + artifacts: File, + staging_root: File, + staging: File, + request: String, manifest: ArtifactManifest, written: BTreeSet, published: bool, @@ -275,17 +277,16 @@ pub struct ArtifactPublication { impl ArtifactPublication { pub fn new(workspace_root: &Path, request_id: [u8; 16], manifest: ArtifactManifest) -> Result { - let artifacts = ensure_directory_path(workspace_root, &[ARTIFACT_ROOT, ARTIFACT_DIRECTORY])?; - let staging_root = ensure_directory_path(&artifacts, &[STAGING_DIRECTORY])?; + let workspace = open_directory(workspace_root).map_err(|error| format!("open artifact workspace: {error}"))?; + let bunkerbox = ensure_directory_at(&workspace, ARTIFACT_ROOT)?; + let artifacts = ensure_directory_at(&bunkerbox, ARTIFACT_DIRECTORY)?; + let staging_root = ensure_directory_at(&artifacts, STAGING_DIRECTORY)?; let request = hex_id(request_id); - let staging = staging_root.join(&request); - let final_path = artifacts.join(&request); - if fs::symlink_metadata(&staging).is_ok() || fs::symlink_metadata(&final_path).is_ok() { + if entry_exists_at(&artifacts, &request).map_err(|error| format!("inspect artifact destination: {error}"))? { return Err(format!("artifact destination already exists for request {request}")); } - fs::create_dir(&staging).map_err(|error| format!("create artifact staging directory: {error}"))?; - set_private_mode(&staging)?; - Ok(Self { staging, final_path, manifest, written: BTreeSet::new(), published: false }) + let staging = create_directory_at(&staging_root, &request, 0o700).map_err(|error| format!("create artifact staging directory: {error}"))?; + Ok(Self { artifacts, staging_root, staging, request, manifest, written: BTreeSet::new(), published: false }) } pub fn begin(&self, index: usize) -> Result { @@ -293,16 +294,7 @@ impl ArtifactPublication { return Err(format!("artifact entry was already written: {index}")); } let entry = self.manifest.entry(index).ok_or_else(|| "artifact index is out of range".to_string())?.clone(); - let components = artifact_components(&entry.path)?; - let (name, parents) = components.split_last().ok_or_else(|| "artifact path is empty".to_string())?; - let parent = ensure_directory_components(&self.staging, parents)?; - let path = parent.join(name); - let file = OpenOptions::new() - .write(true) - .create_new(true) - .mode(0o600) - .custom_flags(libc::O_NOFOLLOW | libc::O_CLOEXEC) - .open(&path) + let file = create_relative_file(&self.staging, &entry.path, 0o600) .map_err(|error| format!("create artifact staging file {}: {error}", entry.path))?; Ok(ArtifactWriter { index, entry, file, written: 0, hasher: Sha256::new() }) } @@ -333,12 +325,8 @@ impl ArtifactPublication { if self.written.len() != self.manifest.entries.len() { return Err("artifact publication is missing entries".to_string()); } - if fs::symlink_metadata(&self.final_path).is_ok() { - return Err("artifact destination appeared before publication".to_string()); - } - fs::rename(&self.staging, &self.final_path).map_err(|error| format!("publish artifacts: {error}"))?; - let parent = self.final_path.parent().ok_or_else(|| "artifact destination has no parent".to_string())?; - File::open(parent).and_then(|file| file.sync_all()).map_err(|error| format!("flush artifact destination: {error}"))?; + rename_noreplace(&self.staging_root, &self.request, &self.artifacts, &self.request).map_err(|error| format!("publish artifacts: {error}"))?; + self.artifacts.sync_all().map_err(|error| format!("flush artifact destination: {error}"))?; self.published = true; Ok(()) } @@ -347,7 +335,7 @@ impl ArtifactPublication { impl Drop for ArtifactPublication { fn drop(&mut self) { if !self.published { - let _ = fs::remove_dir_all(&self.staging); + cleanup_publication(&self.staging_root, &self.staging, &self.request, &self.manifest); } } } @@ -434,44 +422,122 @@ fn create_unique_directory(parent: &Path, label: &str) -> Result Result { - let mut current = root.to_path_buf(); - for component in components { - if component.is_empty() || *component == "." || *component == ".." || component.contains('/') { - return Err("artifact destination contains an invalid component".to_string()); - } - current.push(component); - match fs::symlink_metadata(¤t) { - Ok(metadata) if metadata.file_type().is_dir() => {} - Ok(_) => return Err(format!("artifact destination component is not a directory: {}", current.display())), - Err(error) if error.kind() == io::ErrorKind::NotFound => { - fs::create_dir(¤t).map_err(|create_error| format!("create artifact destination directory: {create_error}"))?; - set_private_mode(¤t)?; - } - Err(error) => return Err(format!("inspect artifact destination: {error}")), +fn ensure_directory_at(parent: &File, name: &str) -> Result { + let name = component_name(name).map_err(|error| format!("invalid artifact directory component: {error}"))?; + let created = unsafe { libc::mkdirat(parent.as_raw_fd(), name.as_ptr(), 0o700) } == 0; + if !created { + let error = io::Error::last_os_error(); + if error.kind() != io::ErrorKind::AlreadyExists { + return Err(format!("create artifact directory: {error}")); } } - Ok(current) + let name = name.to_str().map_err(|_| "artifact directory component is not UTF-8".to_string())?; + let directory = open_directory_at(parent, name).map_err(|error| format!("open artifact directory: {error}"))?; + if created { + set_private_mode_fd(&directory)?; + } + Ok(directory) } -fn ensure_directory_components(root: &Path, components: &[&str]) -> Result { - let mut current = root.to_path_buf(); - for component in components { - if component.is_empty() || *component == "." || *component == ".." || component.contains('/') { - return Err("artifact path contains an invalid component".to_string()); - } - current.push(component); - match fs::symlink_metadata(¤t) { - Ok(metadata) if metadata.file_type().is_dir() => {} - Ok(_) => return Err(format!("artifact parent is not a directory: {}", current.display())), - Err(error) if error.kind() == io::ErrorKind::NotFound => { - fs::create_dir(¤t).map_err(|create_error| format!("create artifact parent: {create_error}"))?; - set_private_mode(¤t)?; - } - Err(error) => return Err(format!("inspect artifact parent: {error}")), +fn create_directory_at(parent: &File, name: &str, mode: u32) -> io::Result { + let name = component_name(name)?; + let result = unsafe { libc::mkdirat(parent.as_raw_fd(), name.as_ptr(), mode as libc::mode_t) }; + if result != 0 { + return Err(io::Error::last_os_error()); + } + let name = name.to_str().map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "artifact directory component is not UTF-8"))?; + let directory = match open_directory_at(parent, name) { + Ok(directory) => directory, + Err(error) => { + let _ = unlink_at(parent, name, libc::AT_REMOVEDIR); + return Err(error); + } + }; + if let Err(error) = set_private_mode_fd_io(&directory, mode) { + let _ = unlink_at(parent, name, libc::AT_REMOVEDIR); + return Err(error); + } + Ok(directory) +} + +fn entry_exists_at(parent: &File, name: &str) -> io::Result { + let name = component_name(name)?; + let mut metadata = unsafe { std::mem::zeroed::() }; + let result = unsafe { libc::fstatat(parent.as_raw_fd(), name.as_ptr(), &mut metadata, libc::AT_SYMLINK_NOFOLLOW) }; + if result == 0 { + return Ok(true); + } + let error = io::Error::last_os_error(); + if error.kind() == io::ErrorKind::NotFound { + Ok(false) + } else { + Err(error) + } +} + +fn rename_noreplace(from_parent: &File, from_name: &str, to_parent: &File, to_name: &str) -> io::Result<()> { + let from_name = component_name(from_name)?; + let to_name = component_name(to_name)?; + #[cfg(target_os = "linux")] + { + let result = + unsafe { libc::renameat2(from_parent.as_raw_fd(), from_name.as_ptr(), to_parent.as_raw_fd(), to_name.as_ptr(), libc::RENAME_NOREPLACE) }; + if result == 0 { + Ok(()) + } else { + Err(io::Error::last_os_error()) } } - Ok(current) + #[cfg(not(target_os = "linux"))] + { + let _ = (from_parent, from_name, to_parent, to_name); + Err(io::Error::new(io::ErrorKind::Unsupported, "artifact publication requires renameat2")) + } +} + +fn cleanup_publication(staging_root: &File, staging: &File, request: &str, manifest: &ArtifactManifest) { + let mut directories = Vec::new(); + for entry in &manifest.entries { + let Ok(components) = artifact_components(&entry.path) else { continue }; + let _ = unlink_relative(staging, &components, 0); + for count in 1..components.len() { + directories.push(components[..count].join("/")); + } + } + directories.sort_by_key(|path| std::cmp::Reverse(path.split('/').count())); + directories.dedup(); + for directory in directories { + if let Ok(components) = artifact_components(&directory) { + let _ = unlink_relative(staging, &components, libc::AT_REMOVEDIR); + } + } + let _ = unlink_at(staging_root, request, libc::AT_REMOVEDIR); +} + +fn unlink_relative(root: &File, components: &[&str], flags: i32) -> io::Result<()> { + let (name, parents) = components.split_last().ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "artifact path is empty"))?; + let mut parent = root.try_clone()?; + for component in parents { + parent = open_directory_at(&parent, component)?; + } + unlink_at(&parent, name, flags) +} + +fn unlink_at(parent: &File, name: &str, flags: i32) -> io::Result<()> { + let name = component_name(name)?; + let result = unsafe { libc::unlinkat(parent.as_raw_fd(), name.as_ptr(), flags) }; + if result == 0 { + Ok(()) + } else { + Err(io::Error::last_os_error()) + } +} + +fn component_name(name: &str) -> io::Result { + if name.is_empty() || name == "." || name == ".." || name.contains('/') { + return Err(io::Error::new(io::ErrorKind::InvalidInput, "invalid path component")); + } + CString::new(name.as_bytes()).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "NUL in path component")) } fn open_directory(path: &Path) -> io::Result { @@ -594,6 +660,18 @@ fn set_private_mode(path: &Path) -> Result<(), String> { fs::set_permissions(path, fs::Permissions::from_mode(0o700)).map_err(|error| format!("set private artifact mode: {error}")) } +fn set_private_mode_fd(file: &File) -> Result<(), String> { + set_private_mode_fd_io(file, 0o700).map_err(|error| format!("set private artifact mode: {error}")) +} + +fn set_private_mode_fd_io(file: &File, mode: u32) -> io::Result<()> { + if unsafe { libc::fchmod(file.as_raw_fd(), mode as libc::mode_t) } != 0 { + Err(io::Error::last_os_error()) + } else { + Ok(()) + } +} + fn hex_id(bytes: [u8; 16]) -> String { bytes.iter().map(|byte| format!("{byte:02x}")).collect() } diff --git a/src/ssh.rs b/src/ssh.rs index b3d1562..45ca274 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -690,7 +690,7 @@ async fn build_and_finish(connection: &mut WorkerConnection, plan: BuildPlan<'_> check_correlation(received_request, received_session, request_id, session_id, "worker artifact manifest")?; (artifact_set_id, artifact_manifest_from_worker(entries, total_bytes, &artifact_policy, artifact_limits)?) } - WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Artifact, kind, message)), + WorkerMessage::Error { kind, message, .. } => return Err(worker_artifact_manifest_error(kind, message)), _ => return Err(worker_protocol("unexpected worker artifact manifest response")), }; let retrieval = timeout( @@ -730,12 +730,13 @@ async fn fetch_and_publish_artifacts( message: "artifact index does not fit worker protocol".to_string(), })?, }) - .await?; + .await + .map_err(artifact_transfer_error)?; let mut writer = publication .begin(index) .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; loop { - match connection.read().await? { + match connection.read().await.map_err(artifact_transfer_error)? { WorkerMessage::ArtifactChunk { request_id: received_request, session_id: received_session, @@ -744,7 +745,8 @@ async fn fetch_and_publish_artifacts( offset, data, } => { - check_correlation(received_request, received_session, request_id, session_id, "worker artifact chunk")?; + check_correlation(received_request, received_session, request_id, session_id, "worker artifact chunk") + .map_err(artifact_transfer_error)?; if received_set != artifact_set_id || entry_index != index as u32 { return Err(RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, @@ -770,7 +772,8 @@ async fn fetch_and_publish_artifacts( artifact_set_id: received_set, entry_index, } => { - check_correlation(received_request, received_session, request_id, session_id, "worker artifact completion")?; + check_correlation(received_request, received_session, request_id, session_id, "worker artifact completion") + .map_err(artifact_transfer_error)?; if received_set != artifact_set_id || entry_index != index as u32 { return Err(RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, @@ -782,8 +785,10 @@ async fn fetch_and_publish_artifacts( .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; break; } - WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Artifact, kind, message)), - _ => return Err(worker_protocol("unexpected worker artifact transfer response")), + WorkerMessage::Error { kind, message, .. } => { + return Err(artifact_transfer_error(worker_error(WorkerOperation::Artifact, kind, message))); + } + _ => return Err(artifact_transfer_error(worker_protocol("unexpected worker artifact transfer response"))), } } } @@ -893,6 +898,21 @@ fn worker_error(operation: WorkerOperation, kind: WorkerErrorKind, message: Stri RemoteBackendError::Transport { class, message } } +fn worker_artifact_manifest_error(kind: WorkerErrorKind, message: String) -> RemoteBackendError { + if matches!(kind, WorkerErrorKind::Artifact) { + RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message } + } else { + worker_error(WorkerOperation::Artifact, kind, message) + } +} + +fn artifact_transfer_error(error: RemoteBackendError) -> RemoteBackendError { + match error { + RemoteBackendError::Transport { message, .. } => RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message }, + other => other, + } +} + fn worker_protocol(message: impl Into) -> RemoteBackendError { RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, message: message.into() } } From 7fee20dd479168639ebb68b07dedd6afaad84f05 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Mon, 17 Aug 2026 23:48:39 +0200 Subject: [PATCH 34/52] Add unit tests for artifacts anchoring --- src/artifact_ut.rs | 27 ++++++++++++++++ src/ssh_ut.rs | 77 ++++++++++++++++++++++++++++++++++++++++++++-- 2 files changed, 101 insertions(+), 3 deletions(-) diff --git a/src/artifact_ut.rs b/src/artifact_ut.rs index 33d93c7..bf7d793 100644 --- a/src/artifact_ut.rs +++ b/src/artifact_ut.rs @@ -73,6 +73,33 @@ fn publication_rejects_collisions_and_cleans_partial_staging() { .is_err()); } +#[cfg(unix)] +#[test] +fn publication_stays_anchored_when_workspace_component_is_swapped() { + let temp = tempdir().unwrap(); + let workspace = temp.path().join("workspace"); + let outside = temp.path().join("outside"); + let original_bunkerbox = temp.path().join("original-bunkerbox"); + fs::create_dir(&workspace).unwrap(); + fs::create_dir(&outside).unwrap(); + let policy = ArtifactPolicy::new(vec!["nested/result".to_string()]).unwrap(); + let entry = ArtifactEntry::new("nested/result", 0o644, 1, sha256(b"x")).unwrap(); + let manifest = ArtifactManifest::new(vec![entry], 1, &policy, limits()).unwrap(); + + let mut publication = ArtifactPublication::new(&workspace, [3; 16], manifest).unwrap(); + fs::rename(workspace.join(ARTIFACT_ROOT), &original_bunkerbox).unwrap(); + std::os::unix::fs::symlink(&outside, workspace.join(ARTIFACT_ROOT)).unwrap(); + + let mut writer = publication.begin(0).unwrap(); + writer.write_chunk(b"x").unwrap(); + publication.complete(writer).unwrap(); + publication.publish().unwrap(); + + let request = "03030303030303030303030303030303"; + assert_eq!(fs::read(original_bunkerbox.join(format!("artifacts/{request}/nested/result"))).unwrap(), b"x"); + assert!(!outside.join(format!("artifacts/{request}")).exists()); +} + #[cfg(unix)] #[test] fn local_capture_rejects_symlink_outputs() { diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index 8d4f3aa..42c00d3 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -28,6 +28,8 @@ enum ScriptMode { WrongVersion, ProtocolError, ArtifactSuccess, + ArtifactManifestError, + ArtifactTransferError, } struct ScriptedFactory { @@ -186,7 +188,11 @@ where ) .await .unwrap(); - let exit_code = if matches!(mode, ScriptMode::ArtifactSuccess) { 0 } else { 7 }; + let exit_code = if matches!(mode, ScriptMode::ArtifactSuccess | ScriptMode::ArtifactManifestError | ScriptMode::ArtifactTransferError) { + 0 + } else { + 7 + }; worker_protocol::write_message_versioned( &mut writer, &WorkerMessage::completed(request_id, session_id, WorkerOperation::Build, exit_code), @@ -194,7 +200,21 @@ where ) .await .unwrap(); - if matches!(mode, ScriptMode::ArtifactSuccess) { + if matches!(mode, ScriptMode::ArtifactManifestError) { + worker_protocol::write_message_versioned( + &mut writer, + &WorkerMessage::error( + request_id, + session_id, + WorkerOperation::Artifact, + WorkerErrorKind::Artifact, + "configured artifact is missing", + ), + hello_version, + ) + .await + .unwrap(); + } else if matches!(mode, ScriptMode::ArtifactSuccess | ScriptMode::ArtifactTransferError) { let digest: [u8; 32] = Sha256::digest(b"data").into(); worker_protocol::write_message_versioned( &mut writer, @@ -215,6 +235,7 @@ where { return 77; } + let data = if matches!(mode, ScriptMode::ArtifactTransferError) { b"da".to_vec() } else { b"data".to_vec() }; worker_protocol::write_message_versioned( &mut writer, &WorkerMessage::ArtifactChunk { @@ -223,7 +244,7 @@ where artifact_set_id: WorkerArtifactSetId([9; 16]), entry_index: 0, offset: 0, - data: b"data".to_vec(), + data, }, hello_version, ) @@ -463,3 +484,53 @@ async fn artifact_capable_worker_manifest_is_fetched_verified_published_and_clea assert!(messages.iter().any(|message| matches!(message, WorkerMessage::FetchArtifact { entry_index: 0, .. }))); assert!(messages.iter().any(|message| matches!(message, WorkerMessage::Cleanup { .. }))); } + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn worker_artifact_capture_failure_is_typed_as_manifest_failure() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::Success, ScriptMode::ArtifactManifestError]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()) + .unwrap() + .with_process_factory(factory) + .with_artifacts(ArtifactPolicy::new(vec!["result".to_string()]).unwrap(), ArtifactLimits::default()); + + let (sync_tx, sync_rx) = tokio::sync::mpsc::channel(64); + backend.execute(authorize_sync(&fixture), sync_tx).await.unwrap(); + let snapshot_id = collect(sync_rx) + .await + .into_iter() + .find_map(|event| match event { + RemoteBackendEvent::SyncCompleted { snapshot_id } => Some(snapshot_id), + _ => None, + }) + .unwrap(); + + let (build_tx, _build_rx) = tokio::sync::mpsc::channel(64); + let error = backend.execute(authorize_build(&fixture, snapshot_id), build_tx).await.unwrap_err(); + assert!(matches!(error, RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, .. })); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn worker_artifact_fetch_failure_is_typed_as_transfer_failure() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::Success, ScriptMode::ArtifactTransferError]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()) + .unwrap() + .with_process_factory(factory) + .with_artifacts(ArtifactPolicy::new(vec!["result".to_string()]).unwrap(), ArtifactLimits::default()); + + let (sync_tx, sync_rx) = tokio::sync::mpsc::channel(64); + backend.execute(authorize_sync(&fixture), sync_tx).await.unwrap(); + let snapshot_id = collect(sync_rx) + .await + .into_iter() + .find_map(|event| match event { + RemoteBackendEvent::SyncCompleted { snapshot_id } => Some(snapshot_id), + _ => None, + }) + .unwrap(); + + let (build_tx, _build_rx) = tokio::sync::mpsc::channel(64); + let error = backend.execute(authorize_build(&fixture, snapshot_id), build_tx).await.unwrap_err(); + assert!(matches!(error, RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, .. })); +} From b41d9fcd18bbaae7529dde48127ea0ca8cceb921 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 01:56:22 +0200 Subject: [PATCH 35/52] Refactor remote worker --- crates/bunkerbox-worker/src/main.rs | 74 +++- crates/bunkerbox-worker/src/main_ut.rs | 40 ++ crates/bunkerbox-worker/src/platform.rs | 1 + crates/bunkerbox-worker/src/process.rs | 36 +- crates/bunkerbox-worker/src/process_ut.rs | 31 ++ crates/bunkerbox-worker/src/storage.rs | 493 +++++++++++++++++++--- crates/bunkerbox-worker/src/storage_ut.rs | 31 ++ crates/bunkerbox-worker/src/worker.rs | 68 ++- 8 files changed, 694 insertions(+), 80 deletions(-) create mode 100644 crates/bunkerbox-worker/src/main_ut.rs diff --git a/crates/bunkerbox-worker/src/main.rs b/crates/bunkerbox-worker/src/main.rs index 030ba64..8cb112d 100644 --- a/crates/bunkerbox-worker/src/main.rs +++ b/crates/bunkerbox-worker/src/main.rs @@ -3,12 +3,23 @@ mod process; mod storage; mod worker; +use std::collections::BTreeSet; use std::path::PathBuf; +use std::time::Duration; + +use storage::WorkerStateLimits; + +struct WorkerConfig { + root: PathBuf, + limits: WorkerStateLimits, + build_timeout: Duration, + max_output_bytes: u64, +} fn main() { match parse_args(std::env::args().skip(1).collect()) { - Ok(root) => { - if let Err(error) = worker::run_stdio(&root) { + Ok(config) => { + if let Err(error) = worker::run_stdio_with_config(&config.root, config.limits, config.build_timeout, config.max_output_bytes) { write_diagnostic(&error); std::process::exit(70); } @@ -20,11 +31,18 @@ fn main() { } } -fn parse_args(args: Vec) -> Result { +fn parse_args(args: Vec) -> Result { let mut stdio = false; let mut root = None; + let mut limits = WorkerStateLimits::default(); + let mut build_timeout = Duration::from_secs(30); + let mut max_output_bytes = 64 * 1024 * 1024; + let mut seen = BTreeSet::new(); let mut index = 0; while index < args.len() { + if !seen.insert(args[index].clone()) { + return Err(format!("duplicate {}", args[index])); + } match args[index].as_str() { "--stdio" => { if stdio { @@ -44,13 +62,57 @@ fn parse_args(args: Vec) -> Result { root = Some(PathBuf::from(value)); index += 2; } + "--build-timeout-ms" => { + build_timeout = Duration::from_millis(parse_value(&args, &mut index, "--build-timeout-ms")?); + if build_timeout.is_zero() { + return Err("--build-timeout-ms must be positive".to_string()); + } + } + "--max-output-bytes" => { + max_output_bytes = parse_value(&args, &mut index, "--max-output-bytes")?; + if max_output_bytes == 0 || max_output_bytes > process::WORKER_MAX_OUTPUT_BYTES { + return Err(format!("--max-output-bytes must be between 1 and {}", process::WORKER_MAX_OUTPUT_BYTES)); + } + } + "--max-worker-uploads" => limits.max_uploads = parse_count(&args, &mut index, "--max-worker-uploads")?, + "--max-worker-upload-bytes" => limits.max_upload_bytes = parse_value(&args, &mut index, "--max-worker-upload-bytes")?, + "--max-worker-jobs" => limits.max_jobs = parse_count(&args, &mut index, "--max-worker-jobs")?, + "--max-worker-job-bytes" => limits.max_job_bytes = parse_value(&args, &mut index, "--max-worker-job-bytes")?, + "--max-worker-artifact-spools" => limits.max_artifact_spools = parse_count(&args, &mut index, "--max-worker-artifact-spools")?, + "--max-worker-artifact-spool-bytes" => { + limits.max_artifact_spool_bytes = parse_value(&args, &mut index, "--max-worker-artifact-spool-bytes")? + } + "--max-worker-state-entries" => limits.max_state_entries = parse_count(&args, &mut index, "--max-worker-state-entries")?, value => return Err(format!("unknown worker argument: {value}")), } } if !stdio { return Err("--stdio is required".to_string()); } - root.ok_or_else(|| "--workspace-root is required".to_string()) + let root = root.ok_or_else(|| "--workspace-root is required".to_string())?; + if limits.max_uploads == 0 + || limits.max_upload_bytes == 0 + || limits.max_jobs == 0 + || limits.max_job_bytes == 0 + || limits.max_artifact_spools == 0 + || limits.max_artifact_spool_bytes == 0 + || limits.max_state_entries == 0 + { + return Err("worker state limits must be positive".to_string()); + } + limits.validate()?; + Ok(WorkerConfig { root, limits, build_timeout, max_output_bytes }) +} + +fn parse_value(args: &[String], index: &mut usize, flag: &str) -> Result { + let value = args.get(*index + 1).ok_or_else(|| format!("{flag} requires a value"))?; + let value = value.parse::().map_err(|_| format!("{flag} requires an unsigned integer"))?; + *index += 2; + Ok(value) +} + +fn parse_count(args: &[String], index: &mut usize, flag: &str) -> Result { + usize::try_from(parse_value(args, index, flag)?).map_err(|_| format!("{flag} is too large")) } fn write_diagnostic(message: &str) { @@ -64,3 +126,7 @@ fn write_diagnostic(message: &str) { let _ = std::io::Write::write_all(&mut stderr, &message); let _ = std::io::Write::write_all(&mut stderr, b"\n"); } + +#[cfg(test)] +#[path = "main_ut.rs"] +mod tests; diff --git a/crates/bunkerbox-worker/src/main_ut.rs b/crates/bunkerbox-worker/src/main_ut.rs new file mode 100644 index 0000000..b6dfef3 --- /dev/null +++ b/crates/bunkerbox-worker/src/main_ut.rs @@ -0,0 +1,40 @@ +use super::*; + +#[test] +fn trusted_worker_limits_are_parsed_from_fixed_flags() { + let config = parse_args(vec![ + "--stdio".into(), + "--workspace-root".into(), + "/var/lib/bunkerbox-worker".into(), + "--build-timeout-ms".into(), + "2500".into(), + "--max-worker-uploads".into(), + "3".into(), + "--max-worker-upload-bytes".into(), + "4096".into(), + "--max-worker-jobs".into(), + "2".into(), + "--max-worker-job-bytes".into(), + "8192".into(), + "--max-worker-artifact-spools".into(), + "2".into(), + "--max-worker-artifact-spool-bytes".into(), + "16384".into(), + "--max-worker-state-entries".into(), + "99".into(), + ]) + .unwrap(); + assert_eq!(config.root, PathBuf::from("/var/lib/bunkerbox-worker")); + assert_eq!(config.build_timeout, Duration::from_millis(2500)); + assert_eq!(config.limits.max_uploads, 3); + assert_eq!(config.limits.max_job_bytes, 8192); + assert_eq!(config.limits.max_state_entries, 99); +} + +#[test] +fn trusted_worker_arguments_fail_closed() { + assert!(parse_args(vec!["--workspace-root".into(), "/tmp/root".into()]).is_err()); + assert!(parse_args(vec!["--stdio".into(), "--workspace-root".into(), "relative".into()]).is_err()); + assert!(parse_args(vec!["--stdio".into(), "--stdio".into(), "--workspace-root".into(), "/tmp/root".into()]).is_err()); + assert!(parse_args(vec!["--stdio".into(), "--workspace-root".into(), "/tmp/root".into(), "--build-timeout-ms".into(), "0".into()]).is_err()); +} diff --git a/crates/bunkerbox-worker/src/platform.rs b/crates/bunkerbox-worker/src/platform.rs index dc1dc3b..ea69974 100644 --- a/crates/bunkerbox-worker/src/platform.rs +++ b/crates/bunkerbox-worker/src/platform.rs @@ -147,6 +147,7 @@ pub fn list_names(directory: &File) -> io::Result> { unsafe { libc::close(duplicate) }; return Err(io::Error::last_os_error()); } + unsafe { libc::rewinddir(stream) }; let mut names = Vec::new(); loop { let entry = unsafe { libc::readdir(stream) }; diff --git a/crates/bunkerbox-worker/src/process.rs b/crates/bunkerbox-worker/src/process.rs index e08bd87..08f64e1 100644 --- a/crates/bunkerbox-worker/src/process.rs +++ b/crates/bunkerbox-worker/src/process.rs @@ -12,6 +12,7 @@ use std::sync::Arc; use std::thread; use std::time::{Duration, Instant}; +#[allow(dead_code)] pub const WORKER_BUILD_TIMEOUT: Duration = Duration::from_secs(30); pub const WORKER_MAX_OUTPUT_BYTES: u64 = 64 * 1024 * 1024; const OUTPUT_BUFFER_BYTES: usize = 8192; @@ -81,6 +82,15 @@ impl JobWorkspace { pub fn root(&self) -> &File { &self.root } + + pub fn create_with_reservation(parent: &File, reservation: storage::JobReservation) -> Result { + let job = Self::create(parent)?; + if let Err(error) = reservation.commit() { + drop(job); + return Err(error); + } + Ok(job) + } } impl Drop for JobWorkspace { @@ -94,7 +104,7 @@ impl Drop for JobWorkspace { pub fn cleanup_stale_jobs(parent: &File) -> io::Result<()> { for lock_name in platform::list_names(parent)?.into_iter().take(MAX_STALE_JOBS) { let Some(name) = lock_name.strip_suffix(".lock") else { continue }; - if !name.starts_with("job-") && !name.starts_with("artifact-") { + if !name.starts_with("job-") && !name.starts_with("artifact-") && !name.starts_with("reservation-job-") { continue; } let Ok(lock) = platform::open_lock_at(parent, &lock_name) else { continue }; @@ -106,8 +116,17 @@ pub fn cleanup_stale_jobs(parent: &File) -> io::Result<()> { Ok(()) } +#[allow(dead_code)] pub fn execute_build( job: &JobWorkspace, build: &WorkerBuild, request_id: WorkerRequestId, session_id: WorkerSessionId, sink: &S, disconnected: &dyn Fn() -> bool, +) -> Result { + execute_build_with_limits(job, build, request_id, session_id, sink, disconnected, WORKER_BUILD_TIMEOUT, WORKER_MAX_OUTPUT_BYTES) +} + +#[allow(clippy::too_many_arguments)] +pub fn execute_build_with_limits( + job: &JobWorkspace, build: &WorkerBuild, request_id: WorkerRequestId, session_id: WorkerSessionId, sink: &S, disconnected: &dyn Fn() -> bool, + build_timeout: Duration, max_output_bytes: u64, ) -> Result { validate_executable(build.trusted_executable_path())?; let cwd = storage::open_relative_directory(job.root(), build.cwd().as_str())?; @@ -159,9 +178,9 @@ pub fn execute_build( } let stop = Arc::new(AtomicBool::new(false)); - let (events_tx, events_rx) = mpsc::channel(); - let stdout_thread = spawn_pump(stdout, StreamKind::Stdout, events_tx.clone(), stop.clone()); - let stderr_thread = spawn_pump(stderr, StreamKind::Stderr, events_tx, stop.clone()); + let (events_tx, events_rx) = mpsc::sync_channel(32); + let stdout_thread = spawn_pump(stdout, StreamKind::Stdout, events_tx.clone(), stop.clone(), max_output_bytes); + let stderr_thread = spawn_pump(stderr, StreamKind::Stderr, events_tx, stop.clone(), max_output_bytes); let started = Instant::now(); let mut child_status = None; @@ -180,7 +199,7 @@ pub fn execute_build( group_killed = true; final_deadline = Some(Instant::now() + FINAL_DRAIN_TIMEOUT); } - if child_status.is_none() && failure.is_none() && started.elapsed() >= WORKER_BUILD_TIMEOUT { + if child_status.is_none() && failure.is_none() && started.elapsed() >= build_timeout { failure = Some("worker build timed out".to_string()); kill_group(pgid); group_killed = true; @@ -231,7 +250,7 @@ pub fn execute_build( continue; } output_total = output_total.saturating_add(bytes.len() as u64); - if output_total > WORKER_MAX_OUTPUT_BYTES { + if output_total > max_output_bytes { failure = Some("worker combined output limit exceeded".to_string()); kill_group(pgid); group_killed = true; @@ -272,6 +291,7 @@ pub fn execute_build( if !group_killed { kill_group(pgid); } + drop(events_rx); stop.store(true, Ordering::Release); let status = match child_status { Some(status) => status, @@ -286,7 +306,7 @@ pub fn execute_build( } fn spawn_pump( - mut reader: R, stream: StreamKind, sender: mpsc::Sender, stop: Arc, + mut reader: R, stream: StreamKind, sender: mpsc::SyncSender, stop: Arc, max_output_bytes: u64, ) -> thread::JoinHandle<()> { thread::spawn(move || { let mut buffer = [0u8; OUTPUT_BUFFER_BYTES]; @@ -302,7 +322,7 @@ fn spawn_pump( } Ok(count) => { total = total.saturating_add(count as u64); - if total > WORKER_MAX_OUTPUT_BYTES { + if total > max_output_bytes { let _ = sender.send(PumpEvent::Failed(stream, "worker output limit exceeded".to_string())); return; } diff --git a/crates/bunkerbox-worker/src/process_ut.rs b/crates/bunkerbox-worker/src/process_ut.rs index f8cf12a..870aceb 100644 --- a/crates/bunkerbox-worker/src/process_ut.rs +++ b/crates/bunkerbox-worker/src/process_ut.rs @@ -69,3 +69,34 @@ fn disconnect_and_inherited_pipes_do_not_leave_a_build_running() { assert!(error.contains("disconnected")); assert!(started.elapsed() < Duration::from_secs(2)); } + +#[test] +fn configured_build_timeout_overrides_worker_default() { + let temp = tempdir().unwrap(); + fs::set_permissions(temp.path(), fs::Permissions::from_mode(0o700)).unwrap(); + let jobs = platform::open_root(temp.path()).unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + let build = WorkerBuild::new( + "shell", + "/bin/sh", + vec!["-c".to_string(), "sleep 1".to_string()], + "", + Vec::new(), + vec![("PATH".to_string(), "/usr/bin:/bin".to_string())], + WorkerUploadId([5; 16]), + ) + .unwrap(); + let writer = FrameWriter::new(Vec::new()); + let error = execute_build_with_limits( + &job, + &build, + WorkerRequestId([9; 16]), + WorkerSessionId([10; 16]), + &writer, + &|| false, + Duration::from_millis(20), + WORKER_MAX_OUTPUT_BYTES, + ) + .unwrap_err(); + assert!(error.contains("timed out")); +} diff --git a/crates/bunkerbox-worker/src/storage.rs b/crates/bunkerbox-worker/src/storage.rs index f6774bd..c2d519a 100644 --- a/crates/bunkerbox-worker/src/storage.rs +++ b/crates/bunkerbox-worker/src/storage.rs @@ -19,6 +19,7 @@ const MANIFEST_FILE: &str = "manifest"; const COMPLETE_FILE: &str = "complete"; const LOCK_FILE: &str = "lock"; const FILES_DIRECTORY: &str = "files"; +const QUOTA_LOCK_FILE: &str = "quota.lock"; const MANIFEST_MAGIC: [u8; 4] = *b"BBWM"; const MANIFEST_VERSION: u16 = 1; const COPY_BUFFER_BYTES: usize = 64 * 1024; @@ -27,17 +28,99 @@ const MAX_STALE_UPLOADS_PER_SESSION: usize = 256; static NEXT_JOB_ID: AtomicU64 = AtomicU64::new(1); +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct WorkerStateLimits { + pub max_uploads: usize, + pub max_upload_bytes: u64, + pub max_jobs: usize, + pub max_job_bytes: u64, + pub max_artifact_spools: usize, + pub max_artifact_spool_bytes: u64, + pub max_state_entries: usize, +} + +impl Default for WorkerStateLimits { + fn default() -> Self { + Self { + max_uploads: 2, + max_upload_bytes: 1024 * 1024 * 1024, + max_jobs: 1, + max_job_bytes: 512 * 1024 * 1024, + max_artifact_spools: 1, + max_artifact_spool_bytes: 512 * 1024 * 1024, + max_state_entries: 20_000, + } + } +} + +impl WorkerStateLimits { + pub const MAX_COUNT: usize = 1_000; + pub const MAX_BYTES: u64 = 16 * 1024 * 1024 * 1024; + pub const MAX_ENTRIES: usize = 1_000_000; + + pub fn validate(self) -> Result<(), String> { + if self.max_uploads == 0 + || self.max_jobs == 0 + || self.max_artifact_spools == 0 + || self.max_uploads > Self::MAX_COUNT + || self.max_jobs > Self::MAX_COUNT + || self.max_artifact_spools > Self::MAX_COUNT + { + return Err(format!("worker state counts must be between 1 and {}", Self::MAX_COUNT)); + } + if self.max_upload_bytes == 0 + || self.max_job_bytes == 0 + || self.max_artifact_spool_bytes == 0 + || self.max_upload_bytes > Self::MAX_BYTES + || self.max_job_bytes > Self::MAX_BYTES + || self.max_artifact_spool_bytes > Self::MAX_BYTES + { + return Err(format!("worker state byte limits must be between 1 and {}", Self::MAX_BYTES)); + } + if self.max_state_entries == 0 || self.max_state_entries > Self::MAX_ENTRIES { + return Err(format!("worker state entry limit must be between 1 and {}", Self::MAX_ENTRIES)); + } + Ok(()) + } +} + +#[derive(Default)] +struct StateUsage { + uploads: usize, + upload_bytes: u64, + jobs: usize, + job_bytes: u64, + artifact_spools: usize, + artifact_spool_bytes: u64, + entries: usize, +} + pub struct UploadStore { sessions: File, jobs: File, + quota_lock: File, + limits: WorkerStateLimits, } impl UploadStore { + #[allow(dead_code)] pub fn new(root: &File) -> Result { + Self::new_with_limits(root, WorkerStateLimits::default()) + } + + pub fn new_with_limits(root: &File, limits: WorkerStateLimits) -> Result { + limits.validate()?; let state = private_directory(root, STATE_DIRECTORY)?; let sessions = private_directory(&state, SESSIONS_DIRECTORY)?; let jobs = private_directory(&state, JOBS_DIRECTORY)?; - let store = Self { sessions, jobs }; + let quota_lock = match platform::create_file_at(&state, QUOTA_LOCK_FILE, 0o600) { + Ok(file) => file, + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => { + platform::open_lock_at(&state, QUOTA_LOCK_FILE).map_err(|error| format!("open worker quota lock: {error}"))? + } + Err(error) => return Err(format!("create worker quota lock: {error}")), + }; + let store = Self { sessions, jobs, quota_lock, limits }; store.cleanup_stale().map_err(|error| format!("clean stale worker state: {error}"))?; Ok(store) } @@ -50,6 +133,11 @@ impl UploadStore { validate_upload_manifest(&entries).map_err(protocol_error)?; validate_manifest_structure(&entries)?; + let requested_bytes = manifest_file_bytes(&entries)?; + let _quota = self.acquire_quota_lock()?; + let usage = self.state_usage()?; + self.check_upload_quota(&usage, requested_bytes, entries.len())?; + let session = private_directory(&self.sessions, &hex_id(session_id.0))?; let uploads = private_directory(&session, UPLOADS_DIRECTORY)?; let token_name = hex_id(upload_id.0); @@ -125,12 +213,20 @@ impl UploadStore { return Err("worker upload identity mismatch".to_string()); } let files = open_existing_private_directory(&token, FILES_DIRECTORY, "worker upload files")?; - Ok(StoredUpload { token, lock, files, entries: manifest.entries }) + Ok(StoredUpload { + uploads: uploads.try_clone().map_err(|error| format!("clone worker uploads directory: {error}"))?, + token_name, + token, + lock, + files, + entries: manifest.entries, + }) } pub fn cleanup(&self, session_id: WorkerSessionId, upload_id: WorkerUploadId) -> Result<(), String> { require_nonzero_id(session_id.0, "worker session")?; require_nonzero_id(upload_id.0, "worker upload")?; + let _quota = self.acquire_quota_lock()?; let session = open_existing_private_directory(&self.sessions, &hex_id(session_id.0), "worker session")?; let uploads = open_existing_private_directory(&session, UPLOADS_DIRECTORY, "worker uploads")?; let token_name = hex_id(upload_id.0); @@ -146,6 +242,186 @@ impl UploadStore { self.jobs.try_clone().map_err(|error| format!("clone worker jobs directory: {error}")) } + pub fn reserve_job(&self, bytes: u64, entries: usize) -> Result { + let _quota = self.acquire_quota_lock()?; + let usage = self.state_usage()?; + if usage.jobs.saturating_add(1) > self.limits.max_jobs { + return Err("worker job quota exceeded".to_string()); + } + if usage.job_bytes.saturating_add(bytes) > self.limits.max_job_bytes { + return Err("worker job byte quota exceeded".to_string()); + } + if usage.entries.saturating_add(entries) > self.limits.max_state_entries { + return Err("worker state entry quota exceeded".to_string()); + } + for _ in 0..32 { + let name = format!("reservation-job-{bytes}-{entries}-{}-{}.lock", unsafe { libc::getpid() }, next_job_id()); + let marker = match platform::create_file_at(&self.jobs, &name, 0o600) { + Ok(marker) => marker, + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(format!("reserve worker job quota: {error}")), + }; + if !platform::lock_exclusive(&marker).map_err(|error| format!("lock worker job quota: {error}"))? { + let _ = platform::unlink_at(&self.jobs, &name, 0); + continue; + } + return Ok(JobReservation { + parent: self.jobs.try_clone().map_err(|error| format!("clone worker jobs directory: {error}"))?, + name, + marker, + committed: false, + }); + } + Err("could not reserve worker job quota".to_string()) + } + + pub fn validate_job(&self, job: &File) -> Result<(), String> { + let usage = walk_tree(job)?; + if usage.bytes > self.limits.max_job_bytes { + return Err("worker job exceeded its byte quota".to_string()); + } + if usage.entries > self.limits.max_state_entries { + return Err("worker job exceeded the state entry quota".to_string()); + } + Ok(()) + } + + pub fn capture_artifacts_with_interrupt( + &self, job_root: &File, paths: &[WorkerArtifactPath], max_file_bytes: u64, max_total_bytes: u64, disconnected: &dyn Fn() -> bool, + ) -> Result { + ArtifactSpool::capture_with_store(self, job_root, paths, max_file_bytes, max_total_bytes, disconnected) + } + + fn acquire_quota_lock(&self) -> Result { + let lock = self.quota_lock.try_clone().map_err(|error| format!("clone worker quota lock: {error}"))?; + for _ in 0..5000 { + if platform::lock_exclusive(&lock).map_err(|error| format!("lock worker quota: {error}"))? { + return Ok(lock); + } + std::thread::sleep(std::time::Duration::from_millis(1)); + } + Err("worker state quota lock timed out".to_string()) + } + + fn check_upload_quota(&self, usage: &StateUsage, bytes: u64, entries: usize) -> Result<(), String> { + if usage.uploads.saturating_add(1) > self.limits.max_uploads { + return Err("worker upload quota exceeded".to_string()); + } + if usage.upload_bytes.saturating_add(bytes) > self.limits.max_upload_bytes { + return Err("worker upload byte quota exceeded".to_string()); + } + if usage.entries.saturating_add(entries) > self.limits.max_state_entries { + return Err("worker state entry quota exceeded".to_string()); + } + Ok(()) + } + + fn validate_current_usage(&self) -> Result<(), String> { + let usage = self.state_usage()?; + if usage.uploads > self.limits.max_uploads || usage.upload_bytes > self.limits.max_upload_bytes { + return Err("worker upload quota exceeded".to_string()); + } + if usage.jobs > self.limits.max_jobs || usage.job_bytes > self.limits.max_job_bytes { + return Err("worker job quota exceeded".to_string()); + } + if usage.artifact_spools > self.limits.max_artifact_spools || usage.artifact_spool_bytes > self.limits.max_artifact_spool_bytes { + return Err("worker artifact spool quota exceeded".to_string()); + } + if usage.entries > self.limits.max_state_entries { + return Err("worker state entry quota exceeded".to_string()); + } + Ok(()) + } + + fn state_usage(&self) -> Result { + let mut usage = StateUsage::default(); + for session_name in platform::list_names(&self.sessions).map_err(|error| format!("list worker sessions: {error}"))? { + if !is_hex_id(&session_name) { + continue; + } + let Ok(session) = platform::open_dir_at(&self.sessions, &session_name) else { continue }; + let Ok(uploads) = platform::open_dir_at(&session, UPLOADS_DIRECTORY) else { continue }; + for upload_name in platform::list_names(&uploads).map_err(|error| format!("list worker uploads: {error}"))? { + if !is_hex_id(&upload_name) { + continue; + } + usage.uploads = usage.uploads.saturating_add(1); + let Ok(upload) = platform::open_dir_at(&uploads, &upload_name) else { continue }; + if let Ok(manifest_file) = platform::open_file_at(&upload, MANIFEST_FILE) { + if let Ok(manifest) = StoredManifest::read(manifest_file) { + usage.upload_bytes = usage.upload_bytes.saturating_add(manifest_file_bytes(&manifest.entries)?); + usage.entries = usage.entries.saturating_add(manifest.entries.len()); + } + } + } + } + for name in platform::list_names(&self.jobs).map_err(|error| format!("list worker jobs: {error}"))? { + if let Some((bytes, entries)) = reservation_name(&name, "reservation-job-") { + usage.jobs = usage.jobs.saturating_add(1); + usage.job_bytes = usage.job_bytes.saturating_add(bytes); + usage.entries = usage.entries.saturating_add(entries); + continue; + } + if name.ends_with(".lock") { + continue; + } + if name.starts_with("job-") { + usage.jobs = usage.jobs.saturating_add(1); + if let Ok(job) = platform::open_dir_at(&self.jobs, &name) { + let tree = walk_tree(&job)?; + usage.job_bytes = usage.job_bytes.saturating_add(tree.bytes); + usage.entries = usage.entries.saturating_add(tree.entries); + } + } else if name.starts_with("artifact-") { + usage.artifact_spools = usage.artifact_spools.saturating_add(1); + if let Ok(spool) = platform::open_dir_at(&self.jobs, &name) { + let tree = walk_tree(&spool)?; + let (reserved_bytes, reserved_entries) = artifact_reservation_name(&name).unwrap_or((0, 0)); + usage.artifact_spool_bytes = usage.artifact_spool_bytes.saturating_add(tree.bytes.max(reserved_bytes)); + usage.entries = usage.entries.saturating_add(tree.entries.max(reserved_entries)); + } + } + } + Ok(usage) + } + + fn reserve_artifact_root(&self, bytes: u64, entries: usize) -> Result<(File, String, String, File, File), String> { + let _quota = self.acquire_quota_lock()?; + let usage = self.state_usage()?; + if usage.artifact_spools.saturating_add(1) > self.limits.max_artifact_spools { + return Err("worker artifact spool quota exceeded".to_string()); + } + if usage.artifact_spool_bytes.saturating_add(bytes) > self.limits.max_artifact_spool_bytes { + return Err("worker artifact spool byte quota exceeded".to_string()); + } + if usage.entries.saturating_add(entries) > self.limits.max_state_entries { + return Err("worker state entry quota exceeded".to_string()); + } + for _ in 0..32 { + let name = format!("artifact-{}-{}-{bytes}-{entries}", unsafe { libc::getpid() }, next_job_id()); + let lock_name = format!("{name}.lock"); + let lock = match platform::create_file_at(&self.jobs, &lock_name, 0o600) { + Ok(lock) => lock, + Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, + Err(error) => return Err(format!("create worker artifact lock: {error}")), + }; + if !platform::lock_exclusive(&lock).map_err(|error| format!("lock worker artifact spool: {error}"))? { + let _ = platform::unlink_at(&self.jobs, &lock_name, 0); + continue; + } + if let Err(error) = platform::create_dir_at(&self.jobs, &name, 0o700) { + let _ = platform::unlink_at(&self.jobs, &lock_name, 0); + if error.kind() == io::ErrorKind::AlreadyExists { + continue; + } + return Err(format!("create worker artifact spool: {error}")); + } + let root = platform::open_dir_at(&self.jobs, &name).map_err(|error| format!("open worker artifact spool: {error}"))?; + return Ok((self.jobs.try_clone().map_err(|error| format!("clone worker jobs directory: {error}"))?, name, lock_name, lock, root)); + } + Err("could not reserve a worker artifact spool".to_string()) + } + fn cleanup_stale(&self) -> io::Result<()> { let sessions = platform::list_names(&self.sessions)?; for session_name in sessions.into_iter().take(MAX_STALE_SESSIONS) { @@ -172,6 +448,31 @@ impl UploadStore { } } +pub struct JobReservation { + parent: File, + name: String, + marker: File, + committed: bool, +} + +impl JobReservation { + pub fn commit(mut self) -> Result<(), String> { + platform::unlink_at(&self.parent, &self.name, 0).map_err(|error| format!("release worker job reservation: {error}"))?; + self.committed = true; + let _ = &self.marker; + Ok(()) + } +} + +impl Drop for JobReservation { + fn drop(&mut self) { + if !self.committed { + let _ = &self.marker; + let _ = platform::unlink_at(&self.parent, &self.name, 0); + } + } +} + pub struct UploadTransaction { uploads: File, token_name: String, @@ -267,16 +568,33 @@ impl Drop for UploadTransaction { } pub struct StoredUpload { + uploads: File, + token_name: String, token: File, lock: File, files: File, entries: Vec, } +impl Drop for StoredUpload { + fn drop(&mut self) { + let _ = (&self.token, &self.lock, &self.files); + let _ = platform::remove_tree_at(&self.uploads, &self.token_name); + } +} + impl StoredUpload { + #[allow(dead_code)] pub fn materialize(&self, destination: &File) -> Result<(), String> { + self.materialize_with_interrupt(destination, &|| false) + } + + pub fn materialize_with_interrupt(&self, destination: &File, disconnected: &dyn Fn() -> bool) -> Result<(), String> { let _ = (&self.token, &self.lock); for entry in &self.entries { + if disconnected() { + return Err("worker input disconnected during materialization".to_string()); + } match entry.kind() { WorkerEntryKind::Directory => ensure_directory(destination, entry.path().as_str(), entry.mode())?, WorkerEntryKind::File => { @@ -284,18 +602,27 @@ impl StoredUpload { let source_metadata = platform::stat_fd(source.as_raw_fd()).map_err(|error| format!("stat stored worker file: {error}"))?; validate_regular_file(&source_metadata, entry.size(), entry.mode(), entry.path().as_str())?; let target = create_relative_file(destination, entry.path().as_str(), entry.mode())?; - copy_and_verify(&source, &target, entry)?; + copy_and_verify_with_interrupt(&source, &target, entry, disconnected)?; } } } platform::sync_fd(destination).map_err(|error| format!("flush worker job workspace: {error}"))?; Ok(()) } + + pub fn declared_bytes(&self) -> Result { + manifest_file_bytes(&self.entries) + } + + pub fn entry_count(&self) -> usize { + self.entries.len() + } } pub struct ArtifactSpool { parent: File, name: String, + lock_name: String, root: File, lock: File, files: File, @@ -305,19 +632,37 @@ pub struct ArtifactSpool { } impl ArtifactSpool { - pub fn capture(parent: &File, job_root: &File, paths: &[WorkerArtifactPath], max_file_bytes: u64, max_total_bytes: u64) -> Result { + fn capture_with_store( + store: &UploadStore, job_root: &File, paths: &[WorkerArtifactPath], max_file_bytes: u64, max_total_bytes: u64, + disconnected: &dyn Fn() -> bool, + ) -> Result { + let reserved_entries = paths.len().saturating_add(2); + let (parent, name, lock_name, lock, root) = store.reserve_artifact_root(max_total_bytes, reserved_entries)?; + let spool = Self::capture_reserved(parent, name, lock_name, lock, root, job_root, paths, max_file_bytes, max_total_bytes, disconnected)?; + let _quota = store.acquire_quota_lock()?; + if let Err(error) = store.validate_current_usage() { + drop(spool); + return Err(error); + } + Ok(spool) + } + + #[allow(clippy::too_many_arguments)] + fn capture_reserved( + parent: File, name: String, lock_name: String, lock: File, root: File, job_root: &File, paths: &[WorkerArtifactPath], max_file_bytes: u64, + max_total_bytes: u64, disconnected: &dyn Fn() -> bool, + ) -> Result { if max_file_bytes == 0 || max_file_bytes > MAX_WORKER_ARTIFACT_FILE_BYTES { return Err("worker artifact per-file limit is invalid".to_string()); } if max_total_bytes == 0 || max_total_bytes > MAX_WORKER_ARTIFACT_TOTAL_BYTES { return Err("worker artifact total limit is invalid".to_string()); } - let (name, lock_name, lock, root) = reserve_artifact_root(parent)?; let files = match private_directory(&root, FILES_DIRECTORY) { Ok(files) => files, Err(error) => { - let _ = platform::remove_tree_at(parent, &name); - let _ = platform::unlink_at(parent, &lock_name, 0); + let _ = platform::remove_tree_at(&parent, &name); + let _ = platform::unlink_at(&parent, &lock_name, 0); return Err(error); } }; @@ -326,6 +671,9 @@ impl ArtifactSpool { let mut entries = Vec::with_capacity(paths.len()); let mut total_bytes = 0u64; for path in paths { + if disconnected() { + return Err("worker input disconnected during artifact capture".to_string()); + } let source = open_relative_file(job_root, path.as_str())?; let before = platform::stat_fd(source.as_raw_fd()).map_err(|error| format!("stat worker artifact {}: {error}", path.as_str()))?; validate_artifact_source(&before, path.as_str())?; @@ -342,7 +690,7 @@ impl ArtifactSpool { ensure_directory(&files, parents, 0o700)?; } let destination = create_relative_file(&files, path.as_str(), mode)?; - let digest = copy_artifact(&source, &destination, size, path.as_str())?; + let digest = copy_artifact_with_interrupt(&source, &destination, size, path.as_str(), disconnected)?; let after = platform::stat_fd(source.as_raw_fd()).map_err(|error| format!("restat worker artifact {}: {error}", path.as_str()))?; if after.st_dev != before.st_dev || after.st_ino != before.st_ino @@ -359,19 +707,12 @@ impl ArtifactSpool { })(); match result { - Ok((artifact_set_id, entries, total_bytes)) => Ok(Self { - parent: parent.try_clone().map_err(|error| format!("clone worker jobs directory: {error}"))?, - name, - root, - lock, - files, - artifact_set_id, - entries, - total_bytes, - }), + Ok((artifact_set_id, entries, total_bytes)) => { + Ok(Self { parent, name, lock_name, root, lock, files, artifact_set_id, entries, total_bytes }) + } Err(error) => { - let _ = platform::remove_tree_at(parent, &name); - let _ = platform::unlink_at(parent, &lock_name, 0); + let _ = platform::remove_tree_at(&parent, &name); + let _ = platform::unlink_at(&parent, &lock_name, 0); Err(error) } } @@ -402,6 +743,7 @@ impl Drop for ArtifactSpool { fn drop(&mut self) { let _ = (&self.root, &self.lock, &self.files); let _ = platform::remove_tree_at(&self.parent, &self.name); + let _ = platform::unlink_at(&self.parent, &self.lock_name, 0); } } @@ -554,44 +896,6 @@ fn private_directory(parent: &File, name: &str) -> Result { Ok(directory) } -fn reserve_artifact_root(parent: &File) -> Result<(String, String, File, File), String> { - for _ in 0..32 { - let name = format!("artifact-{}-{}", unsafe { libc::getpid() }, next_job_id()); - let lock_name = format!("{name}.lock"); - let lock = match platform::create_file_at(parent, &lock_name, 0o600) { - Ok(lock) => lock, - Err(error) if error.kind() == io::ErrorKind::AlreadyExists => continue, - Err(error) => return Err(format!("create worker artifact lock: {error}")), - }; - if !platform::lock_exclusive(&lock).map_err(|error| format!("lock worker artifact spool: {error}"))? { - let _ = platform::unlink_at(parent, &lock_name, 0); - continue; - } - if let Err(error) = platform::create_dir_at(parent, &name, 0o700) { - let _ = platform::unlink_at(parent, &lock_name, 0); - if error.kind() == io::ErrorKind::AlreadyExists { - continue; - } - return Err(format!("create worker artifact spool: {error}")); - } - let root = match platform::open_dir_at(parent, &name) { - Ok(root) => root, - Err(error) => { - let _ = platform::remove_tree_at(parent, &name); - let _ = platform::unlink_at(parent, &lock_name, 0); - return Err(format!("open worker artifact spool: {error}")); - } - }; - if let Err(error) = platform::validate_private_directory(&root, "worker artifact spool") { - let _ = platform::remove_tree_at(parent, &name); - let _ = platform::unlink_at(parent, &lock_name, 0); - return Err(error); - } - return Ok((name, lock_name, lock, root)); - } - Err("could not reserve a worker artifact spool".to_string()) -} - fn validate_artifact_source(metadata: &libc::stat, path: &str) -> Result<(), String> { if metadata.st_mode & libc::S_IFMT != libc::S_IFREG || metadata.st_nlink != 1 { return Err(format!("worker artifact is not a private regular file: {path}")); @@ -602,13 +906,18 @@ fn validate_artifact_source(metadata: &libc::stat, path: &str) -> Result<(), Str Ok(()) } -fn copy_artifact(source: &File, destination: &File, expected_size: u64, path: &str) -> Result { +fn copy_artifact_with_interrupt( + source: &File, destination: &File, expected_size: u64, path: &str, disconnected: &dyn Fn() -> bool, +) -> Result { let mut source = source.try_clone().map_err(|error| format!("clone worker artifact {path}: {error}"))?; let mut destination = destination.try_clone().map_err(|error| format!("clone worker artifact spool {path}: {error}"))?; let mut hasher = Sha256::new(); let mut copied = 0u64; let mut buffer = vec![0u8; COPY_BUFFER_BYTES]; loop { + if disconnected() { + return Err(format!("worker input disconnected during artifact capture: {path}")); + } let count = source.read(&mut buffer).map_err(|error| format!("read worker artifact {path}: {error}"))?; if count == 0 { break; @@ -735,13 +1044,23 @@ fn validate_manifest_structure(entries: &[WorkerUploadEntry]) -> Result<(), Stri Ok(()) } +#[allow(dead_code)] fn copy_and_verify(source: &File, destination: &File, entry: &WorkerUploadEntry) -> Result<(), String> { + copy_and_verify_with_interrupt(source, destination, entry, &|| false) +} + +fn copy_and_verify_with_interrupt( + source: &File, destination: &File, entry: &WorkerUploadEntry, disconnected: &dyn Fn() -> bool, +) -> Result<(), String> { let mut source = source.try_clone().map_err(|error| format!("clone stored worker file: {error}"))?; let mut destination = destination.try_clone().map_err(|error| format!("clone materialized worker file: {error}"))?; let mut buffer = vec![0u8; COPY_BUFFER_BYTES]; let mut hasher = Sha256::new(); let mut copied = 0u64; loop { + if disconnected() { + return Err("worker input disconnected during materialization".to_string()); + } let count = source.read(&mut buffer).map_err(|error| format!("read stored worker file: {error}"))?; if count == 0 { break; @@ -784,6 +1103,62 @@ fn read_bounded(mut file: File, maximum: usize) -> Result, String> { Ok(bytes) } +fn manifest_file_bytes(entries: &[WorkerUploadEntry]) -> Result { + entries + .iter() + .filter(|entry| entry.kind() == WorkerEntryKind::File) + .try_fold(0u64, |total, entry| total.checked_add(entry.size()).ok_or_else(|| "worker manifest byte total overflow".to_string())) +} + +#[derive(Default)] +struct TreeUsage { + bytes: u64, + entries: usize, +} + +fn walk_tree(directory: &File) -> Result { + let mut usage = TreeUsage::default(); + for name in platform::list_names(directory).map_err(|error| format!("list worker state: {error}"))? { + let metadata = platform::stat_at(directory, &name).map_err(|error| format!("stat worker state entry: {error}"))?; + match metadata.st_mode & libc::S_IFMT { + mode if mode == libc::S_IFDIR => { + usage.entries = usage.entries.saturating_add(1); + let child = platform::open_dir_at(directory, &name).map_err(|error| format!("open worker state directory: {error}"))?; + let child_usage = walk_tree(&child)?; + usage.bytes = usage.bytes.saturating_add(child_usage.bytes); + usage.entries = usage.entries.saturating_add(child_usage.entries); + } + mode if mode == libc::S_IFREG => { + usage.entries = usage.entries.saturating_add(1); + if metadata.st_size < 0 { + return Err("worker state file has a negative size".to_string()); + } + usage.bytes = usage.bytes.saturating_add(metadata.st_size as u64); + } + _ => return Err(format!("worker state contains an unsupported file type: {name}")), + } + } + Ok(usage) +} + +fn reservation_name(name: &str, prefix: &str) -> Option<(u64, usize)> { + let value = name.strip_prefix(prefix)?.trim_end_matches(".lock"); + let mut parts = value.split('-'); + let bytes = parts.next()?.parse().ok()?; + let entries = parts.next()?.parse().ok()?; + Some((bytes, entries)) +} + +fn artifact_reservation_name(name: &str) -> Option<(u64, usize)> { + let value = name.strip_prefix("artifact-")?; + let mut parts = value.split('-'); + let _pid = parts.next()?.parse::().ok()?; + let _id = parts.next()?.parse::().ok()?; + let bytes = parts.next()?.parse().ok()?; + let entries = parts.next()?.parse().ok()?; + Some((bytes, entries)) +} + fn put_u32(bytes: &mut Vec, value: usize) -> Result<(), String> { bytes.extend_from_slice(&u32::try_from(value).map_err(|_| "worker manifest count does not fit in u32".to_string())?.to_le_bytes()); Ok(()) diff --git a/crates/bunkerbox-worker/src/storage_ut.rs b/crates/bunkerbox-worker/src/storage_ut.rs index 4b9984e..b01101a 100644 --- a/crates/bunkerbox-worker/src/storage_ut.rs +++ b/crates/bunkerbox-worker/src/storage_ut.rs @@ -21,6 +21,16 @@ fn store_fixture() -> (TempDir, UploadStore) { (temp, store) } +fn limited_store(limits: WorkerStateLimits) -> (TempDir, UploadStore) { + let temp = tempdir().unwrap(); + let root_path = temp.path().join("root"); + fs::create_dir(&root_path).unwrap(); + fs::set_permissions(&root_path, fs::Permissions::from_mode(0o700)).unwrap(); + let root = platform::open_root(&root_path).unwrap(); + let store = UploadStore::new_with_limits(&root, limits).unwrap(); + (temp, store) +} + fn entries(contents: &[u8]) -> Vec { vec![ WorkerUploadEntry::directory("src", 0o755).unwrap(), @@ -102,3 +112,24 @@ fn replaced_stored_file_symlink_is_rejected_during_materialization() { fn unsupported_entry_kind_is_not_accepted_by_the_manifest_constructor() { assert!(WorkerUploadEntry::new("node", WorkerEntryKind::Directory, 0o755, 1, None).is_err()); } + +#[test] +fn live_upload_reservation_enforces_count_and_releases_on_drop() { + let limits = WorkerStateLimits { max_uploads: 1, ..WorkerStateLimits::default() }; + let (_temp, store) = limited_store(limits); + let contents = b"stored"; + let transaction = store.begin(SESSION, UPLOAD, entries(contents)).unwrap(); + assert!(store.begin(WorkerSessionId([3; 16]), WorkerUploadId([4; 16]), entries(contents)).is_err()); + drop(transaction); + assert!(store.begin(WorkerSessionId([3; 16]), WorkerUploadId([4; 16]), entries(contents)).is_ok()); +} + +#[test] +fn live_job_reservation_enforces_bytes_and_releases_on_drop() { + let limits = WorkerStateLimits { max_jobs: 1, max_job_bytes: 5, ..WorkerStateLimits::default() }; + let (_temp, store) = limited_store(limits); + let reservation = store.reserve_job(5, 1).unwrap(); + assert!(store.reserve_job(1, 1).is_err()); + drop(reservation); + assert!(store.reserve_job(5, 1).is_ok()); +} diff --git a/crates/bunkerbox-worker/src/worker.rs b/crates/bunkerbox-worker/src/worker.rs index 495d132..9d29267 100644 --- a/crates/bunkerbox-worker/src/worker.rs +++ b/crates/bunkerbox-worker/src/worker.rs @@ -42,14 +42,33 @@ impl OutputSink for FrameWriter { pub struct WorkerService { store: UploadStore, + build_timeout: std::time::Duration, + max_output_bytes: u64, } impl WorkerService { + #[allow(dead_code)] pub fn new(root: &std::fs::File) -> Result { - let store = UploadStore::new(root)?; + Self::new_with_limits(root, crate::storage::WorkerStateLimits::default(), process::WORKER_BUILD_TIMEOUT) + } + + #[allow(dead_code)] + pub fn new_with_limits( + root: &std::fs::File, limits: crate::storage::WorkerStateLimits, build_timeout: std::time::Duration, + ) -> Result { + Self::new_with_config(root, limits, build_timeout, process::WORKER_MAX_OUTPUT_BYTES) + } + + pub fn new_with_config( + root: &std::fs::File, limits: crate::storage::WorkerStateLimits, build_timeout: std::time::Duration, max_output_bytes: u64, + ) -> Result { + if max_output_bytes == 0 || max_output_bytes > process::WORKER_MAX_OUTPUT_BYTES { + return Err(format!("worker output limit must be between 1 and {}", process::WORKER_MAX_OUTPUT_BYTES)); + } + let store = UploadStore::new_with_limits(root, limits)?; let jobs = store.jobs_directory()?; process::cleanup_stale_jobs(&jobs).map_err(|error| format!("clean stale worker jobs: {error}"))?; - Ok(Self { store }) + Ok(Self { store, build_timeout, max_output_bytes }) } pub fn run(&self, input: R, writer: &FrameWriter) -> Result<(), String> { @@ -196,27 +215,56 @@ impl WorkerService { } }; let jobs = self.store.jobs_directory()?; - let job = match JobWorkspace::create(&jobs) { + let reservation = match self.store.reserve_job(upload.declared_bytes()?, upload.entry_count()) { + Ok(reservation) => reservation, + Err(error) => { + send_error(writer, request_id, session_id, WorkerOperation::Build, WorkerErrorKind::Build, &error)?; + return Ok(()); + } + }; + let job = match JobWorkspace::create_with_reservation(&jobs, reservation) { Ok(job) => job, Err(error) => { send_error(writer, request_id, session_id, WorkerOperation::Build, WorkerErrorKind::Build, &error)?; return Ok(()); } }; - if let Err(error) = upload.materialize(job.root()) { + if let Err(error) = upload.materialize_with_interrupt(job.root(), &|| input.disconnected()) { send_error(writer, request_id, session_id, WorkerOperation::Upload, WorkerErrorKind::Upload, &error)?; return Ok(()); } - let exit_code = match process::execute_build(&job, &build, request_id, session_id, writer, &|| input.disconnected()) { + let exit_code = match process::execute_build_with_limits( + &job, + &build, + request_id, + session_id, + writer, + &|| input.disconnected(), + self.build_timeout, + self.max_output_bytes, + ) { Ok(exit_code) => exit_code, Err(error) => { send_error(writer, request_id, session_id, WorkerOperation::Build, WorkerErrorKind::Build, &error)?; return Ok(()); } }; + if let Err(error) = self.store.validate_job(job.root()) { + send_error(writer, request_id, session_id, WorkerOperation::Build, WorkerErrorKind::Build, &error)?; + return self.wait_for_cleanup( + input, + writer, + BuildCleanup { request_id, session_id, upload_token: build.upload_token(), protocol_version }, + ); + } let artifact_spool = if exit_code == 0 && !build.artifact_paths().is_empty() { - match ArtifactSpool::capture(&jobs, job.root(), build.artifact_paths(), build.artifact_max_file_bytes(), build.artifact_max_total_bytes()) - { + match self.store.capture_artifacts_with_interrupt( + job.root(), + build.artifact_paths(), + build.artifact_max_file_bytes(), + build.artifact_max_total_bytes(), + &|| input.disconnected(), + ) { Ok(spool) => Some(spool), Err(error) => { send_error(writer, request_id, session_id, WorkerOperation::Artifact, WorkerErrorKind::Artifact, &error)?; @@ -441,9 +489,11 @@ fn truncate_message(message: &str) -> &str { &message[..end] } -pub fn run_stdio(root: &std::path::Path) -> Result<(), String> { +pub fn run_stdio_with_config( + root: &std::path::Path, limits: crate::storage::WorkerStateLimits, build_timeout: std::time::Duration, max_output_bytes: u64, +) -> Result<(), String> { let root = crate::platform::open_root(root)?; - let service = WorkerService::new(&root)?; + let service = WorkerService::new_with_config(&root, limits, build_timeout, max_output_bytes)?; let stdin = io::stdin(); let input = stdin; let writer = Arc::new(FrameWriter::new(io::stdout())); From 6f886303aaaead08d6e0b22368474405475eb2e7 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 01:56:41 +0200 Subject: [PATCH 36/52] Add integration test fixtures --- tests/fixtures/remote-cargo/Cargo.toml | 5 +++++ tests/fixtures/remote-cargo/build.rs | 20 ++++++++++++++++++++ tests/fixtures/remote-cargo/src/main.rs | 7 +++++++ 3 files changed, 32 insertions(+) create mode 100644 tests/fixtures/remote-cargo/Cargo.toml create mode 100644 tests/fixtures/remote-cargo/build.rs create mode 100644 tests/fixtures/remote-cargo/src/main.rs diff --git a/tests/fixtures/remote-cargo/Cargo.toml b/tests/fixtures/remote-cargo/Cargo.toml new file mode 100644 index 0000000..0e30b9a --- /dev/null +++ b/tests/fixtures/remote-cargo/Cargo.toml @@ -0,0 +1,5 @@ +[package] +name = "bunkerbox-remote-cargo-fixture" +version = "0.1.0" +edition = "2021" +build = "build.rs" diff --git a/tests/fixtures/remote-cargo/build.rs b/tests/fixtures/remote-cargo/build.rs new file mode 100644 index 0000000..25091a2 --- /dev/null +++ b/tests/fixtures/remote-cargo/build.rs @@ -0,0 +1,20 @@ +use std::env; +use std::fs; +use std::path::PathBuf; +use std::process::Command; + +fn main() { + let cargo = env::var_os("CARGO").expect("Cargo must provide CARGO to build scripts"); + let cargo_path = PathBuf::from(&cargo); + assert_ne!(cargo_path.file_name().and_then(|name| name.to_str()), Some("bunkerbox-remote")); + assert_ne!(cargo_path, PathBuf::from("/usr/local/bunkerbox/bin/cargo")); + let version = Command::new(&cargo).arg("--version").output().expect("target Cargo must execute nested --version"); + assert!(version.status.success(), "target Cargo --version failed"); + + let out_dir = PathBuf::from(env::var_os("OUT_DIR").expect("Cargo must provide OUT_DIR")); + fs::write(out_dir.join("build_marker.txt"), "bunkerbox-cargo-fixture-build-script\n").expect("write Cargo fixture build marker"); + let artifact = PathBuf::from(env::var_os("CARGO_MANIFEST_DIR").expect("Cargo must provide CARGO_MANIFEST_DIR")) + .join("target/debug/bunkerbox-cargo-fixture-artifact.txt"); + fs::write(artifact, "bunkerbox-cargo-fixture-artifact\n").expect("write Cargo fixture artifact"); + println!("cargo:warning=bunkerbox-cargo-fixture-build-script"); +} diff --git a/tests/fixtures/remote-cargo/src/main.rs b/tests/fixtures/remote-cargo/src/main.rs new file mode 100644 index 0000000..b6ce51b --- /dev/null +++ b/tests/fixtures/remote-cargo/src/main.rs @@ -0,0 +1,7 @@ +const BUILD_MARKER: &str = include_str!(concat!(env!("OUT_DIR"), "/build_marker.txt")); + +fn main() { + print!("{BUILD_MARKER}"); + println!("bunkerbox-cargo-fixture-stdout"); + eprintln!("bunkerbox-cargo-fixture-stderr"); +} From 377af09d29fff647c8ebf472173e4008f413ef94 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 01:57:15 +0200 Subject: [PATCH 37/52] Wire remote proto --- src/artifact.rs | 24 +- src/bin/bunkerbox-remote.rs | 20 +- src/cfg.rs | 11 + src/daemon.rs | 442 +++++++++++++++++++++++++++++++-- src/guest_install.rs | 35 ++- src/loopback.rs | 251 +++++++++++++++---- src/main.rs | 11 +- src/remote.rs | 215 +++++++++++++++- src/remote_client.rs | 12 + src/remote_target.rs | 145 ++++++++++- src/snapshot.rs | 65 ++++- src/ssh.rs | 472 ++++++++++++++++++++++++++++-------- src/vscomm/mod.rs | 39 +++ 13 files changed, 1540 insertions(+), 202 deletions(-) diff --git a/src/artifact.rs b/src/artifact.rs index 5170cd3..1fc77fd 100644 --- a/src/artifact.rs +++ b/src/artifact.rs @@ -1,3 +1,4 @@ +use crate::remote::RemoteExecutionControl; use sha2::{Digest, Sha256}; use std::collections::BTreeSet; use std::ffi::CString; @@ -190,6 +191,12 @@ pub struct LocalArtifactSpool { impl LocalArtifactSpool { pub fn capture(job_root: &Path, parent: &Path, policy: &ArtifactPolicy, limits: ArtifactLimits) -> Result { + Self::capture_with_control(job_root, parent, policy, limits, &RemoteExecutionControl::new()) + } + + pub fn capture_with_control( + job_root: &Path, parent: &Path, policy: &ArtifactPolicy, limits: ArtifactLimits, control: &RemoteExecutionControl, + ) -> Result { policy.validate_limits(limits)?; let path = create_unique_directory(parent, "artifact-spool")?; let root = match open_directory(&path) { @@ -204,6 +211,9 @@ impl LocalArtifactSpool { let mut entries = Vec::with_capacity(policy.paths.len()); let mut total = 0u64; for relative in policy.paths() { + if control.is_cancelled() { + return Err("artifact capture cancelled".to_string()); + } let source = open_regular_file(job_root, relative)?; let metadata = source.metadata().map_err(|error| format!("stat artifact {relative}: {error}"))?; if metadata.nlink() != 1 { @@ -219,7 +229,7 @@ impl LocalArtifactSpool { } let mode = metadata.mode() & 0o777; let destination = create_relative_file(&root, relative, mode)?; - let digest = copy_and_hash(&source, &destination, size, relative)?; + let digest = copy_and_hash(&source, &destination, size, relative, control)?; let after = source.metadata().map_err(|error| format!("restat artifact {relative}: {error}"))?; if after.dev() != metadata.dev() || after.ino() != metadata.ino() @@ -309,9 +319,16 @@ impl ArtifactPublication { } pub fn copy_from_reader(&mut self, index: usize, reader: &mut R) -> Result<(), String> { + self.copy_from_reader_with_control(index, reader, &RemoteExecutionControl::new()) + } + + pub fn copy_from_reader_with_control(&mut self, index: usize, reader: &mut R, control: &RemoteExecutionControl) -> Result<(), String> { let mut writer = self.begin(index)?; let mut buffer = [0u8; COPY_BUFFER_BYTES]; loop { + if control.is_cancelled() { + return Err("artifact publication copy cancelled".to_string()); + } let count = reader.read(&mut buffer).map_err(|error| format!("read artifact spool: {error}"))?; if count == 0 { break; @@ -631,13 +648,16 @@ fn create_relative_file(root: &File, relative: &str, mode: u32) -> Result Result<[u8; 32], String> { +fn copy_and_hash(source: &File, destination: &File, expected_size: u64, path: &str, control: &RemoteExecutionControl) -> Result<[u8; 32], String> { let mut source = source.try_clone().map_err(|error| format!("clone artifact source {path}: {error}"))?; let mut destination = destination.try_clone().map_err(|error| format!("clone artifact spool file {path}: {error}"))?; let mut hasher = Sha256::new(); let mut copied = 0u64; let mut buffer = [0u8; COPY_BUFFER_BYTES]; loop { + if control.is_cancelled() { + return Err(format!("artifact capture cancelled: {path}")); + } let count = source.read(&mut buffer).map_err(|error| format!("read artifact {path}: {error}"))?; if count == 0 { break; diff --git a/src/bin/bunkerbox-remote.rs b/src/bin/bunkerbox-remote.rs index da8ff45..a93da01 100644 --- a/src/bin/bunkerbox-remote.rs +++ b/src/bin/bunkerbox-remote.rs @@ -1,9 +1,9 @@ -use bunkerbox::guest_install::install_remote_make_link; +use bunkerbox::guest_install::{install_remote_cargo_link, install_remote_make_link}; #[cfg(test)] use bunkerbox::remote::RemoteSnapshotId; use bunkerbox::remote_client::{ execute_remote_request_to, logical_workspace_cwd, new_request_id, remote_build_request, remote_diagnostic_sync_request, remote_environment_names, - remote_session_from_env, remote_sync_request, remote_tool_enabled, selected_remote_environment, RemoteCompletion, + remote_session_from_env, remote_sync_request, remote_tool_enabled, selected_remote_environment_for_tool, RemoteCompletion, }; #[cfg(test)] use bunkerbox::vscomm::RequestId; @@ -37,11 +37,11 @@ fn run() -> Result { env::args_os().next().and_then(|value| Path::new(&value).file_name().and_then(|name| name.to_str()).map(str::to_owned)).unwrap_or_default(); let args = env::args().skip(1).collect::>(); - if invoked_as == "make" { - return run_transparent_make(&args); + if matches!(invoked_as.as_str(), "make" | "cargo") { + return run_transparent_tool(&invoked_as, &args); } if invoked_as != "bunkerbox-remote" { - return Err("bunkerbox-remote must be invoked directly or through the managed make symlink".to_string()); + return Err("bunkerbox-remote must be invoked directly or through a managed make or cargo symlink".to_string()); } if args.len() == 1 && args[0] == "install" { install_remote_links()?; @@ -69,14 +69,14 @@ fn run_explicit(args: &[String]) -> Result { } } -fn run_transparent_make(args: &[String]) -> Result { +fn run_transparent_tool(tool: &str, args: &[String]) -> Result { let cwd = logical_workspace_cwd(&env::current_dir().map_err(|error| format!("current directory: {error}"))?)?; let session = remote_session_from_env()?; - run_build_with_sync(cwd, "make".to_string(), args.to_vec(), session) + run_build_with_sync(cwd, tool.to_string(), args.to_vec(), session) } fn run_build_with_sync(cwd: String, tool: String, args: Vec, session: WorkspaceSessionId) -> Result { - let environment = selected_remote_environment(remote_environment_names()); + let environment = selected_remote_environment_for_tool(&tool, remote_environment_names()); run_build_with_sync_using(cwd, tool, args, environment, session, execute_request_over_vsock) } @@ -140,7 +140,9 @@ fn build_request( fn install_remote_links() -> Result<(), String> { fs::create_dir_all(VSCOMM_BIN_DIR).map_err(|error| format!("mkdir {VSCOMM_BIN_DIR}: {error}"))?; let executable = env::current_exe().map_err(|error| format!("failed to locate remote binary: {error}"))?; - install_remote_make_link(Path::new(VSCOMM_BIN_DIR), &executable, remote_tool_enabled("make")) + let bin_dir = Path::new(VSCOMM_BIN_DIR); + install_remote_make_link(bin_dir, &executable, remote_tool_enabled("make"))?; + install_remote_cargo_link(bin_dir, &executable, remote_tool_enabled("cargo")) } fn connect_toolchain() -> Result { diff --git a/src/cfg.rs b/src/cfg.rs index 907ed51..3fb1807 100644 --- a/src/cfg.rs +++ b/src/cfg.rs @@ -101,6 +101,8 @@ pub struct RuntimeConfig { pub session_cleanup: Option>, #[serde(default)] pub command: Option>, + #[serde(default)] + pub remote_max_active_builds: Option, } impl RuntimeConfig { @@ -143,6 +145,15 @@ impl RuntimeConfig { self.session_mb.unwrap_or(50) } + pub fn remote_max_active_builds(&self) -> Result { + let value = self.remote_max_active_builds.unwrap_or(2); + let value = usize::try_from(value).map_err(|_| "remote_max_active_builds is too large".to_string())?; + if value == 0 || value > crate::remote::RemoteAdmissionLimits::MAX { + return Err(format!("remote_max_active_builds must be between 1 and {}", crate::remote::RemoteAdmissionLimits::MAX)); + } + Ok(value) + } + pub fn effective_session_cleanup(&self) -> Vec { let user = self.session_cleanup.as_deref().unwrap_or(&[]); DEFAULT_SESSION_CLEANUP.iter().map(|s| s.to_string()).chain(user.iter().cloned()).collect() diff --git a/src/daemon.rs b/src/daemon.rs index 65932b8..c7396b0 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -4,8 +4,8 @@ use crate::logging; use crate::loopback::{LoopbackBackend, RunRemoteSession}; use crate::proxy::{FilterProxy, UnixProxyHandle}; use crate::remote::{ - RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteEnvironmentPolicy, - RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteToolPolicy, + RemoteAdmissionLimits, RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, + RemoteEnvironmentPolicy, RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteTargetId, RemoteToolPolicy, }; use crate::remote_target::SshTarget; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; @@ -13,6 +13,7 @@ use crate::ssh::SshBackend; use crate::vscomm::{validate_exec_request, validate_process_path, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; use crate::workspace::WorkspaceCwd; use rand::Rng; +use std::collections::BTreeMap; use std::fs::File; use std::io::{BufRead, BufReader}; use std::os::fd::{AsRawFd, FromRawFd, RawFd}; @@ -22,6 +23,7 @@ use std::process::Stdio; use std::sync::{Arc, Mutex}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::process::Command; +use tokio::sync::{mpsc, OwnedSemaphorePermit, Semaphore}; const BWRAP_STATUS_FD: RawFd = 3; @@ -51,15 +53,240 @@ pub enum RemoteDispatchError { EventSinkClosed, } +#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)] +struct ExecutionKey { + session_id: crate::remote::WorkspaceSessionId, + request_id: crate::remote::RequestId, +} + +struct AdmissionPermit { + _global: OwnedSemaphorePermit, + _target: OwnedSemaphorePermit, +} + +struct RemoteAdmission { + global: Arc, + targets: Mutex>>, + target_limit: usize, +} + +impl RemoteAdmission { + fn new(limits: RemoteAdmissionLimits) -> Self { + Self { global: Arc::new(Semaphore::new(limits.global_active)), targets: Mutex::new(BTreeMap::new()), target_limit: limits.target_active } + } + + fn acquire(&self, target_id: RemoteTargetId) -> Result { + let global = self.global.clone().try_acquire_owned().map_err(|_| RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::Busy, + message: "remote global active-execution limit reached".to_string(), + })?; + let target_semaphore = self + .targets + .lock() + .map_err(|_| RemoteBackendError::Failed("remote target admission lock poisoned".to_string()))? + .entry(target_id) + .or_insert_with(|| Arc::new(Semaphore::new(self.target_limit))) + .clone(); + let target = match target_semaphore.try_acquire_owned() { + Ok(permit) => permit, + Err(_) => { + drop(global); + return Err(RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::Busy, + message: "remote target active-execution limit reached".to_string(), + }); + } + }; + Ok(AdmissionPermit { _global: global, _target: target }) + } +} + +struct ExecutionState { + cancellation_selected: bool, + terminal: bool, +} + +struct ExecutionEntry { + control: crate::remote::RemoteExecutionControl, + state: Mutex, + _admission: AdmissionPermit, +} + +impl ExecutionEntry { + fn new(control: crate::remote::RemoteExecutionControl, admission: AdmissionPermit) -> Self { + Self { control, state: Mutex::new(ExecutionState { cancellation_selected: false, terminal: false }), _admission: admission } + } + + fn request_cancel(&self) -> CancelSelection { + let mut state = match self.state.lock() { + Ok(state) => state, + Err(_) => return CancelSelection::Rejected, + }; + if state.terminal + || matches!(self.control.phase(), crate::remote::RemoteLifecyclePhase::Finalizing | crate::remote::RemoteLifecyclePhase::Terminal) + { + return CancelSelection::Rejected; + } + if state.cancellation_selected { + return CancelSelection::AlreadySelected; + } + state.cancellation_selected = true; + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Cancelling); + self.control.cancel(); + CancelSelection::Selected + } + + fn cancellation_selected(&self) -> bool { + self.state.lock().map(|state| state.cancellation_selected).unwrap_or(true) + } + + fn accept_event(&self, event: &RemoteBackendEvent) -> bool { + let mut state = match self.state.lock() { + Ok(state) => state, + Err(_) => return false, + }; + if state.terminal || (state.cancellation_selected && !matches!(event, RemoteBackendEvent::Cancelled)) { + return false; + } + if is_terminal_event(event) { + state.terminal = true; + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Terminal); + } + true + } + + fn accept_cancelled(&self) -> bool { + let mut state = match self.state.lock() { + Ok(state) => state, + Err(_) => return false, + }; + if state.terminal || !state.cancellation_selected { + return false; + } + state.terminal = true; + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Terminal); + true + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum CancelSelection { + Selected, + AlreadySelected, + Rejected, +} + +struct RemoteExecutionRegistry { + entries: Mutex>>, +} + +impl RemoteExecutionRegistry { + fn new() -> Self { + Self { entries: Mutex::new(BTreeMap::new()) } + } + + fn insert(&self, key: ExecutionKey, entry: Arc) -> Result<(), RemoteBackendError> { + let mut entries = self.entries.lock().map_err(|_| RemoteBackendError::Failed("remote execution registry lock poisoned".to_string()))?; + if entries.contains_key(&key) { + return Err(RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::Busy, + message: "remote request ID is already active for this session".to_string(), + }); + } + entries.insert(key, entry); + Ok(()) + } + + fn get(&self, key: &ExecutionKey) -> Option> { + self.entries.lock().ok().and_then(|entries| entries.get(key).cloned()) + } + + fn remove(&self, key: &ExecutionKey) { + if let Ok(mut entries) = self.entries.lock() { + entries.remove(key); + } + } + + fn cancel_all(&self) { + if let Ok(entries) = self.entries.lock() { + for entry in entries.values() { + let _ = entry.request_cancel(); + } + } + } +} + +struct ExecutionGuard { + registry: Arc, + key: ExecutionKey, + entry: Arc, + finished: bool, +} + +impl ExecutionGuard { + fn finish(mut self) { + self.finished = true; + self.registry.remove(&self.key); + } +} + +impl Drop for ExecutionGuard { + fn drop(&mut self) { + if !self.finished { + let _ = self.entry.request_cancel(); + self.registry.remove(&self.key); + } + } +} + +fn is_terminal_event(event: &RemoteBackendEvent) -> bool { + matches!( + event, + RemoteBackendEvent::SyncCompleted { .. } + | RemoteBackendEvent::Completed { .. } + | RemoteBackendEvent::Error { .. } + | RemoteBackendEvent::Cancelled + ) +} + pub struct RemoteBroker { policy: RemoteAuthorizationPolicy, context: RemoteExecutionContext, backend: Arc, + registry: Arc, + admission: Arc, + cleanup_timeout: std::time::Duration, } impl RemoteBroker { pub fn new(policy: RemoteAuthorizationPolicy, context: RemoteExecutionContext, backend: Arc) -> Self { - Self { policy, context, backend } + Self { + policy, + context, + backend, + registry: Arc::new(RemoteExecutionRegistry::new()), + admission: Arc::new(RemoteAdmission::new(RemoteAdmissionLimits::default())), + cleanup_timeout: RemoteResourcePolicy::default().cleanup_timeout, + } + } + + pub fn with_admission_limits(mut self, limits: RemoteAdmissionLimits) -> Self { + self.admission = Arc::new(RemoteAdmission::new(limits)); + self + } + + pub fn with_cleanup_timeout(mut self, cleanup_timeout: std::time::Duration) -> Self { + self.cleanup_timeout = cleanup_timeout; + self + } + + pub fn cancel_request(&self, request_id: crate::remote::RequestId) -> bool { + let key = ExecutionKey { session_id: self.context.workspace_session_id, request_id }; + self.registry.get(&key).is_some_and(|entry| !matches!(entry.request_cancel(), CancelSelection::Rejected)) + } + + pub fn cleanup_timeout(&self) -> std::time::Duration { + self.cleanup_timeout } pub async fn dispatch(&self, request: RemoteRequest, events: tokio::sync::mpsc::Sender) -> Result<(), RemoteDispatchError> { @@ -70,13 +297,138 @@ impl RemoteBroker { return Err(RemoteDispatchError::Unauthorized(error)); } }; - match self.backend.execute(authorized, events.clone()).await { - Ok(()) => Ok(()), + + if let crate::remote::RemoteOperation::Cancel { target_request_id } = authorized.request().operation() { + return self.dispatch_cancel(authorized.request_id(), *target_request_id, events).await; + } + + let request_id = authorized.request_id(); + let key = ExecutionKey { session_id: self.context.workspace_session_id, request_id }; + let admission = match self.admission.acquire(self.context.target) { + Ok(admission) => admission, Err(error) => { events.send(error.event()).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; - Err(RemoteDispatchError::Backend(error)) + return Err(RemoteDispatchError::Backend(error)); + } + }; + let control = crate::remote::RemoteExecutionControl::new(); + let entry = Arc::new(ExecutionEntry::new(control.clone(), admission)); + self.registry.insert(key, entry.clone()).map_err(RemoteDispatchError::Backend)?; + let guard = ExecutionGuard { registry: self.registry.clone(), key, entry: entry.clone(), finished: false }; + let (backend_events, mut backend_rx) = mpsc::channel(64); + let mut backend = Box::pin(self.backend.execute(authorized, control.clone(), backend_events)); + let mut backend_result = None; + let mut terminal_sent = false; + let mut sink_closed = false; + let mut terminal_deadline = Box::pin(tokio::time::sleep(self.cleanup_timeout)); + + loop { + if backend_result.is_some() && backend_rx.is_empty() { + break; + } + tokio::select! { + result = &mut backend, if backend_result.is_none() => { + backend_result = Some(result); + } + event = backend_rx.recv() => { + let Some(event) = event else { continue }; + if !entry.accept_event(&event) { + continue; + } + let terminal = is_terminal_event(&event); + if !sink_closed { + tokio::select! { + result = events.send(event) => { + if result.is_err() { + sink_closed = true; + let _ = entry.request_cancel(); + } + } + _ = control.cancelled(), if !terminal => { + sink_closed = true; + let _ = entry.request_cancel(); + } + } + } + if terminal { + terminal_sent = true; + terminal_deadline.as_mut().reset(tokio::time::Instant::now() + self.cleanup_timeout); + } + } + _ = control.cancelled(), if !terminal_sent => { + if entry.accept_cancelled() { + terminal_sent = true; + terminal_deadline.as_mut().reset(tokio::time::Instant::now() + self.cleanup_timeout); + if !sink_closed && events.send(RemoteBackendEvent::Cancelled).await.is_err() { + sink_closed = true; + } + } + } + _ = &mut terminal_deadline, if terminal_sent && backend_result.is_none() => { + backend_result = Some(Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Cleanup })); + } } } + + let result = backend_result.unwrap_or(Ok(())); + if !terminal_sent { + if entry.cancellation_selected() { + if entry.accept_cancelled() && !sink_closed && events.send(RemoteBackendEvent::Cancelled).await.is_err() { + sink_closed = true; + } + } else { + let error = match &result { + Ok(()) => RemoteBackendError::Failed("remote backend completed without a terminal event".to_string()), + Err(error) => error.clone(), + }; + let event = error.event(); + if entry.accept_event(&event) && !sink_closed && events.send(event).await.is_err() { + sink_closed = true; + } + } + } + + drop(backend); + guard.finish(); + if sink_closed { + return Err(RemoteDispatchError::EventSinkClosed); + } + match result { + Ok(()) => Ok(()), + Err(error) => Err(RemoteDispatchError::Backend(error)), + } + } + + async fn dispatch_cancel( + &self, _cancel_request_id: crate::remote::RequestId, target_request_id: crate::remote::RequestId, + events: tokio::sync::mpsc::Sender, + ) -> Result<(), RemoteDispatchError> { + let key = ExecutionKey { session_id: self.context.workspace_session_id, request_id: target_request_id }; + let Some(entry) = self.registry.get(&key) else { + let error = RemoteAuthorizationError::CancelTargetUnavailable; + events.send(error.event()).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; + return Err(RemoteDispatchError::Unauthorized(error)); + }; + if matches!(entry.control.phase(), crate::remote::RemoteLifecyclePhase::Finalizing | crate::remote::RemoteLifecyclePhase::Terminal) { + let error = RemoteAuthorizationError::CancelTargetFinalizing; + events.send(error.event()).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; + return Err(RemoteDispatchError::Unauthorized(error)); + } + match entry.request_cancel() { + CancelSelection::Selected | CancelSelection::AlreadySelected => { + events.send(RemoteBackendEvent::Completed { exit_code: 0 }).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; + Ok(()) + } + CancelSelection::Rejected => { + let error = RemoteAuthorizationError::CancelTargetFinalizing; + events.send(error.event()).await.map_err(|_| RemoteDispatchError::EventSinkClosed)?; + Err(RemoteDispatchError::Unauthorized(error)) + } + } + } + + pub fn cancel_all(&self) { + self.registry.cancel_all(); } } @@ -105,6 +457,7 @@ pub struct RemoteDaemonConfig { resources: RemoteResourcePolicy, artifact_policy: ArtifactPolicy, artifact_limits: ArtifactLimits, + admission_limits: RemoteAdmissionLimits, } enum RemoteBackendSelection { @@ -123,20 +476,29 @@ impl RemoteDaemonConfig { resources: RemoteResourcePolicy::default(), artifact_policy: ArtifactPolicy::default(), artifact_limits: ArtifactLimits::default(), + admission_limits: RemoteAdmissionLimits::default(), } } pub fn ssh(session: Arc, target: SshTarget) -> Result { crate::ssh::SshLaunchSpec::from_target(&target)?; + let target_resources = target.resources(); Ok(Self { session, allowed_tools: Vec::new(), tool_policies: None, environment: None, backend: RemoteBackendSelection::Ssh { target: Box::new(target) }, - resources: RemoteResourcePolicy::default(), + resources: RemoteResourcePolicy { + sync_timeout: target_resources.sync_timeout(), + build_timeout: target_resources.build_timeout(), + idle_output_timeout: target_resources.idle_output_timeout(), + cleanup_timeout: target_resources.cleanup_timeout(), + max_output_bytes: target_resources.max_output_bytes(), + }, artifact_policy: ArtifactPolicy::default(), artifact_limits: ArtifactLimits::default(), + admission_limits: RemoteAdmissionLimits::default(), }) } @@ -163,12 +525,19 @@ impl RemoteDaemonConfig { self.artifact_limits = limits; self } + + pub fn with_admission_limits(mut self, limits: RemoteAdmissionLimits) -> Self { + self.admission_limits = limits; + self + } } struct RemoteComponents { context: RemoteExecutionContext, policy: RemoteAuthorizationPolicy, backend: Arc, + admission_limits: RemoteAdmissionLimits, + cleanup_timeout: std::time::Duration, } impl VsockDaemon { @@ -176,7 +545,17 @@ impl VsockDaemon { passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, remote: RemoteDaemonConfig, ) -> Result { - let RemoteDaemonConfig { session, allowed_tools, tool_policies, environment, backend, resources, artifact_policy, artifact_limits } = remote; + let RemoteDaemonConfig { + session, + allowed_tools, + tool_policies, + environment, + backend, + resources, + artifact_policy, + artifact_limits, + admission_limits, + } = remote; let remote_policy = match (tool_policies, environment) { (Some(tool_policies), Some(environment)) => { RemoteAuthorizationPolicy::from_policies(session.target(), session.session_id(), tool_policies, environment)? @@ -190,15 +569,20 @@ impl VsockDaemon { RemoteBackendSelection::Loopback { tools, target_environment } => Arc::new( LoopbackBackend::new(session, tools) .with_target_environment(target_environment) - .with_timeout(resources.build_timeout) - .with_output_limit(resources.max_output_bytes) + .with_resources(resources) .with_artifacts(artifact_policy.clone(), artifact_limits), ), RemoteBackendSelection::Ssh { target } => { Arc::new(SshBackend::new(session, *target)?.with_artifacts(artifact_policy.clone(), artifact_limits)) } }; - let remote_components = RemoteComponents { context: remote_context, policy: remote_policy, backend }; + let remote_components = RemoteComponents { + context: remote_context, + policy: remote_policy, + backend, + admission_limits, + cleanup_timeout: resources.cleanup_timeout, + }; Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, remote_components) } @@ -244,7 +628,11 @@ impl VsockDaemon { } let connections = Arc::new(Mutex::new(Vec::new())); - let remote_broker = Arc::new(RemoteBroker::new(remote.policy, remote.context, remote.backend)); + let remote_broker = Arc::new( + RemoteBroker::new(remote.policy, remote.context, remote.backend) + .with_admission_limits(remote.admission_limits) + .with_cleanup_timeout(remote.cleanup_timeout), + ); let session = Arc::new(VsockSession { passthrough: Arc::new(passthrough), env_mode, @@ -316,12 +704,22 @@ async fn daemon_loop( } } + session.remote_broker.cancel_all(); let tasks = connections.lock().map_err(|_| "connection task lock poisoned".to_string())?.drain(..).collect::>(); - for task in &tasks { - task.abort(); - } - for task in tasks { - let _ = task.await; + let cleanup_deadline = session.remote_broker.cleanup_timeout(); + let mut tasks = tasks; + if tokio::time::timeout(cleanup_deadline, async { + for task in &mut tasks { + let _ = task.await; + } + }) + .await + .is_err() + { + for task in tasks { + task.abort(); + let _ = task.await; + } } Ok(()) @@ -376,7 +774,10 @@ pub async fn dispatch_remote_frame(frame: Frame, broke }; let response = crate::vscomm::RemoteEvent::from_backend_event(request_id, event).to_frame().map_err(|err| format!("encode remote event: {err}"))?; - write_frame(writer, &response).await?; + if let Err(error) = write_frame(writer, &response).await { + broker.cancel_request(request_id); + return Err(error); + } continue; } @@ -387,7 +788,10 @@ pub async fn dispatch_remote_frame(frame: Frame, broke let response = crate::vscomm::RemoteEvent::from_backend_event(request_id, event) .to_frame() .map_err(|err| format!("encode remote event: {err}"))?; - write_frame(writer, &response).await?; + if let Err(error) = write_frame(writer, &response).await { + broker.cancel_request(request_id); + return Err(error); + } } } } diff --git a/src/guest_install.rs b/src/guest_install.rs index 35d2807..2663103 100644 --- a/src/guest_install.rs +++ b/src/guest_install.rs @@ -3,11 +3,22 @@ use std::fs; use std::os::unix::fs::{symlink, PermissionsExt}; use std::path::{Path, PathBuf}; -const REMOTE_MAKE_OWNER: &str = "bunkerbox-remote"; +const REMOTE_WRAPPER_OWNER: &str = "bunkerbox-remote"; pub fn install_remote_make_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { - let target = bin_dir.join("make"); - let managed = is_managed_remote_make_link(&target); + install_remote_link("make", bin_dir, executable, enabled) +} + +pub fn install_remote_cargo_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { + install_remote_link("cargo", bin_dir, executable, enabled) +} + +fn install_remote_link(command: &str, bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { + if !is_supported_remote_command(command) { + return Err(format!("unsupported remote wrapper command: {command}")); + } + let target = bin_dir.join(command); + let managed = is_managed_remote_link(&target); if enabled { if let Ok(link) = fs::read_link(&target) { @@ -17,13 +28,13 @@ pub fn install_remote_make_link(bin_dir: &Path, executable: &Path, enabled: bool } if fs::symlink_metadata(&target).is_ok() { if !managed { - return Err(format!("cannot install remote make wrapper over existing {}", target.display())); + return Err(format!("cannot install remote {command} wrapper over existing {}", target.display())); } - fs::remove_file(&target).map_err(|error| format!("remove existing remote make wrapper: {error}"))?; + fs::remove_file(&target).map_err(|error| format!("remove existing remote {command} wrapper: {error}"))?; } - symlink(executable, &target).map_err(|error| format!("symlink remote make wrapper: {error}"))?; + symlink(executable, &target).map_err(|error| format!("symlink remote {command} wrapper: {error}"))?; } else if managed { - fs::remove_file(&target).map_err(|error| format!("remove disabled remote make wrapper: {error}"))?; + fs::remove_file(&target).map_err(|error| format!("remove disabled remote {command} wrapper: {error}"))?; } Ok(()) } @@ -36,7 +47,7 @@ pub fn install_vscomm_links(commands: impl IntoIterator, bin_dir: continue; } let target = bin_dir.join(&command); - if command == "make" && is_managed_remote_make_link(&target) { + if is_supported_remote_command(&command) && is_managed_remote_link(&target) { continue; } if command_exists_in_path_except(&command, vscomm_path, path) { @@ -56,12 +67,16 @@ pub fn install_vscomm_links(commands: impl IntoIterator, bin_dir: Ok(()) } -fn is_managed_remote_make_link(target: &Path) -> bool { +fn is_managed_remote_link(target: &Path) -> bool { let Ok(metadata) = fs::symlink_metadata(target) else { return false }; if !metadata.file_type().is_symlink() { return false; } - fs::read_link(target).ok().and_then(|link| link.file_name().map(OsStr::to_owned)).is_some_and(|name| name == REMOTE_MAKE_OWNER) + fs::read_link(target).ok().and_then(|link| link.file_name().map(OsStr::to_owned)).is_some_and(|name| name == REMOTE_WRAPPER_OWNER) +} + +fn is_supported_remote_command(command: &str) -> bool { + matches!(command, "make" | "cargo") } fn command_exists_in_path_except(command: &str, except: &Path, path: &str) -> bool { diff --git a/src/loopback.rs b/src/loopback.rs index c41ecb5..ea1b99e 100644 --- a/src/loopback.rs +++ b/src/loopback.rs @@ -1,7 +1,7 @@ use crate::artifact::{ArtifactLimits, ArtifactPolicy, ArtifactPublication, LocalArtifactSpool}; use crate::remote::{ - AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFuture, RemoteOperation, RemoteResourcePolicy, - RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, WorkspaceSessionId, + AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionControl, RemoteFuture, RemoteOperation, + RemoteResourcePolicy, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, WorkspaceSessionId, }; use crate::snapshot::{SnapshotBuilder, SnapshotEntry, SnapshotExport, SnapshotHandle, SnapshotStore}; use rand::RngCore; @@ -15,7 +15,7 @@ use std::sync::{Arc, Mutex}; use std::time::Duration; use tokio::io::AsyncReadExt; use tokio::process::Command; -use tokio::sync::mpsc; +use tokio::sync::{mpsc, Notify}; use tokio::time::sleep; pub const LOOPBACK_BUILD_TIMEOUT: Duration = crate::remote::DEFAULT_REMOTE_BUILD_TIMEOUT; @@ -95,16 +95,32 @@ impl RunRemoteSession { } fn sync_snapshot_for_request(&self, retain_capability: bool) -> Result, String> { + self.sync_snapshot_for_request_with_control(retain_capability, RemoteExecutionControl::new()) + } + + pub(crate) fn sync_snapshot_for_request_with_control( + &self, retain_capability: bool, control: RemoteExecutionControl, + ) -> Result, String> { let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; - let snapshot = self.snapshot_builder.build_root(&self.workspace_root, self.session_id)?; + let snapshot = self.snapshot_builder.build_root_with_control(&self.workspace_root, self.session_id, control.clone())?; let handle = snapshot.handle().clone(); if !retain_capability { self.discard_unclaimed_snapshot(&handle)?; return Ok(None); } + if control.is_cancelled() { + self.discard_unclaimed_snapshot(&handle)?; + return Err("snapshot synchronization cancelled".to_string()); + } match self.register_snapshot(handle.clone()) { - Ok(snapshot_id) => Ok(Some(snapshot_id)), + Ok(snapshot_id) => { + if control.is_cancelled() { + self.abort_snapshot_capability(snapshot_id)?; + return Err("snapshot synchronization cancelled".to_string()); + } + Ok(Some(snapshot_id)) + } Err(error) => { let cleanup = self.discard_unclaimed_snapshot(&handle); if let Err(cleanup_error) = cleanup { @@ -204,9 +220,9 @@ impl RunRemoteSession { } } - fn materialize_snapshot(&self, handle: &SnapshotHandle, destination: &Path) -> Result<(), String> { + fn materialize_snapshot_with_control(&self, handle: &SnapshotHandle, destination: &Path, control: RemoteExecutionControl) -> Result<(), String> { let _operation = self.snapshot_operation.lock().map_err(|_| "remote session operation lock poisoned".to_string())?; - self.snapshot_store.materialize(handle, destination).map(|_| ()) + self.snapshot_store.materialize_with_control(handle, destination, control).map(|_| ()) } fn new_job_path(&self) -> Result { @@ -277,6 +293,12 @@ impl SnapshotExportClaim { pub(crate) fn read_file_bounded(&self, entry: &SnapshotEntry, max_bytes: u64) -> Result, String> { self.snapshot.read_file_bounded(entry, max_bytes) } + + pub(crate) fn read_file_bounded_with_control( + &self, entry: &SnapshotEntry, max_bytes: u64, control: &RemoteExecutionControl, + ) -> Result, String> { + self.snapshot.read_file_bounded_with_control(entry, max_bytes, control) + } } impl Drop for SnapshotExportClaim { @@ -340,6 +362,11 @@ impl LoopbackBackend { self } + pub fn with_resources(mut self, resources: RemoteResourcePolicy) -> Self { + self.resources = resources; + self + } + pub fn with_target_environment(mut self, environment: BTreeMap) -> Self { let mut trusted = trusted_target_environment(); trusted.extend(environment); @@ -356,7 +383,7 @@ impl LoopbackBackend { impl RemoteBackend for LoopbackBackend { fn execute<'a>( - &'a self, request: AuthorizedRemoteRequest, events: mpsc::Sender, + &'a self, request: AuthorizedRemoteRequest, control: RemoteExecutionControl, events: mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { let session = self.session.clone(); let tools = self.tools.clone(); @@ -368,13 +395,20 @@ impl RemoteBackend for LoopbackBackend { let options = LoopbackBuildOptions { resources, artifact_policy, artifact_limits, request_id }; Box::pin(async move { match request.request().operation() { - RemoteOperation::Sync(sync) => execute_sync(session, sync.retain_capability(), events).await, - RemoteOperation::Build(build) => execute_build(session, tools, target_environment, options, build, events).await, + RemoteOperation::Sync(sync) => execute_sync(session, sync.retain_capability(), resources, control, events).await, + RemoteOperation::Build(build) => execute_build(session, tools, target_environment, options, build, control, events).await, + RemoteOperation::Cancel { .. } => Err(RemoteBackendError::Failed("cancel is handled by the remote broker".to_string())), } }) } } +impl LoopbackBackend { + pub async fn execute(&self, request: AuthorizedRemoteRequest, events: mpsc::Sender) -> Result<(), RemoteBackendError> { + ::execute(self, request, RemoteExecutionControl::new(), events).await + } +} + #[derive(Clone)] struct LoopbackBuildOptions { resources: RemoteResourcePolicy, @@ -384,12 +418,27 @@ struct LoopbackBuildOptions { } async fn execute_sync( - session: Arc, retain_capability: bool, events: mpsc::Sender, + session: Arc, retain_capability: bool, resources: RemoteResourcePolicy, control: RemoteExecutionControl, + events: mpsc::Sender, ) -> Result<(), RemoteBackendError> { + control.set_phase(crate::remote::RemoteLifecyclePhase::Syncing); send_event(&events, RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }).await?; - let sync = tokio::task::spawn_blocking(move || session.sync_snapshot_for_request(retain_capability)) - .await - .map_err(|error| RemoteBackendError::Failed(format!("snapshot worker failed: {error}")))?; + let snapshot_control = control.clone(); + let mut sync_task = tokio::task::spawn_blocking(move || session.sync_snapshot_for_request_with_control(retain_capability, snapshot_control)); + let sync = tokio::select! { + result = tokio::time::timeout(resources.sync_timeout, &mut sync_task) => match result { + Ok(result) => result.map_err(|error| RemoteBackendError::Failed(format!("snapshot worker failed: {error}")))?, + Err(_) => { + control.cancel(); + let _ = tokio::time::timeout(resources.cleanup_timeout, &mut sync_task).await; + return Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Sync }); + } + }, + _ = control.cancelled() => { + let _ = tokio::time::timeout(resources.cleanup_timeout, &mut sync_task).await; + return Err(RemoteBackendError::Cancelled); + }, + }; match sync.map_err(RemoteBackendError::Failed)? { Some(snapshot_id) => send_event(&events, RemoteBackendEvent::SyncCompleted { snapshot_id }).await, None => send_event(&events, RemoteBackendEvent::Completed { exit_code: 0 }).await, @@ -398,29 +447,47 @@ async fn execute_sync( async fn execute_build( session: Arc, tools: Arc>, target_environment: Arc>, - options: LoopbackBuildOptions, build: &crate::remote::RemoteBuild, events: mpsc::Sender, + options: LoopbackBuildOptions, build: &crate::remote::RemoteBuild, control: RemoteExecutionControl, events: mpsc::Sender, ) -> Result<(), RemoteBackendError> { let LoopbackBuildOptions { resources, artifact_policy, artifact_limits, request_id } = options; + control.set_phase(crate::remote::RemoteLifecyclePhase::Transferring); + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } let snapshot = session.claim_snapshot(build.snapshot_id()).map_err(RemoteBackendError::Failed)?; let executable = tools .get(build.tool().as_str()) .ok_or_else(|| RemoteBackendError::Spawn(format!("loopback tool is not configured: {}", build.tool().as_str())))?; let job_path = session.new_job_path().map_err(RemoteBackendError::Failed)?; - let _job = JobGuard { path: job_path.clone() }; + let mut job = JobGuard { path: Some(job_path.clone()) }; let materialization_claim = snapshot.clone_for_worker().map_err(RemoteBackendError::Failed)?; let destination = job_path.clone(); let snapshot_handle = snapshot.handle().clone(); - tokio::task::spawn_blocking({ + let materialize_control = control.clone(); + let mut materialize = tokio::task::spawn_blocking({ let session = session.clone(); move || { - let result = session.materialize_snapshot(&snapshot_handle, &destination); + let result = session.materialize_snapshot_with_control(&snapshot_handle, &destination, materialize_control); drop(materialization_claim); result } - }) - .await - .map_err(|error| RemoteBackendError::Failed(format!("materialization worker failed: {error}")))? - .map_err(RemoteBackendError::Failed)?; + }); + tokio::select! { + result = tokio::time::timeout(resources.sync_timeout, &mut materialize) => match result { + Ok(result) => result + .map_err(|error| RemoteBackendError::Failed(format!("materialization worker failed: {error}")))? + .map_err(RemoteBackendError::Failed)?, + Err(_) => { + control.cancel(); + let _ = tokio::time::timeout(resources.cleanup_timeout, &mut materialize).await; + return Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Sync }); + } + }, + _ = control.cancelled() => { + let _ = tokio::time::timeout(resources.cleanup_timeout, &mut materialize).await; + return Err(RemoteBackendError::Cancelled); + }, + } let cwd = job_path.join(build.cwd().as_str()); let cwd_metadata = fs::symlink_metadata(&cwd).map_err(|error| RemoteBackendError::Failed(format!("remote cwd is unavailable: {error}")))?; @@ -452,14 +519,33 @@ async fn execute_build( } command.kill_on_drop(true); let mut child = command.spawn().map_err(|error| RemoteBackendError::Spawn(format!("spawn loopback tool: {error}")))?; + control.set_phase(crate::remote::RemoteLifecyclePhase::Building); let process_group = child.id().map(|pid| ProcessGroupGuard { pgid: pid as i32, active: true }); let stdout = child.stdout.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stdout".to_string()))?; let stderr = child.stderr.take().ok_or_else(|| RemoteBackendError::Failed("loopback child has no stderr".to_string()))?; let output_bytes = Arc::new(AtomicU64::new(0)); - let mut stdout_task = tokio::spawn(pump(stdout, RemoteStream::Stdout, events.clone(), output_bytes.clone(), resources.max_output_bytes)); - let mut stderr_task = tokio::spawn(pump(stderr, RemoteStream::Stderr, events.clone(), output_bytes, resources.max_output_bytes)); + let output_activity = Arc::new(Notify::new()); + let mut stdout_task = tokio::spawn(pump( + stdout, + RemoteStream::Stdout, + events.clone(), + output_bytes.clone(), + output_activity.clone(), + control.clone(), + resources.max_output_bytes, + )); + let mut stderr_task = tokio::spawn(pump( + stderr, + RemoteStream::Stderr, + events.clone(), + output_bytes, + output_activity.clone(), + control.clone(), + resources.max_output_bytes, + )); let mut child_wait = Box::pin(child.wait()); let mut timeout_sleep = Box::pin(sleep(resources.build_timeout)); + let mut idle_output_sleep = Box::pin(sleep(resources.idle_output_timeout)); let mut post_exit_drain_sleep = Box::pin(sleep(POST_CHILD_EXIT_DRAIN_TIMEOUT)); let mut final_drain_sleep = Box::pin(sleep(POST_CHILD_EXIT_FINAL_DRAIN_TIMEOUT)); let mut child_status = None; @@ -516,7 +602,20 @@ async fn execute_build( } } _ = &mut timeout_sleep, if child_status.is_none() => { - failure.get_or_insert(RemoteBackendError::Timeout); + failure.get_or_insert(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Build }); + kill_process_group(process_group.as_ref()); + group_killed = true; + } + _ = &mut idle_output_sleep, if child_status.is_none() => { + failure.get_or_insert(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::IdleOutput }); + kill_process_group(process_group.as_ref()); + group_killed = true; + } + _ = output_activity.notified(), if child_status.is_none() => { + idle_output_sleep.as_mut().reset(tokio::time::Instant::now() + resources.idle_output_timeout); + } + _ = control.cancelled(), if child_status.is_none() || !stdout_done || !stderr_done => { + failure.get_or_insert(RemoteBackendError::Cancelled); kill_process_group(process_group.as_ref()); group_killed = true; } @@ -548,51 +647,79 @@ async fn execute_build( process_group.active = false; } if let Some(error) = failure { - return Err(error); - } - let status = child_status.unwrap()?; + control.set_phase(crate::remote::RemoteLifecyclePhase::Cleanup); + let cleanup_error = job.cleanup(resources.cleanup_timeout).await.err(); + return Err(cleanup_error.unwrap_or(error)); + } + let status = match child_status.unwrap() { + Ok(status) => status, + Err(error) => { + control.set_phase(crate::remote::RemoteLifecyclePhase::Cleanup); + let cleanup_error = job.cleanup(resources.cleanup_timeout).await.err(); + return Err(cleanup_error.unwrap_or(error)); + } + }; let exit_code = status.code().unwrap_or(-1); if exit_code == 0 && artifact_policy.is_enabled() { + control.set_phase(crate::remote::RemoteLifecyclePhase::ArtifactHandling); let job_root = job_path.clone(); let workspace_root = session.workspace_root().to_path_buf(); let spool_parent = session.jobs_root().to_path_buf(); - let policy = artifact_policy; + let policy = artifact_policy.clone(); let limits = artifact_limits; - let retrieval = tokio::time::timeout( - limits.timeout, - tokio::task::spawn_blocking(move || { - let spool = LocalArtifactSpool::capture(&job_root, &spool_parent, &policy, limits) + let artifact_control = control.clone(); + let artifact_result = tokio::select! { + result = tokio::time::timeout( + limits.timeout, + tokio::task::spawn_blocking(move || { + let control = artifact_control; + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } + let spool = LocalArtifactSpool::capture_with_control(&job_root, &spool_parent, &policy, limits, &control) .map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactManifest, message: error })?; let manifest = spool.manifest().clone(); let mut publication = ArtifactPublication::new(&workspace_root, request_id, manifest) .map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactTransfer, message: error })?; for index in 0..spool.manifest().entries().len() { + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } let mut source = spool.open_entry(index).map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactTransfer, message: error, })?; - publication.copy_from_reader(index, &mut source).map_err(|error| RemoteBackendError::Transport { + publication.copy_from_reader_with_control(index, &mut source, &control).map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactTransfer, message: error, })?; } + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } publication .publish() .map_err(|error| RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::ArtifactTransfer, message: error }) - }), - ) - .await; - match retrieval { - Ok(Ok(result)) => result?, - Ok(Err(error)) => return Err(RemoteBackendError::Failed(format!("artifact worker failed: {error}"))), - Err(_) => { - return Err(RemoteBackendError::Transport { - class: crate::remote::RemoteFailureClass::ArtifactTransfer, - message: "artifact retrieval timed out".to_string(), - }) - } + }), + ) => match result { + Ok(Ok(result)) => result, + Ok(Err(error)) => Err(RemoteBackendError::Failed(format!("artifact worker failed: {error}"))), + Err(_) => { + control.cancel(); + Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Artifact }) + } + }, + _ = control.cancelled() => Err(RemoteBackendError::Cancelled), + }; + if let Err(error) = artifact_result { + control.set_phase(crate::remote::RemoteLifecyclePhase::Cleanup); + let cleanup_error = job.cleanup(resources.cleanup_timeout).await.err(); + return Err(if matches!(error, RemoteBackendError::Cancelled) { error } else { cleanup_error.unwrap_or(error) }); } } + control.set_phase(crate::remote::RemoteLifecyclePhase::Cleanup); + job.cleanup(resources.cleanup_timeout).await?; + control.set_phase(crate::remote::RemoteLifecyclePhase::Finalizing); send_event(&events, RemoteBackendEvent::Completed { exit_code }).await } @@ -603,11 +730,15 @@ enum RemoteStream { } async fn pump( - mut reader: R, stream: RemoteStream, events: mpsc::Sender, output_bytes: Arc, max_output_bytes: u64, + mut reader: R, stream: RemoteStream, events: mpsc::Sender, output_bytes: Arc, output_activity: Arc, + control: RemoteExecutionControl, max_output_bytes: u64, ) -> Result<(), RemoteBackendError> { let mut buffer = [0u8; 8192]; loop { - let count = reader.read(&mut buffer).await.map_err(|error| RemoteBackendError::Failed(format!("read loopback output: {error}")))?; + let count = tokio::select! { + result = reader.read(&mut buffer) => result.map_err(|error| RemoteBackendError::Failed(format!("read loopback output: {error}")))?, + _ = control.cancelled() => return Err(RemoteBackendError::Cancelled), + }; if count == 0 { return Ok(()); } @@ -615,6 +746,7 @@ async fn pump( if total > max_output_bytes { return Err(RemoteBackendError::OutputLimit { limit: max_output_bytes }); } + output_activity.notify_waiters(); let event = match stream { RemoteStream::Stdout => RemoteBackendEvent::Stdout(buffer[..count].to_vec()), RemoteStream::Stderr => RemoteBackendEvent::Stderr(buffer[..count].to_vec()), @@ -648,12 +780,31 @@ async fn send_event(events: &mpsc::Sender, event: RemoteBack } struct JobGuard { - path: PathBuf, + path: Option, +} + +impl JobGuard { + async fn cleanup(&mut self, cleanup_timeout: Duration) -> Result<(), RemoteBackendError> { + let Some(path) = self.path.take() else { return Ok(()) }; + let result = tokio::time::timeout(cleanup_timeout, tokio::task::spawn_blocking(move || fs::remove_dir_all(path))).await; + match result { + Ok(Ok(Ok(()))) => Ok(()), + Ok(Ok(Err(error))) if error.kind() == std::io::ErrorKind::NotFound => Ok(()), + Ok(Ok(Err(error))) => Err(RemoteBackendError::Transport { + class: crate::remote::RemoteFailureClass::Cleanup, + message: format!("remove loopback job: {error}"), + }), + Ok(Err(error)) => Err(RemoteBackendError::Failed(format!("loopback cleanup worker failed: {error}"))), + Err(_) => Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Cleanup }), + } + } } impl Drop for JobGuard { fn drop(&mut self) { - let _ = fs::remove_dir_all(&self.path); + if let Some(path) = self.path.take() { + let _ = fs::remove_dir_all(path); + } } } diff --git a/src/main.rs b/src/main.rs index a436c3a..e6aa1ad 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,5 +1,5 @@ use bunkerbox::cfg::{ProjectConfig, RemoteToolSpec, WorkspaceMode}; -use bunkerbox::remote::{RemoteEnvironmentPolicy, RemoteToolPolicy}; +use bunkerbox::remote::{RemoteAdmissionLimits, RemoteEnvironmentPolicy, RemoteToolPolicy}; use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, loopback, overlay, remote_target, snapshot, tui, vscomm, workspace}; use rand::RngCore; use std::ffi::OsString; @@ -354,6 +354,8 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option daemon::RemoteDaemonConfig::loopback(session.clone(), Vec::new(), tools), remote_target::BackendMode::Ssh => { @@ -368,7 +370,10 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option bunkerbox::remote::RemoteTargetId { } fn remote_tool_names(entries: &[RemoteToolSpec]) -> Vec { - entries.iter().filter(|tool| tool.name == "make").map(|tool| tool.name.clone()).collect() + entries.iter().filter(|tool| matches!(tool.name.as_str(), "make" | "cargo")).map(|tool| tool.name.clone()).collect() } fn write_run_handoff(file: &mut File, path: &Path, session_id: vscomm::WorkspaceSessionId) -> Result<(), String> { diff --git a/src/remote.rs b/src/remote.rs index b8fc6e3..13c8ea5 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -3,9 +3,11 @@ use std::collections::{BTreeMap, BTreeSet}; use std::future::Future; use std::pin::Pin; +use std::sync::atomic::{AtomicBool, AtomicU8, Ordering}; use std::sync::Arc; use std::time::Duration; use tokio::sync::mpsc; +use tokio::sync::Notify; pub const MAX_REMOTE_STRING_BYTES: usize = 4 * 1024; pub const MAX_REMOTE_TOOL_BYTES: usize = 256; @@ -18,9 +20,15 @@ pub const DEFAULT_REMOTE_BUILD_TIMEOUT: Duration = Duration::from_secs(30); pub const DEFAULT_REMOTE_OUTPUT_BYTES: u64 = 64 * 1024 * 1024; pub const DEFAULT_REMOTE_ENVIRONMENT: &[&str] = &["CC", "CXX", "AR", "RUSTFLAGS", "CFLAGS", "CXXFLAGS", "MAKEFLAGS"]; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct RequestId(pub [u8; 16]); +impl RequestId { + pub fn is_zero(self) -> bool { + self.0 == [0; 16] + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct RemoteSnapshotId([u8; 16]); @@ -38,10 +46,10 @@ impl RemoteSnapshotId { } } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct WorkspaceSessionId(pub [u8; 16]); -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct RemoteTargetId(pub [u8; 16]); #[derive(Debug, Clone, PartialEq, Eq)] @@ -138,13 +146,22 @@ impl RemoteBuild { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RemoteResourcePolicy { + pub sync_timeout: Duration, pub build_timeout: Duration, + pub idle_output_timeout: Duration, + pub cleanup_timeout: Duration, pub max_output_bytes: u64, } impl Default for RemoteResourcePolicy { fn default() -> Self { - Self { build_timeout: DEFAULT_REMOTE_BUILD_TIMEOUT, max_output_bytes: DEFAULT_REMOTE_OUTPUT_BYTES } + Self { + sync_timeout: Duration::from_secs(60), + build_timeout: DEFAULT_REMOTE_BUILD_TIMEOUT, + idle_output_timeout: Duration::from_secs(5 * 60), + cleanup_timeout: Duration::from_secs(5), + max_output_bytes: DEFAULT_REMOTE_OUTPUT_BYTES, + } } } @@ -231,10 +248,168 @@ impl RemoteSync { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RemoteAdmissionLimits { + pub global_active: usize, + pub target_active: usize, +} + +impl Default for RemoteAdmissionLimits { + fn default() -> Self { + Self { global_active: 2, target_active: 1 } + } +} + +impl RemoteAdmissionLimits { + pub const MAX: usize = 64; + + pub fn new(global_active: usize, target_active: usize) -> Result { + if global_active == 0 || target_active == 0 { + return Err("remote active-build limits must be positive".to_string()); + } + if global_active > Self::MAX || target_active > Self::MAX { + return Err(format!("remote active-build limits must not exceed {}", Self::MAX)); + } + Ok(Self { global_active, target_active }) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RemoteLifecyclePhase { + Admitting, + Syncing, + Connecting, + Transferring, + Building, + ArtifactHandling, + Cancelling, + Cleanup, + Finalizing, + Terminal, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RemoteTimeoutCause { + Connection, + Sync, + Build, + IdleOutput, + Artifact, + Cleanup, +} + +impl RemoteTimeoutCause { + pub const fn as_str(self) -> &'static str { + match self { + Self::Connection => "connection", + Self::Sync => "sync", + Self::Build => "build", + Self::IdleOutput => "idle output", + Self::Artifact => "artifact", + Self::Cleanup => "cleanup", + } + } +} + +#[derive(Clone)] +pub struct RemoteCancellation { + cancelled: Arc, + notify: Arc, +} + +impl Default for RemoteCancellation { + fn default() -> Self { + Self { cancelled: Arc::new(AtomicBool::new(false)), notify: Arc::new(Notify::new()) } + } +} + +impl RemoteCancellation { + pub fn cancel(&self) -> bool { + if self.cancelled.swap(true, Ordering::AcqRel) { + false + } else { + self.notify.notify_waiters(); + true + } + } + + pub fn is_cancelled(&self) -> bool { + self.cancelled.load(Ordering::Acquire) + } + + pub async fn cancelled(&self) { + while !self.is_cancelled() { + self.notify.notified().await; + } + } +} + +#[derive(Clone)] +pub struct RemoteExecutionControl { + cancellation: RemoteCancellation, + phase: Arc, +} + +impl Default for RemoteExecutionControl { + fn default() -> Self { + Self::new() + } +} + +impl RemoteExecutionControl { + pub fn new() -> Self { + Self { cancellation: RemoteCancellation::default(), phase: Arc::new(AtomicU8::new(0)) } + } + + pub fn cancel(&self) -> bool { + self.cancellation.cancel() + } + + pub fn is_cancelled(&self) -> bool { + self.cancellation.is_cancelled() + } + + pub async fn cancelled(&self) { + self.cancellation.cancelled().await; + } + + pub fn phase(&self) -> RemoteLifecyclePhase { + match self.phase.load(Ordering::Acquire) { + 1 => RemoteLifecyclePhase::Syncing, + 2 => RemoteLifecyclePhase::Connecting, + 3 => RemoteLifecyclePhase::Transferring, + 4 => RemoteLifecyclePhase::Building, + 5 => RemoteLifecyclePhase::ArtifactHandling, + 6 => RemoteLifecyclePhase::Cancelling, + 7 => RemoteLifecyclePhase::Cleanup, + 8 => RemoteLifecyclePhase::Finalizing, + 9 => RemoteLifecyclePhase::Terminal, + _ => RemoteLifecyclePhase::Admitting, + } + } + + pub fn set_phase(&self, phase: RemoteLifecyclePhase) { + let value = match phase { + RemoteLifecyclePhase::Admitting => 0, + RemoteLifecyclePhase::Syncing => 1, + RemoteLifecyclePhase::Connecting => 2, + RemoteLifecyclePhase::Transferring => 3, + RemoteLifecyclePhase::Building => 4, + RemoteLifecyclePhase::ArtifactHandling => 5, + RemoteLifecyclePhase::Cancelling => 6, + RemoteLifecyclePhase::Cleanup => 7, + RemoteLifecyclePhase::Finalizing => 8, + RemoteLifecyclePhase::Terminal => 9, + }; + self.phase.fetch_max(value, Ordering::AcqRel); + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum RemoteOperation { Sync(RemoteSync), Build(RemoteBuild), + Cancel { target_request_id: RequestId }, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -261,6 +436,10 @@ impl RemoteRequest { Self { request_id, workspace_session_id, operation: RemoteOperation::Build(build) } } + pub fn cancel(request_id: RequestId, workspace_session_id: WorkspaceSessionId, target_request_id: RequestId) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Cancel { target_request_id } } + } + pub fn request_id(&self) -> RequestId { self.request_id } @@ -302,6 +481,10 @@ impl AuthorizedRemoteRequest { #[derive(Debug, Clone, PartialEq, Eq)] pub enum RemoteAuthorizationError { + InvalidRequestId, + InvalidCancelTarget, + CancelTargetUnavailable, + CancelTargetFinalizing, SessionMismatch, TargetNotAllowed, ToolNotAllowed(String), @@ -362,6 +545,9 @@ impl RemoteAuthorizationPolicy { } pub fn authorize(&self, context: &RemoteExecutionContext, request: RemoteRequest) -> Result { + if request.request_id.is_zero() { + return Err(RemoteAuthorizationError::InvalidRequestId); + } if context.workspace_session_id != self.allowed_session || request.workspace_session_id != context.workspace_session_id { return Err(RemoteAuthorizationError::SessionMismatch); } @@ -371,6 +557,12 @@ impl RemoteAuthorizationPolicy { let request = match request.operation() { RemoteOperation::Sync(_) => request, + RemoteOperation::Cancel { target_request_id } => { + if target_request_id.is_zero() { + return Err(RemoteAuthorizationError::InvalidCancelTarget); + } + request + } RemoteOperation::Build(build) => { if self.snapshot_authority.as_ref().is_none_or(|authority| !authority.snapshot_available(self.allowed_session, build.snapshot_id())) { return Err(RemoteAuthorizationError::SnapshotNotAllowed); @@ -381,7 +573,14 @@ impl RemoteAuthorizationPolicy { if !tool_policy.allows_arbitrary_argv() && !build.argv().is_empty() { return Err(RemoteAuthorizationError::ToolArgumentsNotAllowed(build.tool().as_str().to_string())); } - let environment = self.environment.filter(build.env())?; + let environment = if build.tool().as_str() == "cargo" { + if let Some((name, _)) = build.env().first() { + return Err(RemoteAuthorizationError::ForbiddenEnvironment(name.clone())); + } + Vec::new() + } else { + self.environment.filter(build.env())? + }; let filtered = RemoteBuild::new(build.cwd.clone(), build.tool.clone(), build.argv.clone(), environment, build.snapshot_id()) .map_err(RemoteAuthorizationError::InvalidEnvironment)?; RemoteRequest::build(request.request_id, request.workspace_session_id, filtered) @@ -408,6 +607,7 @@ pub enum RemoteBackendError { Failed(String), Spawn(String), Transport { class: RemoteFailureClass, message: String }, + Deadline { cause: RemoteTimeoutCause }, Timeout, OutputLimit { limit: u64 }, Cancelled, @@ -425,6 +625,7 @@ pub enum RemoteFailureClass { SnapshotTransfer, ArtifactManifest, ArtifactTransfer, + Busy, Disconnect, Cleanup, } @@ -442,6 +643,7 @@ impl RemoteFailureClass { Self::SnapshotTransfer => "snapshot transfer", Self::ArtifactManifest => "artifact manifest", Self::ArtifactTransfer => "artifact transfer", + Self::Busy => "busy", Self::Disconnect => "disconnect", Self::Cleanup => "cleanup", } @@ -453,6 +655,7 @@ impl RemoteBackendError { match self { Self::Failed(message) | Self::Spawn(message) => RemoteBackendEvent::Error { message: message.clone() }, Self::Transport { class, message } => RemoteBackendEvent::Error { message: format!("remote {} failure: {message}", class.as_str()) }, + Self::Deadline { cause } => RemoteBackendEvent::Error { message: format!("remote {} deadline exceeded", cause.as_str()) }, Self::Timeout => RemoteBackendEvent::Error { message: "remote backend timed out".to_string() }, Self::OutputLimit { limit } => RemoteBackendEvent::Error { message: format!("remote output exceeded limit of {limit} bytes") }, Self::Cancelled => RemoteBackendEvent::Cancelled, @@ -464,7 +667,7 @@ pub type RemoteFuture<'a, T> = Pin + Send + 'a>>; pub trait RemoteBackend: Send + Sync { fn execute<'a>( - &'a self, request: AuthorizedRemoteRequest, events: mpsc::Sender, + &'a self, request: AuthorizedRemoteRequest, control: RemoteExecutionControl, events: mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>>; } diff --git a/src/remote_client.rs b/src/remote_client.rs index c67a281..cb026e1 100644 --- a/src/remote_client.rs +++ b/src/remote_client.rs @@ -51,6 +51,14 @@ pub fn selected_remote_environment(names: impl IntoIterator) -> V .collect() } +pub fn selected_remote_environment_for_tool(tool: &str, names: impl IntoIterator) -> Vec<(String, String)> { + if tool == "cargo" { + Vec::new() + } else { + selected_remote_environment(names) + } +} + fn never_forward_environment(name: &str) -> bool { let upper = name.to_ascii_uppercase(); matches!(upper.as_str(), "PATH" | "HOME" | "SSH_AUTH_SOCK" | "SSH_AGENT_PID" | "GITHUB_TOKEN" | "GITLAB_TOKEN" | "NPM_TOKEN" | "KUBECONFIG") @@ -83,6 +91,10 @@ pub fn remote_build_request( Ok(RemoteRequest::build(request_id, session_id, build)) } +pub fn remote_cancel_request(request_id: RequestId, session_id: WorkspaceSessionId, target_request_id: RequestId) -> RemoteRequest { + RemoteRequest::cancel(request_id, session_id, target_request_id) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum RemoteCompletion { Synced(RemoteSnapshotId), diff --git a/src/remote_target.rs b/src/remote_target.rs index 5cc0ef7..d39b54a 100644 --- a/src/remote_target.rs +++ b/src/remote_target.rs @@ -40,8 +40,68 @@ pub struct ResourceLimits { pub connect_timeout: Duration, pub sync_timeout: Duration, pub build_timeout: Duration, + pub idle_output_timeout: Duration, + pub cleanup_timeout: Duration, pub max_output_bytes: u64, + pub max_active_builds: usize, pub artifact: ArtifactLimits, + pub worker: WorkerStateLimits, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct WorkerStateLimits { + pub max_uploads: usize, + pub max_upload_bytes: u64, + pub max_jobs: usize, + pub max_job_bytes: u64, + pub max_artifact_spools: usize, + pub max_artifact_spool_bytes: u64, + pub max_state_entries: usize, +} + +impl Default for WorkerStateLimits { + fn default() -> Self { + Self { + max_uploads: 2, + max_upload_bytes: 1024 * 1024 * 1024, + max_jobs: 1, + max_job_bytes: 512 * 1024 * 1024, + max_artifact_spools: 1, + max_artifact_spool_bytes: 512 * 1024 * 1024, + max_state_entries: 20_000, + } + } +} + +impl WorkerStateLimits { + pub const MAX_COUNT: usize = 1_000; + pub const MAX_BYTES: u64 = 16 * 1024 * 1024 * 1024; + pub const MAX_ENTRIES: usize = 1_000_000; + + pub fn new( + max_uploads: usize, max_upload_bytes: u64, max_jobs: usize, max_job_bytes: u64, max_artifact_spools: usize, max_artifact_spool_bytes: u64, + max_state_entries: usize, + ) -> Result { + if max_uploads == 0 || max_jobs == 0 || max_artifact_spools == 0 { + return Err("worker state counts must be positive".to_string()); + } + if max_uploads > Self::MAX_COUNT || max_jobs > Self::MAX_COUNT || max_artifact_spools > Self::MAX_COUNT { + return Err(format!("worker state counts must not exceed {}", Self::MAX_COUNT)); + } + if max_upload_bytes == 0 + || max_job_bytes == 0 + || max_artifact_spool_bytes == 0 + || max_upload_bytes > Self::MAX_BYTES + || max_job_bytes > Self::MAX_BYTES + || max_artifact_spool_bytes > Self::MAX_BYTES + { + return Err(format!("worker state byte limits must be between 1 and {}", Self::MAX_BYTES)); + } + if max_state_entries == 0 || max_state_entries > Self::MAX_ENTRIES { + return Err(format!("worker state entry limit must be between 1 and {}", Self::MAX_ENTRIES)); + } + Ok(Self { max_uploads, max_upload_bytes, max_jobs, max_job_bytes, max_artifact_spools, max_artifact_spool_bytes, max_state_entries }) + } } impl ResourceLimits { @@ -57,6 +117,14 @@ impl ResourceLimits { self.build_timeout } + pub fn idle_output_timeout(&self) -> Duration { + self.idle_output_timeout + } + + pub fn cleanup_timeout(&self) -> Duration { + self.cleanup_timeout + } + pub fn max_output(&self) -> u64 { self.max_output_bytes } @@ -65,9 +133,17 @@ impl ResourceLimits { self.max_output_bytes } + pub fn max_active_builds(&self) -> usize { + self.max_active_builds + } + pub fn artifact_limits(&self) -> ArtifactLimits { self.artifact } + + pub fn worker_state_limits(&self) -> WorkerStateLimits { + self.worker + } } /// An SSH target after all configuration and local-file checks have passed. @@ -526,6 +602,10 @@ struct RawResources { build_timeout: RawQuantity, #[serde(rename = "max-output-bytes", alias = "max-output", alias = "max_output_bytes", alias = "max_output")] max_output: RawQuantity, + #[serde(default, rename = "idle-output-timeout-seconds", alias = "idle-output-timeout", alias = "idle_output_timeout")] + idle_output_timeout: Option, + #[serde(default, rename = "cleanup-timeout-seconds", alias = "cleanup-timeout", alias = "cleanup_timeout")] + cleanup_timeout: Option, #[serde(default, rename = "artifact-timeout-seconds", alias = "artifact-timeout", alias = "artifact_timeout")] artifact_timeout: Option, #[serde(default, rename = "max-artifact-bytes", alias = "max-artifact-bytes-per-file", alias = "max_artifact_bytes")] @@ -534,6 +614,22 @@ struct RawResources { max_artifact_total_bytes: Option, #[serde(default, rename = "max-artifact-entries", alias = "max_artifact_entries")] max_artifact_entries: Option, + #[serde(default, rename = "max-worker-uploads", alias = "max_worker_uploads")] + max_worker_uploads: Option, + #[serde(default, rename = "max-worker-upload-bytes", alias = "max_worker_upload_bytes")] + max_worker_upload_bytes: Option, + #[serde(default, rename = "max-worker-jobs", alias = "max_worker_jobs")] + max_worker_jobs: Option, + #[serde(default, rename = "max-worker-job-bytes", alias = "max_worker_job_bytes")] + max_worker_job_bytes: Option, + #[serde(default, rename = "max-worker-artifact-spools", alias = "max_worker_artifact_spools")] + max_worker_artifact_spools: Option, + #[serde(default, rename = "max-worker-artifact-spool-bytes", alias = "max_worker_artifact_spool_bytes")] + max_worker_artifact_spool_bytes: Option, + #[serde(default, rename = "max-worker-state-entries", alias = "max_worker_state_entries")] + max_worker_state_entries: Option, + #[serde(default, rename = "max-active-builds", alias = "max_active_builds")] + max_active_builds: Option, } #[derive(Deserialize)] @@ -612,8 +708,20 @@ fn validate_resources(raw: RawResources) -> Result { let connect_timeout = parse_duration("connect-timeout", raw.connect_timeout)?; let sync_timeout = parse_duration("sync-timeout", raw.sync_timeout)?; let build_timeout = parse_duration("build-timeout", raw.build_timeout)?; + validate_lifecycle_duration("connect-timeout", connect_timeout)?; + validate_lifecycle_duration("sync-timeout", sync_timeout)?; + validate_lifecycle_duration("build-timeout", build_timeout)?; let max_output_bytes = parse_size("max-output", raw.max_output)?; + let max_active_builds = parse_count("max-active-builds", raw.max_active_builds, 1)?; + if max_active_builds == 0 || max_active_builds > 64 { + return Err("max-active-builds must be between 1 and 64".to_string()); + } let defaults = ArtifactLimits::default(); + let idle_output_timeout = + raw.idle_output_timeout.map_or(Ok(Duration::from_secs(5 * 60)), |value| parse_duration("idle-output-timeout", value))?; + let cleanup_timeout = raw.cleanup_timeout.map_or(Ok(Duration::from_secs(5)), |value| parse_duration("cleanup-timeout", value))?; + validate_lifecycle_duration("idle-output-timeout", idle_output_timeout)?; + validate_lifecycle_duration("cleanup-timeout", cleanup_timeout)?; let artifact_timeout = raw.artifact_timeout.map_or(Ok(defaults.timeout), |value| parse_duration("artifact-timeout", value))?; let max_artifact_bytes = raw.max_artifact_bytes.map_or(Ok(defaults.max_file_bytes), |value| parse_size("max-artifact-bytes", value))?; let max_artifact_total_bytes = @@ -622,7 +730,42 @@ fn validate_resources(raw: RawResources) -> Result { .max_artifact_entries .map_or(Ok(defaults.max_entries), |value| usize::try_from(value).map_err(|_| "max-artifact-entries is too large".to_string()))?; let artifact = ArtifactLimits::new(artifact_timeout, max_artifact_entries, max_artifact_bytes, max_artifact_total_bytes)?; - Ok(ResourceLimits { connect_timeout, sync_timeout, build_timeout, max_output_bytes, artifact }) + let worker_defaults = WorkerStateLimits::default(); + let worker = WorkerStateLimits::new( + parse_count("max-worker-uploads", raw.max_worker_uploads, worker_defaults.max_uploads)?, + raw.max_worker_upload_bytes.map_or(Ok(worker_defaults.max_upload_bytes), |value| parse_size("max-worker-upload-bytes", value))?, + parse_count("max-worker-jobs", raw.max_worker_jobs, worker_defaults.max_jobs)?, + raw.max_worker_job_bytes.map_or(Ok(worker_defaults.max_job_bytes), |value| parse_size("max-worker-job-bytes", value))?, + parse_count("max-worker-artifact-spools", raw.max_worker_artifact_spools, worker_defaults.max_artifact_spools)?, + raw.max_worker_artifact_spool_bytes + .map_or(Ok(worker_defaults.max_artifact_spool_bytes), |value| parse_size("max-worker-artifact-spool-bytes", value))?, + parse_count("max-worker-state-entries", raw.max_worker_state_entries, worker_defaults.max_state_entries)?, + )?; + Ok(ResourceLimits { + connect_timeout, + sync_timeout, + build_timeout, + idle_output_timeout, + cleanup_timeout, + max_output_bytes, + max_active_builds, + artifact, + worker, + }) +} + +fn parse_count(field: &str, value: Option, default: usize) -> Result { + match value { + Some(value) => usize::try_from(value).map_err(|_| format!("{field} is too large")), + None => Ok(default), + } +} + +fn validate_lifecycle_duration(field: &str, value: Duration) -> Result<(), String> { + if value.is_zero() || value > Duration::from_secs(24 * 60 * 60) { + return Err(format!("{field} must be between 1 second and 24 hours")); + } + Ok(()) } fn parse_duration(field: &str, quantity: RawQuantity) -> Result { diff --git a/src/snapshot.rs b/src/snapshot.rs index ac4be5e..e462462 100644 --- a/src/snapshot.rs +++ b/src/snapshot.rs @@ -1,5 +1,5 @@ use crate::cfg::ProjectConfig; -use crate::remote::WorkspaceSessionId; +use crate::remote::{RemoteExecutionControl, WorkspaceSessionId}; use crate::workspace::WorkspaceHandle; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; @@ -266,7 +266,16 @@ impl SnapshotStore { } pub fn materialize(&self, handle: &SnapshotHandle, destination: &Path) -> Result { + self.materialize_with_control(handle, destination, RemoteExecutionControl::new()) + } + + pub fn materialize_with_control( + &self, handle: &SnapshotHandle, destination: &Path, control: RemoteExecutionControl, + ) -> Result { let snapshot = self.resolve(handle)?; + if control.is_cancelled() { + return Err("snapshot materialization cancelled".to_string()); + } if fs::symlink_metadata(destination).is_ok() { return Err("materialization destination already exists".to_string()); } @@ -280,6 +289,9 @@ impl SnapshotStore { let destination_root = open_directory(destination)?; for entry in snapshot.entries() { + if control.is_cancelled() { + return Err("snapshot materialization cancelled".to_string()); + } match entry.kind { SnapshotEntryKind::Directory => ensure_destination_directory(&destination_root, entry.path.as_str(), entry.mode)?, SnapshotEntryKind::RegularFile => { @@ -291,10 +303,14 @@ impl SnapshotStore { entry.size, entry.content_digest.ok_or_else(|| "regular file has no digest".to_string())?, entry.path.as_str(), + &control, )?; } } } + if control.is_cancelled() { + return Err("snapshot materialization cancelled".to_string()); + } cleanup.committed = true; Ok(MaterializedWorkspace { root: destination.to_path_buf() }) @@ -302,6 +318,12 @@ impl SnapshotStore { #[allow(dead_code)] fn read_staged_file(&self, handle: &SnapshotHandle, entry: &SnapshotEntry, max_bytes: u64) -> Result, String> { + self.read_staged_file_with_control(handle, entry, max_bytes, &RemoteExecutionControl::new()) + } + + fn read_staged_file_with_control( + &self, handle: &SnapshotHandle, entry: &SnapshotEntry, max_bytes: u64, control: &RemoteExecutionControl, + ) -> Result, String> { if entry.kind != SnapshotEntryKind::RegularFile { return Err(format!("snapshot export entry is not a regular file: {}", entry.path.as_str())); } @@ -333,6 +355,9 @@ impl SnapshotStore { let mut hasher = Sha256::new(); let mut read_bytes = 0u64; loop { + if control.is_cancelled() { + return Err("snapshot export cancelled".to_string()); + } let count = (&source).read(&mut buffer).map_err(|error| format!("read snapshot export file {}: {error}", entry.path.as_str()))?; if count == 0 { break; @@ -396,6 +421,12 @@ impl SnapshotExport { } pub(crate) fn read_file_bounded(&self, entry: &SnapshotEntry, max_bytes: u64) -> Result, String> { + self.read_file_bounded_with_control(entry, max_bytes, &RemoteExecutionControl::new()) + } + + pub(crate) fn read_file_bounded_with_control( + &self, entry: &SnapshotEntry, max_bytes: u64, control: &RemoteExecutionControl, + ) -> Result, String> { let manifest_entry = self .snapshot .entries @@ -405,7 +436,7 @@ impl SnapshotExport { if manifest_entry != entry { return Err(format!("snapshot export entry does not match the manifest: {}", entry.path.as_str())); } - self.store.read_staged_file(self.snapshot.handle(), manifest_entry, max_bytes) + self.store.read_staged_file_with_control(self.snapshot.handle(), manifest_entry, max_bytes, control) } } @@ -449,6 +480,12 @@ impl SnapshotBuilder { } pub(crate) fn build_root(&self, workspace_root: &Path, session_id: WorkspaceSessionId) -> Result { + self.build_root_with_control(workspace_root, session_id, RemoteExecutionControl::new()) + } + + pub(crate) fn build_root_with_control( + &self, workspace_root: &Path, session_id: WorkspaceSessionId, control: RemoteExecutionControl, + ) -> Result { validate_limits(&self.limits)?; if session_id.0 == [0; 16] { return Err("snapshot requires an authoritative nonzero workspace session".to_string()); @@ -476,6 +513,7 @@ impl SnapshotBuilder { root_device, stage_files: files_root, next_buffer: vec![0; SNAPSHOT_COPY_BUFFER_BYTES], + control, }; walk_directory(root.as_raw_fd(), "", 0, &self.limits, &self.exclusions, &mut state)?; state.entries.sort_by(|left, right| left.path.as_str().cmp(right.path.as_str())); @@ -485,12 +523,18 @@ impl SnapshotBuilder { if manifest.len() > self.limits.max_manifest_bytes { return Err(format!("snapshot manifest exceeds maximum size {}", self.limits.max_manifest_bytes)); } + if state.control.is_cancelled() { + return Err("snapshot creation cancelled".to_string()); + } fs::write(stage.join("manifest.json"), manifest).map_err(|error| format!("write snapshot manifest: {error}"))?; let final_session = store_root.join(hex(&session_id.0)); fs::create_dir_all(&final_session).map_err(|error| format!("create snapshot session directory: {error}"))?; set_mode(&final_session, 0o700)?; let final_path = final_session.join(hex(&id.0)); + if state.control.is_cancelled() { + return Err("snapshot creation cancelled".to_string()); + } if !final_path.exists() { fs::rename(&stage, &final_path).map_err(|error| format!("publish snapshot: {error}"))?; } else { @@ -528,16 +572,23 @@ struct WalkState { root_device: libc::dev_t, stage_files: PathBuf, next_buffer: Vec, + control: RemoteExecutionControl, } fn walk_directory( directory_fd: RawFd, parent: &str, depth: usize, limits: &SnapshotLimits, exclusions: &SnapshotExclusionPolicy, state: &mut WalkState, ) -> Result<(), String> { + if state.control.is_cancelled() { + return Err("snapshot creation cancelled".to_string()); + } check_deadline(state.started, limits)?; if depth > limits.max_depth { return Err(format!("snapshot exceeds maximum depth {}", limits.max_depth)); } for name in read_directory_names(directory_fd)? { + if state.control.is_cancelled() { + return Err("snapshot creation cancelled".to_string()); + } check_deadline(state.started, limits)?; let component = name.to_str().ok_or_else(|| "snapshot contains a non-UTF-8 path component".to_string())?; validate_component(component, limits.max_component_bytes)?; @@ -607,6 +658,9 @@ fn copy_and_hash_file( let mut hasher = Sha256::new(); let mut read_bytes = 0u64; loop { + if state.control.is_cancelled() { + return Err("snapshot creation cancelled".to_string()); + } check_deadline(state.started, limits)?; let count = (&file).read(&mut state.next_buffer).map_err(|error| format!("read snapshot file {relative}: {error}"))?; if count == 0 { @@ -833,13 +887,18 @@ fn open_relative_file(root: &File, relative: &str, flags: i32) -> Result Result<(), String> { +fn copy_materialized_file( + source: &File, destination: &File, expected_size: u64, expected_digest: [u8; 32], path: &str, control: &RemoteExecutionControl, +) -> Result<(), String> { let mut source = source.try_clone().map_err(|error| format!("clone snapshot content {path}: {error}"))?; let mut destination = destination.try_clone().map_err(|error| format!("clone materialized file {path}: {error}"))?; let mut buffer = vec![0u8; SNAPSHOT_COPY_BUFFER_BYTES]; let mut hasher = Sha256::new(); let mut copied = 0u64; loop { + if control.is_cancelled() { + return Err(format!("snapshot materialization cancelled: {path}")); + } let count = source.read(&mut buffer).map_err(|error| format!("read snapshot content {path}: {error}"))?; if count == 0 { break; diff --git a/src/ssh.rs b/src/ssh.rs index 45ca274..57dce03 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -1,8 +1,8 @@ use crate::artifact::{ArtifactLimits, ArtifactManifest, ArtifactPolicy, ArtifactPublication}; use crate::loopback::{RunRemoteSession, SnapshotExportClaim}; use crate::remote::{ - AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteFailureClass, RemoteFuture, RemoteOperation, - RemoteSnapshotId, + AuthorizedRemoteRequest, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionControl, RemoteFailureClass, RemoteFuture, + RemoteOperation, RemoteSnapshotId, }; use crate::remote_target::{ResourceLimits, SshTarget}; use crate::snapshot::SnapshotEntryKind; @@ -13,10 +13,11 @@ use crate::worker_protocol::{ }; use rand::RngCore; use std::collections::BTreeMap; +use std::future::Future; use std::io; use std::path::{Path, PathBuf}; use std::sync::{Arc, Mutex}; -use std::time::Duration; +use std::time::{Duration, Instant}; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite}; use tokio::process::Command; use tokio::task::JoinHandle; @@ -24,7 +25,6 @@ use tokio::time::timeout; const SSH_PROGRAM: &str = "/usr/bin/ssh"; const MAX_SSH_DIAGNOSTIC_BYTES: usize = 16 * 1024; -const CLEANUP_TIMEOUT: Duration = Duration::from_secs(1); const PROCESS_REAP_TIMEOUT: Duration = Duration::from_secs(2); type WorkerReader = Box; @@ -43,7 +43,21 @@ impl SshLaunchSpec { return Err("SSH target is not fully validated".to_string()); } - let remote_command = format!("exec {} --stdio --workspace-root {}", shell_quote(target.worker_path()), shell_quote(target.workspace_root())); + let worker = target.resources().worker_state_limits(); + let build_timeout_ms = target.resources().build_timeout().as_millis().max(1); + let remote_command = format!( + "exec {} --stdio --workspace-root {} --build-timeout-ms {} --max-worker-uploads {} --max-worker-upload-bytes {} --max-worker-jobs {} --max-worker-job-bytes {} --max-worker-artifact-spools {} --max-worker-artifact-spool-bytes {} --max-worker-state-entries {}", + shell_quote(target.worker_path()), + shell_quote(target.workspace_root()), + build_timeout_ms, + worker.max_uploads, + worker.max_upload_bytes, + worker.max_jobs, + worker.max_job_bytes, + worker.max_artifact_spools, + worker.max_artifact_spool_bytes, + worker.max_state_entries, + ); let connect_timeout = target.resources().connect_timeout().as_secs().max(1).to_string(); let args = vec![ "-F".to_string(), @@ -246,7 +260,7 @@ impl SshBackend { impl RemoteBackend for SshBackend { fn execute<'a>( - &'a self, request: AuthorizedRemoteRequest, events: tokio::sync::mpsc::Sender, + &'a self, request: AuthorizedRemoteRequest, control: RemoteExecutionControl, events: tokio::sync::mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { let operation = request.request().operation().clone(); let request_id = request.request_id(); @@ -257,15 +271,24 @@ impl RemoteBackend for SshBackend { let artifact_policy = self.artifact_policy.clone(); let artifact_limits = self.artifact_limits; Box::pin(async move { - let backend = SshExecution { session, target, factory, uploads, artifact_policy, artifact_limits }; + let backend = SshExecution { session, target, factory, uploads, artifact_policy, artifact_limits, control: control.clone() }; match operation { RemoteOperation::Sync(sync) => backend.execute_sync(request_id.0, sync.retain_capability(), events).await, RemoteOperation::Build(build) => backend.execute_build(request_id.0, &build, events).await, + RemoteOperation::Cancel { .. } => Err(RemoteBackendError::Failed("cancel is handled by the remote broker".to_string())), } }) } } +impl SshBackend { + pub async fn execute( + &self, request: AuthorizedRemoteRequest, events: tokio::sync::mpsc::Sender, + ) -> Result<(), RemoteBackendError> { + ::execute(self, request, RemoteExecutionControl::new(), events).await + } +} + struct SshExecution { session: Arc, target: SshTarget, @@ -273,14 +296,38 @@ struct SshExecution { uploads: Arc>>, artifact_policy: ArtifactPolicy, artifact_limits: ArtifactLimits, + control: RemoteExecutionControl, } impl SshExecution { async fn execute_sync( &self, request_id: [u8; 16], retain_capability: bool, events: tokio::sync::mpsc::Sender, ) -> Result<(), RemoteBackendError> { + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Syncing); send_event(&events, RemoteBackendEvent::SyncProgress { completed_bytes: 0, total_bytes: None }).await?; - let snapshot_id = self.session.sync_snapshot().map_err(RemoteBackendError::Failed)?; + let sync_deadline = Instant::now() + self.target.resources().sync_timeout(); + let session = self.session.clone(); + let snapshot_control = self.control.clone(); + let mut snapshot_task = tokio::task::spawn_blocking(move || session.sync_snapshot_for_request_with_control(true, snapshot_control)); + let snapshot_id = tokio::select! { + result = tokio::time::timeout(remaining(sync_deadline), &mut snapshot_task) => { + match result { + Ok(result) => result + .map_err(|error| RemoteBackendError::Failed(format!("snapshot worker failed: {error}")))? + .map_err(RemoteBackendError::Failed)? + .ok_or_else(|| RemoteBackendError::Failed("remote snapshot capability was not retained".to_string()))?, + Err(_) => { + self.control.cancel(); + let _ = tokio::time::timeout(self.target.resources().cleanup_timeout(), &mut snapshot_task).await; + return Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Sync }); + } + } + } + _ = self.control.cancelled() => { + let _ = tokio::time::timeout(self.target.resources().cleanup_timeout(), &mut snapshot_task).await; + return Err(RemoteBackendError::Cancelled); + } + }; let export = match self.session.claim_snapshot_for_export(snapshot_id) { Ok(export) => export, Err(error) => { @@ -302,7 +349,19 @@ impl SshExecution { let _ = self.session.abort_snapshot_capability(snapshot_id); return Err(error); } - let mut connection = match WorkerConnection::spawn(&self.factory, &self.target) { + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Connecting); + let mut connection = match phase( + async { + let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; + connection.handshake(WorkerRequestId(request_id), session_id_for(&self.session), WORKER_PROTOCOL_VERSION).await?; + Ok(connection) + }, + remaining(sync_deadline).min(self.target.resources().connect_timeout()), + &self.control, + crate::remote::RemoteTimeoutCause::Connection, + ) + .await + { Ok(connection) => connection, Err(error) => { drop(export); @@ -310,6 +369,7 @@ impl SshExecution { return Err(error); } }; + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Transferring); let session_id = session_id_for(&self.session); let operation = upload_and_finish( &mut connection, @@ -321,23 +381,27 @@ impl SshExecution { export: &export, events: &events, cleanup: !retain_capability, + control: &self.control, + handshake: false, }, ); - let result = match timeout(self.target.resources().sync_timeout(), operation).await { - Ok(result) => result, - Err(_) => { - connection.kill_and_reap().await; - Err(RemoteBackendError::Timeout) - } - }; + let result = phase(operation, remaining(sync_deadline), &self.control, crate::remote::RemoteTimeoutCause::Sync).await; drop(export); if let Err(error) = result { - connection.kill_and_reap().await; + connection.kill_and_reap(self.target.resources().cleanup_timeout()).await; + self.cleanup_upload(request_id, upload_id).await; let _ = self.session.abort_snapshot_capability(snapshot_id); return Err(error); } + if self.control.is_cancelled() { + connection.kill_and_reap(self.target.resources().cleanup_timeout()).await; + self.cleanup_upload(request_id, upload_id).await; + let _ = self.session.abort_snapshot_capability(snapshot_id); + return Err(RemoteBackendError::Cancelled); + } + if retain_capability { self.uploads .lock() @@ -345,6 +409,7 @@ impl SshExecution { .insert(snapshot_id, upload_id); if let Err(error) = send_event(&events, RemoteBackendEvent::SyncCompleted { snapshot_id }).await { self.remove_upload(snapshot_id); + self.cleanup_upload(request_id, upload_id).await; let _ = self.session.abort_snapshot_capability(snapshot_id); return Err(error); } @@ -358,6 +423,14 @@ impl SshExecution { async fn execute_build( &self, request_id: [u8; 16], build: &crate::remote::RemoteBuild, events: tokio::sync::mpsc::Sender, ) -> Result<(), RemoteBackendError> { + if self.control.is_cancelled() { + let upload_id = self.uploads.lock().ok().and_then(|mut uploads| uploads.remove(&build.snapshot_id())); + if let Some(upload_id) = upload_id { + self.cleanup_upload(request_id, upload_id).await; + } + let _ = self.session.abort_snapshot_capability(build.snapshot_id()); + return Err(RemoteBackendError::Cancelled); + } let snapshot_id = build.snapshot_id(); let upload_id = self .uploads @@ -368,11 +441,18 @@ impl SshExecution { class: RemoteFailureClass::SnapshotTransfer, message: "remote snapshot upload is unavailable".to_string(), })?; - let claim = self.session.claim_snapshot(snapshot_id).map_err(RemoteBackendError::Failed)?; + let claim = match self.session.claim_snapshot(snapshot_id) { + Ok(claim) => claim, + Err(error) => { + self.cleanup_upload(request_id, upload_id).await; + return Err(RemoteBackendError::Failed(error)); + } + }; let executable = match self.target.tools().get(build.tool().as_str()) { Some(executable) => executable.clone(), None => { drop(claim); + self.cleanup_upload(request_id, upload_id).await; return Err(RemoteBackendError::Transport { class: RemoteFailureClass::WorkerUnavailable, message: format!("remote tool is not configured: {}", build.tool().as_str()), @@ -381,25 +461,69 @@ impl SshExecution { }; let guest_env = build.env().to_vec(); let target_env = self.target.environment().iter().map(|(key, value)| (key.clone(), value.clone())).collect::>(); - let worker_build = - WorkerBuild::new(build.tool().as_str(), executable, build.argv().to_vec(), build.cwd().as_str(), guest_env, target_env, upload_id) - .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, message: error.to_string() })?; + let worker_build = match WorkerBuild::new( + build.tool().as_str(), + executable, + build.argv().to_vec(), + build.cwd().as_str(), + guest_env, + target_env, + upload_id, + ) { + Ok(worker_build) => worker_build, + Err(error) => { + drop(claim); + self.cleanup_upload(request_id, upload_id).await; + return Err(RemoteBackendError::Transport { class: RemoteFailureClass::WorkerProtocol, message: error.to_string() }); + } + }; let worker_build = if self.artifact_policy.is_enabled() { - let paths = self - .artifact_policy - .paths() - .iter() - .map(|path| WorkerArtifactPath::new(path.clone())) - .collect::, _>>() - .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error.to_string() })?; - worker_build - .with_artifacts(paths, self.artifact_limits.max_file_bytes, self.artifact_limits.max_total_bytes) - .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error.to_string() })? + let paths = match self.artifact_policy.paths().iter().map(|path| WorkerArtifactPath::new(path.clone())).collect::, _>>() { + Ok(paths) => paths, + Err(error) => { + drop(claim); + self.cleanup_upload(request_id, upload_id).await; + return Err(RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error.to_string() }); + } + }; + match worker_build.with_artifacts(paths, self.artifact_limits.max_file_bytes, self.artifact_limits.max_total_bytes) { + Ok(worker_build) => worker_build, + Err(error) => { + drop(claim); + self.cleanup_upload(request_id, upload_id).await; + return Err(RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, message: error.to_string() }); + } + } } else { worker_build }; - let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; + self.control.set_phase(crate::remote::RemoteLifecyclePhase::Connecting); + let mut connection = match phase( + async { + let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; + connection + .handshake( + WorkerRequestId(request_id), + WorkerSessionId(self.session.session_id().0), + if self.artifact_policy.is_enabled() { WORKER_ARTIFACT_PROTOCOL_VERSION } else { WORKER_PROTOCOL_VERSION }, + ) + .await?; + Ok(connection) + }, + self.target.resources().connect_timeout(), + &self.control, + crate::remote::RemoteTimeoutCause::Connection, + ) + .await + { + Ok(connection) => connection, + Err(error) => { + self.cleanup_upload(request_id, upload_id).await; + drop(claim); + return Err(error); + } + }; let session_id = WorkerSessionId(self.session.session_id().0); let operation = build_and_finish( &mut connection, @@ -414,18 +538,28 @@ impl SshExecution { artifact_limits: self.artifact_limits, workspace_root: self.session.workspace_root().to_path_buf(), protocol_version: if self.artifact_policy.is_enabled() { WORKER_ARTIFACT_PROTOCOL_VERSION } else { WORKER_PROTOCOL_VERSION }, + control: &self.control, + handshake: false, }, ); - let result = match timeout(self.target.resources().build_timeout(), operation).await { - Ok(result) => result, - Err(_) => { - connection.kill_and_reap().await; - Err(RemoteBackendError::Timeout) - } - }; + let result = operation.await; if let Err(error) = result { - let _ = timeout(CLEANUP_TIMEOUT, connection.cleanup(WorkerRequestId(request_id), session_id, upload_id)).await; - connection.kill_and_reap().await; + let cleanup_succeeded = if !matches!(error, RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Cleanup }) { + phase( + connection.cleanup(WorkerRequestId(request_id), session_id, upload_id), + self.target.resources().cleanup_timeout(), + &self.control, + crate::remote::RemoteTimeoutCause::Cleanup, + ) + .await + .is_ok() + } else { + false + }; + connection.kill_and_reap(self.target.resources().cleanup_timeout()).await; + if !cleanup_succeeded { + self.cleanup_upload(request_id, upload_id).await; + } return Err(error); } drop(claim); @@ -437,6 +571,41 @@ impl SshExecution { uploads.remove(&snapshot_id); } } + + async fn cleanup_upload(&self, request_id: [u8; 16], upload_id: WorkerUploadId) { + let cleanup_control = RemoteExecutionControl::new(); + let cleanup_timeout = self.target.resources().cleanup_timeout(); + let session_id = session_id_for(&self.session); + let mut connection = match phase( + async { + let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; + connection.handshake(WorkerRequestId(request_id), session_id, WORKER_PROTOCOL_VERSION).await?; + Ok(connection) + }, + cleanup_timeout, + &cleanup_control, + crate::remote::RemoteTimeoutCause::Cleanup, + ) + .await + { + Ok(connection) => connection, + Err(_) => return, + }; + + let cleaned = phase( + connection.cleanup(WorkerRequestId(request_id), session_id, upload_id), + cleanup_timeout, + &cleanup_control, + crate::remote::RemoteTimeoutCause::Cleanup, + ) + .await + .is_ok(); + if cleaned { + let _ = phase(connection.finish(), cleanup_timeout, &cleanup_control, crate::remote::RemoteTimeoutCause::Cleanup).await; + } else { + connection.kill_and_reap(cleanup_timeout).await; + } + } } struct WorkerConnection { @@ -551,9 +720,9 @@ impl WorkerConnection { Ok(()) } - async fn kill_and_reap(&mut self) { + async fn kill_and_reap(&mut self, cleanup_timeout: Duration) { self.process.kill_group(); - let _ = timeout(PROCESS_REAP_TIMEOUT, self.process.wait()).await; + let _ = timeout(cleanup_timeout.min(PROCESS_REAP_TIMEOUT), self.process.wait()).await; if let Some(task) = self.stderr_task.take() { task.abort(); } @@ -577,6 +746,8 @@ struct UploadPlan<'a> { export: &'a SnapshotExportClaim, events: &'a tokio::sync::mpsc::Sender, cleanup: bool, + control: &'a RemoteExecutionControl, + handshake: bool, } struct BuildPlan<'a> { @@ -590,22 +761,33 @@ struct BuildPlan<'a> { artifact_limits: ArtifactLimits, workspace_root: PathBuf, protocol_version: u16, + control: &'a RemoteExecutionControl, + handshake: bool, } async fn upload_and_finish(connection: &mut WorkerConnection, plan: UploadPlan<'_>) -> Result<(), RemoteBackendError> { - let UploadPlan { request_id, session_id, upload_id, entries, export, events, cleanup } = plan; - connection.handshake(request_id, session_id, WORKER_PROTOCOL_VERSION).await?; + let UploadPlan { request_id, session_id, upload_id, entries, export, events, cleanup, control, handshake } = plan; + if handshake { + connection.handshake(request_id, session_id, WORKER_PROTOCOL_VERSION).await?; + } + control.set_phase(crate::remote::RemoteLifecyclePhase::Transferring); let total_bytes = export.total_file_bytes(); connection.write(&WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries: entries.to_vec() }).await?; let mut completed_bytes = 0u64; for entry in export.entries() { + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } if entry.kind() != SnapshotEntryKind::RegularFile { continue; } let contents = export - .read_file_bounded(entry, MAX_WORKER_FILE_BYTES) + .read_file_bounded_with_control(entry, MAX_WORKER_FILE_BYTES, control) .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::SnapshotTransfer, message: error })?; for (index, chunk) in contents.chunks(MAX_WORKER_CHUNK_BYTES).enumerate() { + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } let offset = u64::try_from(index).unwrap_or(u64::MAX).saturating_mul(MAX_WORKER_CHUNK_BYTES as u64); connection .write(&WorkerMessage::UploadFileChunk { @@ -653,73 +835,130 @@ async fn upload_and_finish(connection: &mut WorkerConnection, plan: UploadPlan<' } async fn build_and_finish(connection: &mut WorkerConnection, plan: BuildPlan<'_>) -> Result<(), RemoteBackendError> { - let BuildPlan { request_id, session_id, build, events, upload_id, resources, artifact_policy, artifact_limits, workspace_root, protocol_version } = - plan; - connection.handshake(request_id, session_id, protocol_version).await?; + let BuildPlan { + request_id, + session_id, + build, + events, + upload_id, + resources, + artifact_policy, + artifact_limits, + workspace_root, + protocol_version, + control, + handshake, + } = plan; + if handshake { + connection.handshake(request_id, session_id, protocol_version).await?; + } + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } + control.set_phase(crate::remote::RemoteLifecyclePhase::Building); + let exit_code = phase( + execute_native_build(connection, request_id, session_id, build, events, resources, control), + resources.build_timeout(), + control, + crate::remote::RemoteTimeoutCause::Build, + ) + .await?; + if exit_code == 0 && artifact_policy.is_enabled() { + control.set_phase(crate::remote::RemoteLifecyclePhase::ArtifactHandling); + let artifact_deadline = Instant::now() + artifact_limits.timeout; + let (artifact_set_id, manifest) = phase( + async { + match connection.read().await? { + WorkerMessage::ArtifactManifest { + request_id: received_request, + session_id: received_session, + artifact_set_id, + entries, + total_bytes, + } => { + check_correlation(received_request, received_session, request_id, session_id, "worker artifact manifest")?; + Ok((artifact_set_id, artifact_manifest_from_worker(entries, total_bytes, &artifact_policy, artifact_limits)?)) + } + WorkerMessage::Error { kind, message, .. } => Err(worker_artifact_manifest_error(kind, message)), + _ => Err(worker_protocol("unexpected worker artifact manifest response")), + } + }, + remaining(artifact_deadline), + control, + crate::remote::RemoteTimeoutCause::Artifact, + ) + .await?; + phase( + fetch_and_publish_artifacts(connection, request_id, session_id, artifact_set_id, &manifest, workspace_root, control), + remaining(artifact_deadline), + control, + crate::remote::RemoteTimeoutCause::Artifact, + ) + .await?; + } + control.set_phase(crate::remote::RemoteLifecyclePhase::Cleanup); + phase(connection.cleanup(request_id, session_id, upload_id), resources.cleanup_timeout(), control, crate::remote::RemoteTimeoutCause::Cleanup) + .await?; + phase(connection.finish(), resources.cleanup_timeout(), control, crate::remote::RemoteTimeoutCause::Cleanup).await?; + control.set_phase(crate::remote::RemoteLifecyclePhase::Finalizing); + send_event(events, RemoteBackendEvent::Completed { exit_code }).await +} + +async fn execute_native_build( + connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, build: WorkerBuild, + events: &tokio::sync::mpsc::Sender, resources: ResourceLimits, control: &RemoteExecutionControl, +) -> Result { connection.write(&WorkerMessage::build(request_id, session_id, build)).await?; let mut output_bytes = 0u64; - let exit_code = loop { - match connection.read().await? { - WorkerMessage::Stdout { request_id: received_request, session_id: received_session, data } => { - check_correlation(received_request, received_session, request_id, session_id, "worker stdout")?; - output_bytes = output_bytes.saturating_add(data.len() as u64); - if output_bytes > resources.max_output_bytes() { - return Err(RemoteBackendError::OutputLimit { limit: resources.max_output_bytes() }); + let mut idle = Box::pin(tokio::time::sleep(resources.idle_output_timeout())); + loop { + tokio::select! { + result = connection.read() => match result? { + WorkerMessage::Stdout { request_id: received_request, session_id: received_session, data } => { + check_correlation(received_request, received_session, request_id, session_id, "worker stdout")?; + if !data.is_empty() { + idle.as_mut().reset(tokio::time::Instant::now() + resources.idle_output_timeout()); + } + output_bytes = output_bytes.saturating_add(data.len() as u64); + if output_bytes > resources.max_output_bytes() { + return Err(RemoteBackendError::OutputLimit { limit: resources.max_output_bytes() }); + } + send_event_control(events, RemoteBackendEvent::Stdout(data), control).await?; } - send_event(events, RemoteBackendEvent::Stdout(data)).await?; - } - WorkerMessage::Stderr { request_id: received_request, session_id: received_session, data } => { - check_correlation(received_request, received_session, request_id, session_id, "worker stderr")?; - output_bytes = output_bytes.saturating_add(data.len() as u64); - if output_bytes > resources.max_output_bytes() { - return Err(RemoteBackendError::OutputLimit { limit: resources.max_output_bytes() }); + WorkerMessage::Stderr { request_id: received_request, session_id: received_session, data } => { + check_correlation(received_request, received_session, request_id, session_id, "worker stderr")?; + if !data.is_empty() { + idle.as_mut().reset(tokio::time::Instant::now() + resources.idle_output_timeout()); + } + output_bytes = output_bytes.saturating_add(data.len() as u64); + if output_bytes > resources.max_output_bytes() { + return Err(RemoteBackendError::OutputLimit { limit: resources.max_output_bytes() }); + } + send_event_control(events, RemoteBackendEvent::Stderr(data), control).await?; } - send_event(events, RemoteBackendEvent::Stderr(data)).await?; - } - WorkerMessage::Completed { request_id: received_request, session_id: received_session, operation: WorkerOperation::Build, exit_code } => { - check_correlation(received_request, received_session, request_id, session_id, "worker completion")?; - break exit_code; - } - WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Build, kind, message)), - _ => return Err(worker_protocol("unexpected worker build response")), - } - }; - if exit_code == 0 && artifact_policy.is_enabled() { - let (artifact_set_id, manifest) = match connection.read().await? { - WorkerMessage::ArtifactManifest { request_id: received_request, session_id: received_session, artifact_set_id, entries, total_bytes } => { - check_correlation(received_request, received_session, request_id, session_id, "worker artifact manifest")?; - (artifact_set_id, artifact_manifest_from_worker(entries, total_bytes, &artifact_policy, artifact_limits)?) - } - WorkerMessage::Error { kind, message, .. } => return Err(worker_artifact_manifest_error(kind, message)), - _ => return Err(worker_protocol("unexpected worker artifact manifest response")), - }; - let retrieval = timeout( - artifact_limits.timeout, - fetch_and_publish_artifacts(connection, request_id, session_id, artifact_set_id, &manifest, workspace_root), - ) - .await; - match retrieval { - Ok(result) => result?, - Err(_) => { - return Err(RemoteBackendError::Transport { - class: RemoteFailureClass::ArtifactTransfer, - message: "artifact retrieval timed out".to_string(), - }) - } + WorkerMessage::Completed { request_id: received_request, session_id: received_session, operation: WorkerOperation::Build, exit_code } => { + check_correlation(received_request, received_session, request_id, session_id, "worker completion")?; + return Ok(exit_code); + } + WorkerMessage::Error { kind, message, .. } => return Err(worker_error(WorkerOperation::Build, kind, message)), + _ => return Err(worker_protocol("unexpected worker build response")), + }, + _ = &mut idle => return Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::IdleOutput }), + _ = control.cancelled() => return Err(RemoteBackendError::Cancelled), } } - connection.cleanup(request_id, session_id, upload_id).await?; - connection.finish().await?; - send_event(events, RemoteBackendEvent::Completed { exit_code }).await } async fn fetch_and_publish_artifacts( connection: &mut WorkerConnection, request_id: WorkerRequestId, session_id: WorkerSessionId, artifact_set_id: WorkerArtifactSetId, - manifest: &ArtifactManifest, workspace_root: PathBuf, + manifest: &ArtifactManifest, workspace_root: PathBuf, control: &RemoteExecutionControl, ) -> Result<(), RemoteBackendError> { let mut publication = ArtifactPublication::new(&workspace_root, request_id.0, manifest.clone()) .map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error })?; for index in 0..manifest.entries().len() { + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } connection .write(&WorkerMessage::FetchArtifact { request_id, @@ -745,6 +984,9 @@ async fn fetch_and_publish_artifacts( offset, data, } => { + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } check_correlation(received_request, received_session, request_id, session_id, "worker artifact chunk") .map_err(artifact_transfer_error)?; if received_set != artifact_set_id || entry_index != index as u32 { @@ -792,6 +1034,10 @@ async fn fetch_and_publish_artifacts( } } } + if control.is_cancelled() { + return Err(RemoteBackendError::Cancelled); + } + control.set_phase(crate::remote::RemoteLifecyclePhase::Finalizing); publication.publish().map_err(|error| RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactTransfer, message: error }) } @@ -856,6 +1102,15 @@ fn send_event<'a>( Box::pin(async move { events.send(event).await.map_err(|_| RemoteBackendError::Cancelled) }) } +async fn send_event_control( + events: &tokio::sync::mpsc::Sender, event: RemoteBackendEvent, control: &RemoteExecutionControl, +) -> Result<(), RemoteBackendError> { + tokio::select! { + result = events.send(event) => result.map_err(|_| RemoteBackendError::Cancelled), + _ = control.cancelled() => Err(RemoteBackendError::Cancelled), + } +} + fn classify_spawn_error(message: String) -> RemoteBackendError { RemoteBackendError::Transport { class: RemoteFailureClass::Connect, message } } @@ -887,6 +1142,25 @@ fn worker_io_error(error: worker_protocol::WorkerProtocolError) -> RemoteBackend } } +async fn phase( + future: F, duration: Duration, control: &RemoteExecutionControl, cause: crate::remote::RemoteTimeoutCause, +) -> Result +where + F: Future>, +{ + if duration.is_zero() { + return Err(RemoteBackendError::Deadline { cause }); + } + tokio::select! { + result = timeout(duration, future) => result.unwrap_or_else(|_| Err(RemoteBackendError::Deadline { cause })), + _ = control.cancelled() => Err(RemoteBackendError::Cancelled), + } +} + +fn remaining(deadline: Instant) -> Duration { + deadline.saturating_duration_since(Instant::now()) +} + fn worker_error(operation: WorkerOperation, kind: WorkerErrorKind, message: String) -> RemoteBackendError { let class = match kind { WorkerErrorKind::WorkerProtocol => RemoteFailureClass::WorkerProtocol, diff --git a/src/vscomm/mod.rs b/src/vscomm/mod.rs index d444dfd..e405fb5 100644 --- a/src/vscomm/mod.rs +++ b/src/vscomm/mod.rs @@ -185,6 +185,7 @@ impl RemoteBuild { pub enum RemoteOperation { Sync(RemoteSync), Build(RemoteBuild), + Cancel { target_request_id: RequestId }, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -211,12 +212,23 @@ impl RemoteRequest { Self { request_id, workspace_session_id, operation: RemoteOperation::Build(build) } } + pub fn cancel(request_id: RequestId, workspace_session_id: WorkspaceSessionId, target_request_id: RequestId) -> Self { + Self { request_id, workspace_session_id, operation: RemoteOperation::Cancel { target_request_id } } + } + pub fn to_frame(&self) -> Result { + if self.request_id.0 == [0; 16] { + return Err("remote request ID must be nonzero".to_string()); + } + if self.workspace_session_id.0 == [0; 16] { + return Err("remote workspace session ID must be nonzero".to_string()); + } let mut writer = WireWriter::new(*b"BBR1"); writer.u16(REMOTE_PROTOCOL_VERSION); writer.u8(match &self.operation { RemoteOperation::Sync(_) => 1, RemoteOperation::Build(_) => 2, + RemoteOperation::Cancel { .. } => 3, }); writer.u8(0); writer.bytes(&self.request_id.0); @@ -226,6 +238,11 @@ impl RemoteRequest { encode_remote_build(&mut writer, build)?; } else if let RemoteOperation::Sync(sync) = &self.operation { writer.u8(u8::from(sync.retain_capability)); + } else if let RemoteOperation::Cancel { target_request_id } = &self.operation { + if target_request_id.0 == [0; 16] { + return Err("remote cancel target request ID must be nonzero".to_string()); + } + writer.bytes(&target_request_id.0); } writer.into_frame(FrameType::RemoteRequest) @@ -242,7 +259,13 @@ impl RemoteRequest { let operation_kind = reader.u8()?; reader.zero_reserved()?; let request_id = RequestId(reader.array16()?); + if request_id.0 == [0; 16] { + return Err("remote request ID must be nonzero".to_string()); + } let workspace_session_id = WorkspaceSessionId(reader.array16()?); + if workspace_session_id.0 == [0; 16] { + return Err("remote workspace session ID must be nonzero".to_string()); + } let operation = match operation_kind { 1 => { let retain_capability = match reader.u8()? { @@ -253,6 +276,13 @@ impl RemoteRequest { RemoteOperation::Sync(RemoteSync { retain_capability }) } 2 => RemoteOperation::Build(decode_remote_build(&mut reader)?), + 3 => { + let target_request_id = RequestId(reader.array16()?); + if target_request_id.0 == [0; 16] { + return Err("remote cancel target request ID must be nonzero".to_string()); + } + RemoteOperation::Cancel { target_request_id } + } value => return Err(format!("unknown remote operation: {value}")), }; let request = Self { request_id, workspace_session_id, operation }; @@ -278,6 +308,9 @@ impl RemoteRequest { let build = remote_domain::RemoteBuild::new(cwd, tool, build.argv, build.env, snapshot_id)?; Ok(remote_domain::RemoteRequest::build(request_id, session_id, build)) } + RemoteOperation::Cancel { target_request_id } => { + Ok(remote_domain::RemoteRequest::cancel(request_id, session_id, remote_domain::RequestId(target_request_id.0))) + } } } } @@ -564,6 +597,9 @@ impl RemoteEvent { } pub fn to_frame(&self) -> Result { + if self.request_id.0 == [0; 16] { + return Err("remote event request ID must be nonzero".to_string()); + } let mut writer = WireWriter::new(*b"BBE1"); writer.u16(REMOTE_PROTOCOL_VERSION); writer.u8(match &self.kind { @@ -615,6 +651,9 @@ impl RemoteEvent { let event_kind = reader.u8()?; reader.zero_reserved()?; let request_id = RequestId(reader.array16()?); + if request_id.0 == [0; 16] { + return Err("remote event request ID must be nonzero".to_string()); + } let kind = match event_kind { 1 => { let completed_bytes = reader.u64()?; From 8932240c56140a84f6f479cb18bd511cac9a030f Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 01:57:23 +0200 Subject: [PATCH 38/52] Add unit tests --- src/artifact_ut.rs | 25 ++++++++++ src/bunkerbox-remote_ut.rs | 37 ++++++++++++++- src/cfg_ut.rs | 44 +++++++++++++++++ src/daemon_ut.rs | 90 ++++++++++++++++++++++++++++++++--- src/guest_install_ut.rs | 67 ++++++++++++++++++++++++++ src/loopback_ut.rs | 97 +++++++++++++++++++++++++++++++++++++- src/main_ut.rs | 10 ++-- src/remote_client_ut.rs | 14 ++++++ src/remote_target_ut.rs | 31 ++++++++++++ src/remote_ut.rs | 60 +++++++++++++++++++++++ src/snapshot_ut.rs | 18 ++++++- src/ssh_ut.rs | 6 ++- src/vscomm/mod_ut.rs | 22 +++++++++ 13 files changed, 508 insertions(+), 13 deletions(-) diff --git a/src/artifact_ut.rs b/src/artifact_ut.rs index bf7d793..ca32287 100644 --- a/src/artifact_ut.rs +++ b/src/artifact_ut.rs @@ -1,4 +1,5 @@ use super::*; +use crate::remote::RemoteExecutionControl; use std::fs; use std::io::Read; use tempfile::tempdir; @@ -113,6 +114,30 @@ fn local_capture_rejects_symlink_outputs() { assert!(LocalArtifactSpool::capture(&job, &jobs, &policy, limits()).is_err()); } +#[test] +fn cancelled_artifact_capture_and_copy_leave_no_publication() { + let temp = tempdir().unwrap(); + let job = temp.path().join("job"); + let workspace = temp.path().join("workspace"); + let jobs = temp.path().join("jobs"); + fs::create_dir(&job).unwrap(); + fs::create_dir(&workspace).unwrap(); + fs::create_dir(&jobs).unwrap(); + fs::write(job.join("result"), b"artifact").unwrap(); + let policy = ArtifactPolicy::new(vec!["result".to_string()]).unwrap(); + let control = RemoteExecutionControl::new(); + control.cancel(); + + assert!(LocalArtifactSpool::capture_with_control(&job, &jobs, &policy, limits(), &control).is_err()); + + let entry = ArtifactEntry::new("result", 0o644, 8, sha256(b"artifact")).unwrap(); + let manifest = ArtifactManifest::new(vec![entry], 8, &policy, limits()).unwrap(); + let mut publication = ArtifactPublication::new(&workspace, [4; 16], manifest).unwrap(); + assert!(publication.copy_from_reader_with_control(0, &mut &b"artifact"[..], &control).is_err()); + drop(publication); + assert!(!workspace.join(".bunkerbox/artifacts/.staging/04040404040404040404040404040404").exists()); +} + fn sha256(value: &[u8]) -> [u8; 32] { use sha2::{Digest, Sha256}; Sha256::digest(value).into() diff --git a/src/bunkerbox-remote_ut.rs b/src/bunkerbox-remote_ut.rs index b5fbacf..292ceb5 100644 --- a/src/bunkerbox-remote_ut.rs +++ b/src/bunkerbox-remote_ut.rs @@ -1,5 +1,5 @@ use super::*; -use bunkerbox::vscomm::{Frame, RemoteEvent, RemoteEventKind}; +use bunkerbox::vscomm::{Frame, RemoteEvent, RemoteEventKind, RemoteOperation}; fn snapshot_id() -> bunkerbox::remote::RemoteSnapshotId { bunkerbox::remote::RemoteSnapshotId::from_bytes([9; 16]) @@ -50,6 +50,14 @@ fn parses_build_tool_and_args_without_joining() { ); } +#[test] +fn parses_cargo_toolchain_and_arguments_without_joining() { + assert_eq!( + parse_command(&["build".into(), "cargo".into(), "+nightly".into(), "build".into(), "literal $(value)".into()]), + Ok(RemoteCommand::Build { tool: "cargo".into(), args: vec!["+nightly".into(), "build".into(), "literal $(value)".into()] }) + ); +} + #[test] fn build_request_preserves_logical_cwd_and_arguments() { let request = build_request( @@ -68,6 +76,23 @@ fn build_request_preserves_logical_cwd_and_arguments() { assert_eq!(build.argv, ["release mode", "$(literal)"]); } +#[test] +fn cargo_build_request_preserves_structured_toolchain_arguments() { + let request = build_request( + RemoteCommand::Build { tool: "cargo".into(), args: vec!["+nightly".into(), "build".into(), "$(literal)".into()] }, + "crates/app".into(), + RequestId([1; 16]), + WorkspaceSessionId([2; 16]), + snapshot_id(), + ) + .unwrap(); + let RemoteOperation::Build(build) = request.operation else { panic!("expected build") }; + assert_eq!(build.tool.as_str(), "cargo"); + assert_eq!(build.cwd.as_str(), "crates/app"); + assert_eq!(build.argv, ["+nightly", "build", "$(literal)"]); + assert!(build.env.is_empty()); +} + #[test] fn sync_success_uses_existing_remote_helper_and_returns_status() { let request_id = RequestId([3; 16]); @@ -207,6 +232,16 @@ fn configured_remote_make_does_not_overwrite_unmanaged_entry() { assert_eq!(std::fs::read(root.path().join("make")).unwrap(), b"native"); } +#[test] +fn configured_remote_cargo_does_not_overwrite_unmanaged_entry() { + let root = tempfile::tempdir().unwrap(); + let executable = root.path().join("bunkerbox-remote"); + std::fs::write(&executable, b"remote").unwrap(); + std::fs::write(root.path().join("cargo"), b"native").unwrap(); + assert!(install_remote_cargo_link(root.path(), &executable, true).is_err()); + assert_eq!(std::fs::read(root.path().join("cargo")).unwrap(), b"native"); +} + #[test] fn transparent_build_syncs_first_and_reuses_that_capability() { let session = WorkspaceSessionId([2; 16]); diff --git a/src/cfg_ut.rs b/src/cfg_ut.rs index 6d4ddde..4181843 100644 --- a/src/cfg_ut.rs +++ b/src/cfg_ut.rs @@ -540,6 +540,10 @@ fn load_or_create_validates_remote_policy_configuration() { let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); assert_eq!(cfg.project.remote.environment, vec!["PROJECT_MODE"]); assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "make".into(), allow_args: false }]); + + write_project_conf(root.path(), "project:\n remote:\n tools:\n - name: cargo\n allow-args: true\n"); + let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); + assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "cargo".into(), allow_args: true }]); } #[test] @@ -548,3 +552,43 @@ fn load_or_create_rejects_forbidden_remote_environment() { write_project_conf(root.path(), "project:\n remote:\n environment:\n - SSH_AUTH_SOCK\n"); assert!(ProjectConfig::load_or_create(root.path()).is_err()); } + +#[test] +fn remote_global_active_build_limit_defaults_and_rejects_invalid_values() { + let config = RuntimeConfig { + oci: PathBuf::from("/usr/bin/oci"), + image: "image".into(), + network: None, + allow: None, + workspace: None, + workspace_quota: None, + workspace_exclude: None, + home: None, + home_path: None, + encrypt: None, + session_mb: None, + session_cleanup: None, + command: None, + remote_max_active_builds: None, + }; + assert_eq!(config.remote_max_active_builds().unwrap(), 2); + for value in [Some(0), Some((crate::remote::RemoteAdmissionLimits::MAX + 1) as u64)] { + let config = RuntimeConfig { + oci: PathBuf::from("/usr/bin/oci"), + image: "image".into(), + network: None, + allow: None, + workspace: None, + workspace_quota: None, + workspace_exclude: None, + home: None, + home_path: None, + encrypt: None, + session_mb: None, + session_cleanup: None, + command: None, + remote_max_active_builds: value, + }; + assert!(config.remote_max_active_builds().is_err()); + } +} diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index 07b0e0f..df7f6cc 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -2,8 +2,9 @@ use super::{dispatch_remote_frame, is_allowed, RemoteBroker, RemoteDispatchError use super::{monitor_bwrap_status, ChildEvent}; use crate::cfg::EnvMode; use crate::remote::{ - AuthorizedRemoteRequest, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteExecutionContext, RemoteFuture, - RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, RequestId, WorkspaceRelativePath, WorkspaceSessionId, + AuthorizedRemoteRequest, RemoteAdmissionLimits, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, + RemoteExecutionContext, RemoteExecutionControl, RemoteFuture, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, + RequestId, WorkspaceRelativePath, WorkspaceSessionId, }; use crate::vscomm::{ Frame, RemoteBuild as WireRemoteBuild, RemoteRequest as WireRemoteRequest, RemoteTool as WireRemoteTool, RequestId as WireRequestId, @@ -15,6 +16,7 @@ use std::sync::{Arc, Mutex}; use std::task::{Context, Poll}; use tokio::io::AsyncWrite; use tokio::sync::mpsc; +use tokio::sync::Notify; #[test] fn bwrap_status_reports_command_start() { @@ -64,7 +66,7 @@ impl RemoteSnapshotAuthority for RejectSnapshotAuthority { impl RemoteBackend for RecordingBackend { fn execute<'a>( - &'a self, request: AuthorizedRemoteRequest, events: mpsc::Sender, + &'a self, request: AuthorizedRemoteRequest, _control: RemoteExecutionControl, events: mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { self.calls.lock().unwrap().push(request); let emit = self.emit.clone(); @@ -85,7 +87,7 @@ struct StreamingBackend { impl RemoteBackend for StreamingBackend { fn execute<'a>( - &'a self, _request: AuthorizedRemoteRequest, events: mpsc::Sender, + &'a self, _request: AuthorizedRemoteRequest, _control: RemoteExecutionControl, events: mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { let event_count = self.event_count; let release = self.release.clone(); @@ -103,7 +105,7 @@ struct HangingBackend; impl RemoteBackend for HangingBackend { fn execute<'a>( - &'a self, _request: AuthorizedRemoteRequest, events: mpsc::Sender, + &'a self, _request: AuthorizedRemoteRequest, _control: RemoteExecutionControl, events: mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { Box::pin(async move { events.send(RemoteBackendEvent::Stdout(b"first".to_vec())).await.map_err(|_| RemoteBackendError::Cancelled)?; @@ -112,6 +114,28 @@ impl RemoteBackend for HangingBackend { } } +struct HoldingBackend { + started: Arc, + release: Arc, +} + +impl RemoteBackend for HoldingBackend { + fn execute<'a>( + &'a self, _request: AuthorizedRemoteRequest, control: RemoteExecutionControl, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + let started = self.started.clone(); + let release = self.release.clone(); + Box::pin(async move { + events.send(RemoteBackendEvent::Stdout(b"active".to_vec())).await.map_err(|_| RemoteBackendError::Cancelled)?; + started.notify_one(); + tokio::select! { + _ = release.notified() => events.send(RemoteBackendEvent::Completed { exit_code: 0 }).await.map_err(|_| RemoteBackendError::Cancelled), + _ = control.cancelled() => Err(RemoteBackendError::Cancelled), + } + }) + } +} + struct FailingWriter; impl AsyncWrite for FailingWriter { @@ -129,8 +153,12 @@ impl AsyncWrite for FailingWriter { } fn remote_request(tool: &str) -> RemoteRequest { + remote_request_with_id(tool, RequestId([1; 16])) +} + +fn remote_request_with_id(tool: &str, request_id: RequestId) -> RemoteRequest { RemoteRequest::build( - RequestId([1; 16]), + request_id, WorkspaceSessionId([2; 16]), crate::remote::RemoteBuild::new( WorkspaceRelativePath::new("src").unwrap(), @@ -354,6 +382,56 @@ async fn writer_failure_cancels_hanging_backend_without_waiting_forever() { assert!(result.is_err()); } +#[tokio::test] +async fn admission_rejects_active_excess_without_calling_backend() { + let started = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + let broker = Arc::new( + remote_broker(Arc::new(HoldingBackend { started: started.clone(), release: release.clone() })) + .with_admission_limits(RemoteAdmissionLimits::new(1, 1).unwrap()), + ); + let (first_tx, mut first_rx) = mpsc::channel(8); + let first = tokio::spawn({ + let broker = broker.clone(); + async move { broker.dispatch(remote_request_with_id("make", RequestId([1; 16])), first_tx).await } + }); + started.notified().await; + + let (second_tx, mut second_rx) = mpsc::channel(8); + let error = broker.dispatch(remote_request_with_id("make", RequestId([2; 16])), second_tx).await.unwrap_err(); + assert!(matches!(error, RemoteDispatchError::Backend(RemoteBackendError::Transport { class: crate::remote::RemoteFailureClass::Busy, .. }))); + assert!(matches!(second_rx.recv().await, Some(RemoteBackendEvent::Error { message }) if message.contains("busy"))); + + release.notify_one(); + assert!(first.await.unwrap().is_ok()); + assert_eq!(first_rx.recv().await, Some(RemoteBackendEvent::Stdout(b"active".to_vec()))); + assert_eq!(first_rx.recv().await, Some(RemoteBackendEvent::Completed { exit_code: 0 })); +} + +#[tokio::test] +async fn cancel_acknowledges_separately_and_emits_one_target_terminal() { + let started = Arc::new(Notify::new()); + let release = Arc::new(Notify::new()); + let broker = Arc::new( + remote_broker(Arc::new(HoldingBackend { started: started.clone(), release })) + .with_admission_limits(RemoteAdmissionLimits::new(2, 1).unwrap()), + ); + let (target_tx, mut target_rx) = mpsc::channel(8); + let target = tokio::spawn({ + let broker = broker.clone(); + async move { broker.dispatch(remote_request_with_id("make", RequestId([1; 16])), target_tx).await } + }); + started.notified().await; + assert_eq!(target_rx.recv().await, Some(RemoteBackendEvent::Stdout(b"active".to_vec()))); + + let (cancel_tx, mut cancel_rx) = mpsc::channel(8); + broker.dispatch(RemoteRequest::cancel(RequestId([8; 16]), WorkspaceSessionId([2; 16]), RequestId([1; 16])), cancel_tx).await.unwrap(); + assert_eq!(cancel_rx.recv().await, Some(RemoteBackendEvent::Completed { exit_code: 0 })); + assert_eq!(target_rx.recv().await, Some(RemoteBackendEvent::Cancelled)); + assert!(target_rx.try_recv().is_err()); + assert!(matches!(target.await.unwrap(), Err(RemoteDispatchError::Backend(RemoteBackendError::Cancelled)))); +} + #[test] fn local_passthrough_authorization_remains_separate() { assert!(is_allowed(&["make *".into()], "make", &["--release".into()])); diff --git a/src/guest_install_ut.rs b/src/guest_install_ut.rs index a3ab65e..3ce24da 100644 --- a/src/guest_install_ut.rs +++ b/src/guest_install_ut.rs @@ -21,6 +21,23 @@ fn remote_make_ownership_survives_local_passthrough_install() { assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); } +#[test] +fn remote_cargo_ownership_survives_local_passthrough_install() { + let root = tempfile::tempdir().unwrap(); + let native = root.path().join("native"); + fs::create_dir(&native).unwrap(); + executable(&native.join("cargo")); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_cargo_link(root.path(), &remote, true).unwrap(); + install_vscomm_links(vec!["cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + + assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), remote); +} + #[test] fn disabled_remote_make_preserves_native_and_vscomm_behavior() { let root = tempfile::tempdir().unwrap(); @@ -40,6 +57,24 @@ fn disabled_remote_make_preserves_native_and_vscomm_behavior() { assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); } +#[test] +fn disabled_remote_cargo_preserves_native_behavior() { + let root = tempfile::tempdir().unwrap(); + let native = root.path().join("native"); + fs::create_dir(&native).unwrap(); + executable(&native.join("cargo")); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_cargo_link(root.path(), &remote, true).unwrap(); + install_remote_cargo_link(root.path(), &remote, false).unwrap(); + install_vscomm_links(vec!["cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + + assert!(fs::symlink_metadata(root.path().join("cargo")).is_err()); +} + #[test] fn remote_make_wins_when_native_make_is_present() { let root = tempfile::tempdir().unwrap(); @@ -57,6 +92,23 @@ fn remote_make_wins_when_native_make_is_present() { assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); } +#[test] +fn remote_cargo_wins_when_native_cargo_is_present() { + let root = tempfile::tempdir().unwrap(); + let native = root.path().join("native"); + fs::create_dir(&native).unwrap(); + executable(&native.join("cargo")); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + executable(&remote); + executable(&vscomm); + + install_remote_cargo_link(root.path(), &remote, true).unwrap(); + install_vscomm_links(vec!["cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + + assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), remote); +} + #[test] fn repeated_install_is_idempotent_and_stale_managed_links_are_replaced() { let root = tempfile::tempdir().unwrap(); @@ -80,3 +132,18 @@ fn repeated_install_is_idempotent_and_stale_managed_links_are_replaced() { install_remote_make_link(root.path(), &remote, false).unwrap(); assert!(fs::symlink_metadata(root.path().join("make")).is_err()); } + +#[test] +fn repeated_cargo_install_is_idempotent_and_repairs_stale_managed_links() { + let root = tempfile::tempdir().unwrap(); + let remote = root.path().join("bunkerbox-remote"); + executable(&remote); + + install_remote_cargo_link(root.path(), &remote, true).unwrap(); + install_remote_cargo_link(root.path(), &remote, true).unwrap(); + fs::remove_file(root.path().join("cargo")).unwrap(); + symlink(root.path().join("old/bunkerbox-remote"), root.path().join("cargo")).unwrap(); + install_remote_cargo_link(root.path(), &remote, true).unwrap(); + + assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), remote); +} diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 8f71b52..9989140 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -4,6 +4,8 @@ use crate::remote::{ RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteFailureClass, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTool, RequestId, WorkspaceRelativePath, }; +use std::env; +use std::os::unix::fs::PermissionsExt; use tempfile::TempDir; fn fixture() -> (TempDir, Arc, RemoteTargetId, WorkspaceSessionId) { @@ -457,7 +459,7 @@ async fn timeout_kills_a_direct_child_process() { let backend = LoopbackBackend::new(session, tools).with_timeout(Duration::from_millis(50)); let (events, _receiver) = mpsc::channel(8); let request = authorized_build(target, session_id, "sleep", vec!["5".to_string()], Vec::new(), snapshot_id); - assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Timeout)); + assert_eq!(backend.execute(request, events).await, Err(RemoteBackendError::Deadline { cause: crate::remote::RemoteTimeoutCause::Build })); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] @@ -511,3 +513,96 @@ async fn loopback_missing_required_artifact_is_a_terminal_artifact_failure() { assert!(!events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { .. }))); assert!(fs::read_dir(&session.jobs_root).unwrap().next().is_none()); } + +fn cargo_fixture_session(temp: &TempDir) -> (Arc, RemoteTargetId, WorkspaceSessionId) { + let source = Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/remote-cargo"); + let workspace = temp.path().join("cargo-workspace"); + fs::create_dir_all(workspace.join("src")).unwrap(); + for relative in ["Cargo.toml", "build.rs", "src/main.rs"] { + fs::copy(source.join(relative), workspace.join(relative)).unwrap(); + } + let session_id = WorkspaceSessionId([11; 16]); + let target = RemoteTargetId([12; 16]); + let snapshot_store = SnapshotStore::new(temp.path().join("cargo-snapshots")); + let exclusions = crate::snapshot::SnapshotExclusionPolicy::from_patterns(Vec::::new()).unwrap(); + let builder = SnapshotBuilder::new(snapshot_store.clone(), crate::snapshot::SnapshotLimits::default(), exclusions); + let session = Arc::new(RunRemoteSession::new(session_id, target, workspace, snapshot_store, builder, temp.path().join("cargo-jobs")).unwrap()); + (session, target, session_id) +} + +fn cargo_executable() -> PathBuf { + env::var_os("PATH") + .into_iter() + .flat_map(|path| env::split_paths(&path).collect::>()) + .map(|directory| directory.join("cargo")) + .find(|path| path.is_file() && fs::metadata(path).is_ok_and(|metadata| metadata.permissions().mode() & 0o111 != 0)) + .expect("Cargo must be available on PATH for the remote Cargo fixture") +} + +fn cargo_target_environment() -> BTreeMap { + let mut environment = BTreeMap::new(); + for name in ["PATH", "HOME", "RUSTUP_HOME", "CARGO_HOME"] { + if let Ok(value) = env::var(name) { + environment.insert(name.to_string(), value); + } + } + environment +} + +fn authorized_cargo_build( + target: RemoteTargetId, session: WorkspaceSessionId, request_id: u8, args: Vec, snapshot_id: RemoteSnapshotId, +) -> crate::remote::AuthorizedRemoteRequest { + let build = RemoteBuild::new(WorkspaceRelativePath::new("").unwrap(), RemoteTool::new("cargo").unwrap(), args, Vec::new(), snapshot_id).unwrap(); + let request = RemoteRequest::build(RequestId([request_id; 16]), session, build); + let policy = RemoteAuthorizationPolicy::new(target, session, vec!["cargo".to_string()]).with_snapshot_authority(Arc::new(TestSnapshotAuthority)); + policy.authorize(&RemoteExecutionContext { target, workspace_session_id: session }, request).unwrap() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn cargo_fixture_runs_build_script_nested_cargo_and_trusted_artifact_flow() { + let temp = tempfile::tempdir().unwrap(); + let (session, target, session_id) = cargo_fixture_session(&temp); + let cargo = cargo_executable(); + let mut tools = BTreeMap::new(); + tools.insert("cargo".to_string(), cargo); + let target_environment = cargo_target_environment(); + let artifact = "target/debug/bunkerbox-cargo-fixture-artifact.txt".to_string(); + let backend = LoopbackBackend::new(session.clone(), tools.clone()) + .with_target_environment(target_environment.clone()) + .with_artifacts(ArtifactPolicy::new(vec![artifact.clone()]).unwrap(), ArtifactLimits::default()); + let snapshot_id = sync_capability(&backend, target, session_id).await; + let (events, receiver) = mpsc::channel(64); + assert_eq!( + backend.execute(authorized_cargo_build(target, session_id, 13, vec!["build".into(), "--offline".into()], snapshot_id), events).await, + Ok(()) + ); + let events = collect_events(receiver).await; + let output = events + .iter() + .filter_map(|event| match event { + RemoteBackendEvent::Stdout(bytes) | RemoteBackendEvent::Stderr(bytes) => Some(bytes.as_slice()), + _ => None, + }) + .flatten() + .copied() + .collect::>(); + assert!(String::from_utf8_lossy(&output).contains("bunkerbox-cargo-fixture-build-script")); + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); + assert!(session + .workspace_root() + .join(".bunkerbox/artifacts/0d0d0d0d0d0d0d0d0d0d0d0d0d0d0d0d/target/debug/bunkerbox-cargo-fixture-artifact.txt") + .is_file()); + + let backend = LoopbackBackend::new(session.clone(), tools) + .with_target_environment(target_environment) + .with_artifacts(ArtifactPolicy::new(vec!["target/debug/missing-cargo-artifact".to_string()]).unwrap(), ArtifactLimits::default()); + let snapshot_id = sync_capability(&backend, target, session_id).await; + let (events, receiver) = mpsc::channel(64); + assert!(matches!( + backend.execute(authorized_cargo_build(target, session_id, 14, vec!["build".into(), "--offline".into()], snapshot_id), events).await, + Err(RemoteBackendError::Transport { class: RemoteFailureClass::ArtifactManifest, .. }) + )); + let events = collect_events(receiver).await; + assert!(!events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { .. }))); + assert!(fs::read_dir(&session.jobs_root).unwrap().next().is_none()); +} diff --git a/src/main_ut.rs b/src/main_ut.rs index 66bcb2b..bb656a8 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -59,10 +59,14 @@ fn run_handoff_rejects_zero_session() { } #[test] -fn remote_tool_names_only_enables_the_make_wrapper() { +fn remote_tool_names_only_enables_the_fixed_make_and_cargo_wrappers() { assert_eq!( - remote_tool_names(&[RemoteToolSpec { name: "make".into(), allow_args: true }, RemoteToolSpec { name: "cargo".into(), allow_args: false },]), - vec!["make"] + remote_tool_names(&[ + RemoteToolSpec { name: "make".into(), allow_args: true }, + RemoteToolSpec { name: "cargo".into(), allow_args: false }, + RemoteToolSpec { name: "cmake".into(), allow_args: true }, + ]), + vec!["make", "cargo"] ); } diff --git a/src/remote_client_ut.rs b/src/remote_client_ut.rs index 93d15df..8aa8e50 100644 --- a/src/remote_client_ut.rs +++ b/src/remote_client_ut.rs @@ -16,3 +16,17 @@ fn selected_environment_uses_only_targeted_names() { std::env::remove_var("BB_TEST_REMOTE_ALLOWED"); std::env::remove_var("BB_TEST_REMOTE_PATH"); } + +#[test] +fn cargo_environment_is_empty_even_when_names_are_requested() { + std::env::set_var("BB_TEST_CARGO_FLAGS", "should-not-forward"); + assert!(selected_remote_environment_for_tool("cargo", vec!["BB_TEST_CARGO_FLAGS".into()]).is_empty()); + std::env::remove_var("BB_TEST_CARGO_FLAGS"); +} + +#[test] +fn cancel_request_uses_a_distinct_request_and_target_identity() { + let request = remote_cancel_request(RequestId([8; 16]), WorkspaceSessionId([2; 16]), RequestId([7; 16])); + assert_eq!(request.request_id, RequestId([8; 16])); + assert_eq!(request.operation, crate::vscomm::RemoteOperation::Cancel { target_request_id: RequestId([7; 16]) }); +} diff --git a/src/remote_target_ut.rs b/src/remote_target_ut.rs index 1a666c8..5a08f33 100644 --- a/src/remote_target_ut.rs +++ b/src/remote_target_ut.rs @@ -371,3 +371,34 @@ fn project_artifact_paths_reject_traversal_globs_and_duplicates() { assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); } } + +#[test] +fn lifecycle_admission_and_worker_limits_are_loaded_with_safe_defaults() { + let fixture = Fixture::new(); + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace( + " max-output-bytes: 67108864\n", + " max-output-bytes: 67108864\n idle-output-timeout-seconds: 7\n cleanup-timeout-seconds: 3\n max-active-builds: 2\n max-worker-uploads: 3\n max-worker-upload-bytes: 2M\n max-worker-jobs: 2\n max-worker-job-bytes: 4M\n max-worker-artifact-spools: 2\n max-worker-artifact-spool-bytes: 5M\n max-worker-state-entries: 99\n", + ); + write_config(&fixture, &yaml); + let target = RemoteTargetConfig::load_from(&fixture.config).unwrap().ssh_target("ssh-one").unwrap().resources(); + assert_eq!(target.idle_output_timeout(), Duration::from_secs(7)); + assert_eq!(target.cleanup_timeout(), Duration::from_secs(3)); + assert_eq!(target.max_active_builds(), 2); + assert_eq!(target.worker_state_limits().max_uploads, 3); + assert_eq!(target.worker_state_limits().max_upload_bytes, 2 * 1024 * 1024); + assert_eq!(target.worker_state_limits().max_state_entries, 99); +} + +#[test] +fn lifecycle_and_admission_limits_reject_zero_or_excessive_values() { + let fixture = Fixture::new(); + for replacement in [ + (" connect-timeout-seconds: 5", " connect-timeout-seconds: 0"), + (" max-output-bytes: 67108864", " max-active-builds: 65\n max-output-bytes: 67108864"), + (" max-output-bytes: 67108864", " cleanup-timeout-seconds: 0\n max-output-bytes: 67108864"), + ] { + let yaml = valid_yaml(&fixture, &fixture.project, "ssh", Some("ssh-one")).replace(replacement.0, replacement.1); + write_config(&fixture, &yaml); + assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); + } +} diff --git a/src/remote_ut.rs b/src/remote_ut.rs index b977ba2..f3f17b8 100644 --- a/src/remote_ut.rs +++ b/src/remote_ut.rs @@ -114,6 +114,36 @@ fn environment_policy_rejects_duplicates_and_control_data() { .is_err()); } +#[test] +fn cargo_requires_an_empty_guest_environment() { + let cargo = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("cargo").unwrap(), + vec!["build".into()], + Vec::new(), + RemoteSnapshotId::from_bytes([9; 16]), + ) + .unwrap(); + let authorized = + policy(vec!["cargo".into()]).authorize(&context(), RemoteRequest::build(RequestId([1; 16]), WorkspaceSessionId([2; 16]), cargo)).unwrap(); + let RemoteOperation::Build(build) = authorized.request().operation() else { panic!("expected build") }; + assert!(build.env().is_empty()); + + let cargo_with_rustflags = RemoteBuild::new( + WorkspaceRelativePath::new("src").unwrap(), + RemoteTool::new("cargo").unwrap(), + vec!["build".into()], + vec![("RUSTFLAGS".into(), "-C opt-level=3".into())], + RemoteSnapshotId::from_bytes([9; 16]), + ) + .unwrap(); + assert_eq!( + policy(vec!["cargo".into()]) + .authorize(&context(), RemoteRequest::build(RequestId([1; 16]), WorkspaceSessionId([2; 16]), cargo_with_rustflags)), + Err(RemoteAuthorizationError::ForbiddenEnvironment("RUSTFLAGS".into())) + ); +} + #[test] fn command_policy_rejects_unapproved_arguments() { let policy = RemoteAuthorizationPolicy::from_policies( @@ -125,4 +155,34 @@ fn command_policy_rejects_unapproved_arguments() { .unwrap() .with_snapshot_authority(std::sync::Arc::new(TestSnapshotAuthority)); assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::ToolArgumentsNotAllowed("make".into()))); + + let cargo_policy = RemoteAuthorizationPolicy::from_policies( + RemoteTargetId([3; 16]), + WorkspaceSessionId([2; 16]), + [("cargo".into(), RemoteToolPolicy::new(false))], + RemoteEnvironmentPolicy::default(), + ) + .unwrap() + .with_snapshot_authority(std::sync::Arc::new(TestSnapshotAuthority)); + assert_eq!(cargo_policy.authorize(&context(), request("cargo")), Err(RemoteAuthorizationError::ToolArgumentsNotAllowed("cargo".into()))); +} + +#[test] +fn request_and_cancel_ids_must_be_nonzero() { + let policy = policy(vec!["make".into()]); + assert_eq!( + policy.authorize(&context(), RemoteRequest::sync(RequestId([0; 16]), WorkspaceSessionId([2; 16]))), + Err(RemoteAuthorizationError::InvalidRequestId) + ); + assert_eq!( + policy.authorize(&context(), RemoteRequest::cancel(RequestId([8; 16]), WorkspaceSessionId([2; 16]), RequestId([0; 16]))), + Err(RemoteAuthorizationError::InvalidCancelTarget) + ); +} + +#[test] +fn admission_limits_have_safe_bounds() { + assert_eq!(RemoteAdmissionLimits::new(2, 1).unwrap(), RemoteAdmissionLimits { global_active: 2, target_active: 1 }); + assert!(RemoteAdmissionLimits::new(0, 1).is_err()); + assert!(RemoteAdmissionLimits::new(RemoteAdmissionLimits::MAX + 1, 1).is_err()); } diff --git a/src/snapshot_ut.rs b/src/snapshot_ut.rs index 26cf80a..b9f408e 100644 --- a/src/snapshot_ut.rs +++ b/src/snapshot_ut.rs @@ -1,6 +1,6 @@ use super::*; use crate::cfg::{ProjectConfig, ProjectSection, RemoteSection}; -use crate::remote::WorkspaceSessionId; +use crate::remote::{RemoteExecutionControl, WorkspaceSessionId}; use std::fs; use std::os::unix::fs::{symlink, PermissionsExt}; use std::os::unix::net::UnixListener; @@ -48,6 +48,22 @@ fn write_file(root: &Path, path: &str, contents: &[u8]) { fs::write(path, contents).unwrap(); } +#[test] +fn cancellation_stops_snapshot_creation_and_materialization_before_publication() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "result", b"snapshot"); + let builder = builder(&store, SnapshotLimits::default(), &[]); + let control = RemoteExecutionControl::new(); + control.cancel(); + assert!(builder.build_root_with_control(source.path(), session(1), control.clone()).is_err()); + + let snapshot = builder.build_root(source.path(), session(1)).unwrap(); + let destination = store.path().join("materialized"); + assert!(builder.store.materialize_with_control(&snapshot.handle, &destination, control).is_err()); + assert!(!destination.exists()); +} + #[test] fn snapshots_nested_modified_and_untracked_files_with_modes() { let source = TempDir::new().unwrap(); diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index 42c00d3..14bc92f 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -399,7 +399,11 @@ async fn scripted_worker_proves_end_to_end_ssh_transport_and_exact_build_result( assert!(joined.contains("IdentityAgent=none")); assert!(!joined.contains("literal $(arg)")); assert!(!joined.contains("snapshot_id")); - assert_eq!(spec.remote_command(), "exec '/usr/local/libexec/bunkerbox-worker' --stdio --workspace-root '/var/tmp/bunkerbox-workers'"); + assert!(spec + .remote_command() + .starts_with("exec '/usr/local/libexec/bunkerbox-worker' --stdio --workspace-root '/var/tmp/bunkerbox-workers'")); + assert!(spec.remote_command().contains("--build-timeout-ms 5000")); + assert!(spec.remote_command().contains("--max-worker-uploads 2")); } } diff --git a/src/vscomm/mod_ut.rs b/src/vscomm/mod_ut.rs index 2712a11..b1da0df 100644 --- a/src/vscomm/mod_ut.rs +++ b/src/vscomm/mod_ut.rs @@ -100,6 +100,28 @@ fn diagnostic_remote_sync_does_not_retain_capability() { assert!(!sync.retain_capability); } +#[test] +fn remote_cancel_round_trips_as_operation_kind_three() { + let (request_id, session_id) = ids(); + let request = RemoteRequest::cancel(request_id, session_id, RequestId([7; 16])); + let frame = request.to_frame().unwrap(); + assert_eq!(frame.payload[6], 3); + assert_eq!(RemoteRequest::from_frame(frame).unwrap(), request); + let domain = request.into_domain().unwrap(); + assert_eq!(domain.operation(), &crate::remote::RemoteOperation::Cancel { target_request_id: crate::remote::RequestId([7; 16]) }); +} + +#[test] +fn remote_cancel_rejects_zero_target() { + let (request_id, session_id) = ids(); + assert!(RemoteRequest::cancel(request_id, session_id, RequestId([0; 16])).to_frame().is_err()); +} + +#[test] +fn remote_event_rejects_zero_request_id() { + assert!(RemoteEvent { request_id: RequestId([0; 16]), kind: RemoteEventKind::Cancelled }.to_frame().is_err()); +} + #[test] fn invalid_remote_sync_capability_flag_is_rejected() { let (request_id, session_id) = ids(); From 8c46b8955ad8ccba5d403db1441f0e206facd10f Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 11:53:10 +0200 Subject: [PATCH 39/52] Refactor generic remote client --- src/bin/bunkerbox-remote.rs | 27 +++--- src/cfg.rs | 4 +- src/guest_install.rs | 181 +++++++++++++++++++++++++++++------- src/main.rs | 2 +- src/remote.rs | 10 ++ src/remote_client.rs | 21 ++++- 6 files changed, 194 insertions(+), 51 deletions(-) diff --git a/src/bin/bunkerbox-remote.rs b/src/bin/bunkerbox-remote.rs index a93da01..58092f6 100644 --- a/src/bin/bunkerbox-remote.rs +++ b/src/bin/bunkerbox-remote.rs @@ -1,15 +1,15 @@ -use bunkerbox::guest_install::{install_remote_cargo_link, install_remote_make_link}; +use bunkerbox::guest_install::{is_managed_remote_wrapper, synchronize_remote_wrappers}; +use bunkerbox::remote::validate_remote_wrapper_name; #[cfg(test)] use bunkerbox::remote::RemoteSnapshotId; use bunkerbox::remote_client::{ - execute_remote_request_to, logical_workspace_cwd, new_request_id, remote_build_request, remote_diagnostic_sync_request, remote_environment_names, - remote_session_from_env, remote_sync_request, remote_tool_enabled, selected_remote_environment_for_tool, RemoteCompletion, + configured_remote_tools, execute_remote_request_to, logical_workspace_cwd, new_request_id, remote_build_request, remote_diagnostic_sync_request, + remote_environment_names, remote_session_from_env, remote_sync_request, selected_remote_environment_for_tool, RemoteCompletion, }; #[cfg(test)] use bunkerbox::vscomm::RequestId; use bunkerbox::vscomm::{RemoteRequest, WorkspaceSessionId, TOOLCHAIN_PORT, VSCOMM_BIN_DIR}; use std::env; -use std::fs; use std::io::{self, Read, Write}; use std::mem; use std::path::Path; @@ -37,11 +37,8 @@ fn run() -> Result { env::args_os().next().and_then(|value| Path::new(&value).file_name().and_then(|name| name.to_str()).map(str::to_owned)).unwrap_or_default(); let args = env::args().skip(1).collect::>(); - if matches!(invoked_as.as_str(), "make" | "cargo") { - return run_transparent_tool(&invoked_as, &args); - } if invoked_as != "bunkerbox-remote" { - return Err("bunkerbox-remote must be invoked directly or through a managed make or cargo symlink".to_string()); + return run_transparent_tool(&invoked_as, &args); } if args.len() == 1 && args[0] == "install" { install_remote_links()?; @@ -70,9 +67,16 @@ fn run_explicit(args: &[String]) -> Result { } fn run_transparent_tool(tool: &str, args: &[String]) -> Result { + let tool = validate_remote_wrapper_name(tool.to_string())?; + if !configured_remote_tools()?.contains(&tool) { + return Err(format!("remote tool is not configured: {tool}")); + } + if !is_managed_remote_wrapper(Path::new(VSCOMM_BIN_DIR), &tool)? { + return Err(format!("remote wrapper is not managed: {tool}")); + } let cwd = logical_workspace_cwd(&env::current_dir().map_err(|error| format!("current directory: {error}"))?)?; let session = remote_session_from_env()?; - run_build_with_sync(cwd, tool.to_string(), args.to_vec(), session) + run_build_with_sync(cwd, tool, args.to_vec(), session) } fn run_build_with_sync(cwd: String, tool: String, args: Vec, session: WorkspaceSessionId) -> Result { @@ -138,11 +142,10 @@ fn build_request( } fn install_remote_links() -> Result<(), String> { - fs::create_dir_all(VSCOMM_BIN_DIR).map_err(|error| format!("mkdir {VSCOMM_BIN_DIR}: {error}"))?; let executable = env::current_exe().map_err(|error| format!("failed to locate remote binary: {error}"))?; let bin_dir = Path::new(VSCOMM_BIN_DIR); - install_remote_make_link(bin_dir, &executable, remote_tool_enabled("make"))?; - install_remote_cargo_link(bin_dir, &executable, remote_tool_enabled("cargo")) + let vscomm_path = bin_dir.join("bunkerbox-vscomm"); + synchronize_remote_wrappers(configured_remote_tools()?, bin_dir, &executable, &vscomm_path) } fn connect_toolchain() -> Result { diff --git a/src/cfg.rs b/src/cfg.rs index 3fb1807..d1f8199 100644 --- a/src/cfg.rs +++ b/src/cfg.rs @@ -4,7 +4,7 @@ use std::path::{Path, PathBuf}; use serde::Deserialize; -use crate::remote::{RemoteEnvironmentPolicy, RemoteTool}; +use crate::remote::{validate_remote_wrapper_name, RemoteEnvironmentPolicy}; use crate::snapshot::SnapshotExclusionPolicy; use crate::vscomm::buildsys::{self, PassthroughMode}; @@ -295,7 +295,7 @@ impl ProjectConfig { SnapshotExclusionPolicy::from_patterns(self.project.remote.exclude.clone())?; let mut tools = std::collections::BTreeSet::new(); for tool in &self.project.remote.tools { - RemoteTool::new(tool.name.clone())?; + validate_remote_wrapper_name(tool.name.clone())?; if !tools.insert(tool.name.clone()) { return Err(format!("duplicate remote tool: {}", tool.name)); } diff --git a/src/guest_install.rs b/src/guest_install.rs index 2663103..9a2571e 100644 --- a/src/guest_install.rs +++ b/src/guest_install.rs @@ -1,42 +1,67 @@ +pub use crate::remote::{validate_remote_wrapper_name, REMOTE_WRAPPER_STATE_FILE}; +use std::collections::BTreeSet; use std::ffi::OsStr; -use std::fs; -use std::os::unix::fs::{symlink, PermissionsExt}; +use std::fs::{self, File, OpenOptions}; +use std::io::{Read, Write}; +use std::os::unix::fs::{symlink, MetadataExt, OpenOptionsExt, PermissionsExt}; use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; const REMOTE_WRAPPER_OWNER: &str = "bunkerbox-remote"; +const VSCOMM_OWNER: &str = "bunkerbox-vscomm"; +const MAX_REMOTE_WRAPPER_STATE_BYTES: u64 = 64 * 1024; +static NEXT_STATE_TEMP: AtomicU64 = AtomicU64::new(1); -pub fn install_remote_make_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { - install_remote_link("make", bin_dir, executable, enabled) -} - -pub fn install_remote_cargo_link(bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { - install_remote_link("cargo", bin_dir, executable, enabled) -} +pub fn synchronize_remote_wrappers( + commands: impl IntoIterator, bin_dir: &Path, executable: &Path, vscomm_path: &Path, +) -> Result<(), String> { + fs::create_dir_all(bin_dir).map_err(|error| format!("mkdir {}: {error}", bin_dir.display()))?; + let (previous, _had_state) = read_managed_wrappers(bin_dir)?; + let current = collect_wrapper_names(commands)?; -fn install_remote_link(command: &str, bin_dir: &Path, executable: &Path, enabled: bool) -> Result<(), String> { - if !is_supported_remote_command(command) { - return Err(format!("unsupported remote wrapper command: {command}")); + for command in previous.difference(¤t) { + let target = bin_dir.join(command); + if is_remote_wrapper_target(&target, executable) { + fs::remove_file(&target).map_err(|error| format!("remove disabled remote wrapper {command}: {error}"))?; + } } - let target = bin_dir.join(command); - let managed = is_managed_remote_link(&target); - if enabled { - if let Ok(link) = fs::read_link(&target) { - if link == executable { - return Ok(()); + for command in ¤t { + let target = bin_dir.join(command); + let Some(link) = read_link_if_present(&target)? else { + symlink(executable, &target).map_err(|error| format!("symlink remote wrapper {command}: {error}"))?; + continue; + }; + + if link == executable { + if !previous.contains(command) { + return Err(format!("cannot install remote wrapper over unmanaged {}", target.display())); } + continue; } - if fs::symlink_metadata(&target).is_ok() { - if !managed { - return Err(format!("cannot install remote {command} wrapper over existing {}", target.display())); + + let recognized_remote = is_remote_link(&link, executable) && previous.contains(command); + let recognized_vscomm = link == vscomm_path || is_stale_vscomm_link(&link, vscomm_path); + if !recognized_remote && !recognized_vscomm { + if previous.contains(command) { + let mut ownership = current.clone(); + ownership.remove(command); + let _ = write_managed_wrappers(bin_dir, &ownership); } - fs::remove_file(&target).map_err(|error| format!("remove existing remote {command} wrapper: {error}"))?; + return Err(format!("cannot install remote wrapper over existing {}", target.display())); } - symlink(executable, &target).map_err(|error| format!("symlink remote {command} wrapper: {error}"))?; - } else if managed { - fs::remove_file(&target).map_err(|error| format!("remove disabled remote {command} wrapper: {error}"))?; + fs::remove_file(&target).map_err(|error| format!("remove existing remote wrapper {command}: {error}"))?; + symlink(executable, &target).map_err(|error| format!("symlink remote wrapper {command}: {error}"))?; } - Ok(()) + + write_managed_wrappers(bin_dir, ¤t) +} + +pub fn is_managed_remote_wrapper(bin_dir: &Path, command: &str) -> Result { + let command = validate_remote_wrapper_name(command.to_string())?; + let (managed, had_state) = read_managed_wrappers(bin_dir)?; + let expected = bin_dir.join(REMOTE_WRAPPER_OWNER); + Ok(had_state && managed.contains(&command) && is_remote_wrapper_target(&bin_dir.join(command), &expected)) } pub fn install_vscomm_links(commands: impl IntoIterator, bin_dir: &Path, vscomm_path: &Path, path: &str) -> Result<(), String> { @@ -46,8 +71,11 @@ pub fn install_vscomm_links(commands: impl IntoIterator, bin_dir: if command.is_empty() { continue; } + if command == "bunkerbox" || command.starts_with("bunkerbox-") || command.starts_with(REMOTE_WRAPPER_STATE_FILE) { + return Err(format!("vscomm command name is reserved: {command}")); + } let target = bin_dir.join(&command); - if is_supported_remote_command(&command) && is_managed_remote_link(&target) { + if validate_remote_wrapper_name(command.clone()).is_ok() && is_managed_remote_wrapper(bin_dir, &command)? { continue; } if command_exists_in_path_except(&command, vscomm_path, path) { @@ -67,16 +95,101 @@ pub fn install_vscomm_links(commands: impl IntoIterator, bin_dir: Ok(()) } -fn is_managed_remote_link(target: &Path) -> bool { - let Ok(metadata) = fs::symlink_metadata(target) else { return false }; - if !metadata.file_type().is_symlink() { - return false; +fn collect_wrapper_names(commands: impl IntoIterator) -> Result, String> { + let mut names = BTreeSet::new(); + for command in commands { + let command = validate_remote_wrapper_name(command)?; + if !names.insert(command.clone()) { + return Err(format!("duplicate remote wrapper name: {command}")); + } } - fs::read_link(target).ok().and_then(|link| link.file_name().map(OsStr::to_owned)).is_some_and(|name| name == REMOTE_WRAPPER_OWNER) + Ok(names) +} + +fn read_managed_wrappers(bin_dir: &Path) -> Result<(BTreeSet, bool), String> { + let path = bin_dir.join(REMOTE_WRAPPER_STATE_FILE); + let Some(mut file) = open_private_state(&path)? else { return Ok((BTreeSet::new(), false)) }; + let length = file.metadata().map_err(|error| format!("stat remote wrapper state: {error}"))?.len(); + if length > MAX_REMOTE_WRAPPER_STATE_BYTES { + return Err("remote wrapper state is too large".to_string()); + } + let mut contents = String::new(); + file.read_to_string(&mut contents).map_err(|error| format!("read remote wrapper state: {error}"))?; + let mut names = BTreeSet::new(); + for line in contents.lines() { + let name = validate_remote_wrapper_name(line.to_string())?; + if !names.insert(name.clone()) { + return Err(format!("duplicate remote wrapper state name: {name}")); + } + } + Ok((names, true)) +} + +fn write_managed_wrappers(bin_dir: &Path, names: &BTreeSet) -> Result<(), String> { + let mut temporary = None; + let mut file = None; + for _ in 0..32 { + let sequence = NEXT_STATE_TEMP.fetch_add(1, Ordering::Relaxed); + let path = bin_dir.join(format!("{REMOTE_WRAPPER_STATE_FILE}.tmp-{}-{sequence}", std::process::id())); + let result = OpenOptions::new().write(true).create_new(true).mode(0o600).open(&path); + match result { + Ok(value) => { + temporary = Some(path); + file = Some(value); + break; + } + Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {} + Err(error) => return Err(format!("create remote wrapper state: {error}")), + } + } + let temporary = temporary.ok_or_else(|| "could not create remote wrapper state temporary file".to_string())?; + let mut file = file.expect("remote wrapper state file exists with its temporary path"); + let contents = names.iter().map(|name| format!("{name}\n")).collect::(); + let result = (|| { + file.write_all(contents.as_bytes()).map_err(|error| format!("write remote wrapper state: {error}"))?; + file.sync_all().map_err(|error| format!("sync remote wrapper state: {error}"))?; + fs::rename(&temporary, bin_dir.join(REMOTE_WRAPPER_STATE_FILE)).map_err(|error| format!("publish remote wrapper state: {error}")) + })(); + if result.is_err() { + let _ = fs::remove_file(&temporary); + } + result +} + +fn open_private_state(path: &Path) -> Result, String> { + let mut options = OpenOptions::new(); + options.read(true).custom_flags(libc::O_NOFOLLOW); + let file = match options.open(path) { + Ok(file) => file, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(format!("open remote wrapper state: {error}")), + }; + let metadata = file.metadata().map_err(|error| format!("stat remote wrapper state: {error}"))?; + if !metadata.file_type().is_file() || metadata.permissions().mode() & 0o077 != 0 || metadata.uid() != unsafe { libc::geteuid() } { + return Err("remote wrapper state is not a private owned regular file".to_string()); + } + Ok(Some(file)) +} + +fn read_link_if_present(path: &Path) -> Result, String> { + match fs::read_link(path) { + Ok(link) => Ok(Some(link)), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None), + Err(error) if error.kind() == std::io::ErrorKind::InvalidInput => Err(format!("remote wrapper entry is not a symlink: {}", path.display())), + Err(error) => Err(format!("read remote wrapper entry {}: {error}", path.display())), + } +} + +fn is_remote_wrapper_target(path: &Path, expected: &Path) -> bool { + fs::read_link(path).ok().is_some_and(|link| is_remote_link(&link, expected)) +} + +fn is_remote_link(link: &Path, expected: &Path) -> bool { + link == expected || link.file_name().is_some_and(|name| name == OsStr::new(REMOTE_WRAPPER_OWNER)) } -fn is_supported_remote_command(command: &str) -> bool { - matches!(command, "make" | "cargo") +fn is_stale_vscomm_link(link: &Path, expected: &Path) -> bool { + link == expected || link.file_name().is_some_and(|name| name == OsStr::new(VSCOMM_OWNER)) } fn command_exists_in_path_except(command: &str, except: &Path, path: &str) -> bool { diff --git a/src/main.rs b/src/main.rs index e6aa1ad..e058d7e 100644 --- a/src/main.rs +++ b/src/main.rs @@ -451,7 +451,7 @@ fn new_target_id() -> bunkerbox::remote::RemoteTargetId { } fn remote_tool_names(entries: &[RemoteToolSpec]) -> Vec { - entries.iter().filter(|tool| matches!(tool.name.as_str(), "make" | "cargo")).map(|tool| tool.name.clone()).collect() + entries.iter().map(|tool| tool.name.clone()).collect() } fn write_run_handoff(file: &mut File, path: &Path, session_id: vscomm::WorkspaceSessionId) -> Result<(), String> { diff --git a/src/remote.rs b/src/remote.rs index 13c8ea5..31aed9e 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -19,6 +19,7 @@ pub const MAX_REMOTE_ENV_VALUE_BYTES: usize = 4 * 1024; pub const DEFAULT_REMOTE_BUILD_TIMEOUT: Duration = Duration::from_secs(30); pub const DEFAULT_REMOTE_OUTPUT_BYTES: u64 = 64 * 1024 * 1024; pub const DEFAULT_REMOTE_ENVIRONMENT: &[&str] = &["CC", "CXX", "AR", "RUSTFLAGS", "CFLAGS", "CXXFLAGS", "MAKEFLAGS"]; +pub const REMOTE_WRAPPER_STATE_FILE: &str = ".bunkerbox-remote-tools"; #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct RequestId(pub [u8; 16]); @@ -97,6 +98,15 @@ impl RemoteTool { } } +pub fn validate_remote_wrapper_name(name: impl Into) -> Result { + let name = name.into(); + RemoteTool::new(name.clone())?; + if name == "bunkerbox" || name.starts_with("bunkerbox-") || name.starts_with(REMOTE_WRAPPER_STATE_FILE) { + return Err(format!("remote wrapper name is reserved: {name}")); + } + Ok(name) +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct RemoteBuild { cwd: WorkspaceRelativePath, diff --git a/src/remote_client.rs b/src/remote_client.rs index cb026e1..5403c77 100644 --- a/src/remote_client.rs +++ b/src/remote_client.rs @@ -1,6 +1,7 @@ -use crate::remote::RemoteSnapshotId; +use crate::remote::{validate_remote_wrapper_name, RemoteSnapshotId}; use crate::vscomm::{RemoteBuild, RemoteEvent, RemoteEventKind, RemoteRequest, RemoteTool, RequestId, WorkspaceRelativePath, WorkspaceSessionId}; use rand::RngCore; +use std::collections::BTreeSet; use std::env; use std::io::{Read, Write}; use std::path::Path; @@ -33,7 +34,23 @@ pub fn new_request_id() -> RequestId { } pub fn remote_tool_enabled(tool: &str) -> bool { - env::var("BUNKERBOX_REMOTE_TOOLS").ok().is_some_and(|tools| tools.split(',').any(|candidate| candidate == tool)) + configured_remote_tools().ok().is_some_and(|tools| tools.iter().any(|candidate| candidate == tool)) +} + +pub fn configured_remote_tools() -> Result, String> { + let Some(raw) = env::var_os("BUNKERBOX_REMOTE_TOOLS") else { return Ok(Vec::new()) }; + let raw = raw.into_string().map_err(|_| "BUNKERBOX_REMOTE_TOOLS is not valid UTF-8".to_string())?; + let mut tools = BTreeSet::new(); + for candidate in raw.split(',') { + if candidate.is_empty() { + return Err("BUNKERBOX_REMOTE_TOOLS contains an empty tool name".to_string()); + } + let candidate = validate_remote_wrapper_name(candidate.to_string())?; + if !tools.insert(candidate.clone()) { + return Err(format!("BUNKERBOX_REMOTE_TOOLS contains duplicate tool: {candidate}")); + } + } + Ok(tools.into_iter().collect()) } pub fn remote_environment_names() -> Vec { From 09ae61e6fd304622e7062ba868589f23d2546fed Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 11:53:19 +0200 Subject: [PATCH 40/52] Update unit tests --- src/bunkerbox-remote_ut.rs | 50 ++++++------- src/cfg_ut.rs | 16 +++++ src/guest_install_ut.rs | 143 ++++++++++++++++++++----------------- src/loopback_ut.rs | 16 +++++ src/main_ut.rs | 4 +- src/remote_client_ut.rs | 11 +++ src/remote_ut.rs | 1 + 7 files changed, 143 insertions(+), 98 deletions(-) diff --git a/src/bunkerbox-remote_ut.rs b/src/bunkerbox-remote_ut.rs index 292ceb5..166549b 100644 --- a/src/bunkerbox-remote_ut.rs +++ b/src/bunkerbox-remote_ut.rs @@ -58,6 +58,14 @@ fn parses_cargo_toolchain_and_arguments_without_joining() { ); } +#[test] +fn parses_arbitrary_tool_and_arguments_without_build_system_logic() { + assert_eq!( + parse_command(&["build".into(), "build-my-car".into(), "--variant".into(), "pink unicorn".into()]), + Ok(RemoteCommand::Build { tool: "build-my-car".into(), args: vec!["--variant".into(), "pink unicorn".into()] }) + ); +} + #[test] fn build_request_preserves_logical_cwd_and_arguments() { let request = build_request( @@ -200,46 +208,30 @@ fn logical_cwd_is_workspace_relative_only() { } #[test] -fn configured_remote_make_installation_precedes_native_path_resolution() { - let root = tempfile::tempdir().unwrap(); - let executable = root.path().join("bunkerbox-remote"); - std::fs::write(&executable, b"remote").unwrap(); - install_remote_make_link(root.path(), &executable, true).unwrap(); - assert_eq!(std::fs::read_link(root.path().join("make")).unwrap(), executable); -} - -#[test] -fn disabled_remote_make_removes_only_its_managed_link() { +fn generic_remote_installation_preserves_unmanaged_entries() { let root = tempfile::tempdir().unwrap(); let executable = root.path().join("bunkerbox-remote"); std::fs::write(&executable, b"remote").unwrap(); - install_remote_make_link(root.path(), &executable, true).unwrap(); - install_remote_make_link(root.path(), &executable, false).unwrap(); - assert!(!root.path().join("make").exists()); + std::fs::write(root.path().join("build-my-car"), b"native").unwrap(); + let vscomm = root.path().join("bunkerbox-vscomm"); + std::fs::write(&vscomm, b"vscomm").unwrap(); - std::fs::write(root.path().join("make"), b"native").unwrap(); - install_remote_make_link(root.path(), &executable, false).unwrap(); - assert_eq!(std::fs::read(root.path().join("make")).unwrap(), b"native"); + assert!(synchronize_remote_wrappers(vec!["build-my-car".into()], root.path(), &executable, &vscomm).is_err()); + assert_eq!(std::fs::read(root.path().join("build-my-car")).unwrap(), b"native"); } #[test] -fn configured_remote_make_does_not_overwrite_unmanaged_entry() { +fn generic_remote_installation_removes_only_managed_wrappers() { let root = tempfile::tempdir().unwrap(); let executable = root.path().join("bunkerbox-remote"); std::fs::write(&executable, b"remote").unwrap(); - std::fs::write(root.path().join("make"), b"native").unwrap(); - assert!(install_remote_make_link(root.path(), &executable, true).is_err()); - assert_eq!(std::fs::read(root.path().join("make")).unwrap(), b"native"); -} + let vscomm = root.path().join("bunkerbox-vscomm"); + std::fs::write(&vscomm, b"vscomm").unwrap(); -#[test] -fn configured_remote_cargo_does_not_overwrite_unmanaged_entry() { - let root = tempfile::tempdir().unwrap(); - let executable = root.path().join("bunkerbox-remote"); - std::fs::write(&executable, b"remote").unwrap(); - std::fs::write(root.path().join("cargo"), b"native").unwrap(); - assert!(install_remote_cargo_link(root.path(), &executable, true).is_err()); - assert_eq!(std::fs::read(root.path().join("cargo")).unwrap(), b"native"); + synchronize_remote_wrappers(vec!["make".into(), "cargo".into()], root.path(), &executable, &vscomm).unwrap(); + synchronize_remote_wrappers(vec!["cargo".into()], root.path(), &executable, &vscomm).unwrap(); + assert!(!root.path().join("make").exists()); + assert!(is_managed_remote_wrapper(root.path(), "cargo").unwrap()); } #[test] diff --git a/src/cfg_ut.rs b/src/cfg_ut.rs index 4181843..8e6d8f7 100644 --- a/src/cfg_ut.rs +++ b/src/cfg_ut.rs @@ -546,6 +546,22 @@ fn load_or_create_validates_remote_policy_configuration() { assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "cargo".into(), allow_args: true }]); } +#[test] +fn load_or_create_accepts_arbitrary_remote_tool_names_but_rejects_control_names() { + let root = TempDir::new().unwrap(); + write_project_conf( + root.path(), + "project:\n remote:\n tools:\n - name: build-my-car\n allow-args: true\n - name: ninja+debug\n allow-args: false\n", + ); + let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); + assert_eq!(cfg.project.remote.tools.len(), 2); + + for name in ["bunkerbox", "bunkerbox-remote", ".bunkerbox-remote-tools"] { + write_project_conf(root.path(), &format!("project:\n remote:\n tools:\n - name: {name}\n allow-args: true\n")); + assert!(ProjectConfig::load_or_create(root.path()).is_err(), "accepted reserved remote tool: {name}"); + } +} + #[test] fn load_or_create_rejects_forbidden_remote_environment() { let root = TempDir::new().unwrap(); diff --git a/src/guest_install_ut.rs b/src/guest_install_ut.rs index 3ce24da..d5fdd81 100644 --- a/src/guest_install_ut.rs +++ b/src/guest_install_ut.rs @@ -1,149 +1,158 @@ use super::*; -use std::os::unix::fs::PermissionsExt; +use std::os::unix::fs::{MetadataExt, PermissionsExt}; fn executable(path: &Path) { fs::write(path, b"binary").unwrap(); fs::set_permissions(path, fs::Permissions::from_mode(0o755)).unwrap(); } +fn synchronize(root: &Path, names: &[&str], remote: &Path, vscomm: &Path) { + synchronize_remote_wrappers(names.iter().map(|name| (*name).to_string()), root, remote, vscomm).unwrap(); +} + #[test] -fn remote_make_ownership_survives_local_passthrough_install() { +fn configured_arbitrary_wrappers_survive_vscomm_install() { let root = tempfile::tempdir().unwrap(); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); executable(&vscomm); - install_remote_make_link(root.path(), &remote, true).unwrap(); - install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, "").unwrap(); + synchronize(root.path(), &["make", "cargo", "build-my-car"], &remote, &vscomm); + install_vscomm_links(vec!["make".into(), "cargo".into(), "build-my-car".into()], root.path(), &vscomm, "").unwrap(); - assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); - assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); + for name in ["make", "cargo", "build-my-car"] { + assert_eq!(fs::read_link(root.path().join(name)).unwrap(), remote); + assert!(is_managed_remote_wrapper(root.path(), name).unwrap()); + } } #[test] -fn remote_cargo_ownership_survives_local_passthrough_install() { +fn removing_a_configured_tool_removes_only_its_managed_wrapper() { let root = tempfile::tempdir().unwrap(); - let native = root.path().join("native"); - fs::create_dir(&native).unwrap(); - executable(&native.join("cargo")); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); executable(&vscomm); - install_remote_cargo_link(root.path(), &remote, true).unwrap(); - install_vscomm_links(vec!["cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + synchronize(root.path(), &["make", "build-my-car"], &remote, &vscomm); + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); - assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), remote); + assert!(fs::symlink_metadata(root.path().join("make")).is_err()); + assert_eq!(fs::read_link(root.path().join("build-my-car")).unwrap(), remote); + assert!(!is_managed_remote_wrapper(root.path(), "make").unwrap()); + assert!(is_managed_remote_wrapper(root.path(), "build-my-car").unwrap()); } #[test] -fn disabled_remote_make_preserves_native_and_vscomm_behavior() { +fn unmanaged_entries_are_never_overwritten() { let root = tempfile::tempdir().unwrap(); - let native = root.path().join("native"); - fs::create_dir(&native).unwrap(); - let native_make = native.join("make"); - executable(&native_make); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); executable(&vscomm); + fs::write(root.path().join("build-my-car"), b"unmanaged").unwrap(); - install_remote_make_link(root.path(), &remote, false).unwrap(); - install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); - - assert_eq!(fs::read_link(root.path().join("make")).unwrap_err().kind(), std::io::ErrorKind::NotFound); - assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); + assert!(synchronize_remote_wrappers(vec!["build-my-car".into()], root.path(), &remote, &vscomm).is_err()); + assert_eq!(fs::read(root.path().join("build-my-car")).unwrap(), b"unmanaged"); } #[test] -fn disabled_remote_cargo_preserves_native_behavior() { +fn stale_managed_links_are_repaired_and_repeated_install_is_idempotent() { let root = tempfile::tempdir().unwrap(); - let native = root.path().join("native"); - fs::create_dir(&native).unwrap(); - executable(&native.join("cargo")); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); executable(&vscomm); - install_remote_cargo_link(root.path(), &remote, true).unwrap(); - install_remote_cargo_link(root.path(), &remote, false).unwrap(); - install_vscomm_links(vec!["cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); + fs::remove_file(root.path().join("build-my-car")).unwrap(); + symlink(root.path().join("old/bunkerbox-remote"), root.path().join("build-my-car")).unwrap(); + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); - assert!(fs::symlink_metadata(root.path().join("cargo")).is_err()); + assert_eq!(fs::read_link(root.path().join("build-my-car")).unwrap(), remote); } #[test] -fn remote_make_wins_when_native_make_is_present() { +fn vscomm_to_remote_transition_is_allowed_only_for_the_configured_name() { let root = tempfile::tempdir().unwrap(); - let native = root.path().join("native"); - fs::create_dir(&native).unwrap(); - executable(&native.join("make")); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); executable(&vscomm); - install_remote_make_link(root.path(), &remote, true).unwrap(); - install_vscomm_links(vec!["make".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + install_vscomm_links(vec!["build-my-car".into()], root.path(), &vscomm, "").unwrap(); + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); - assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); + assert_eq!(fs::read_link(root.path().join("build-my-car")).unwrap(), remote); } #[test] -fn remote_cargo_wins_when_native_cargo_is_present() { +fn changed_state_tracked_link_becomes_unmanaged_and_is_not_removed() { let root = tempfile::tempdir().unwrap(); - let native = root.path().join("native"); - fs::create_dir(&native).unwrap(); - executable(&native.join("cargo")); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); + let other = root.path().join("other"); executable(&remote); executable(&vscomm); + executable(&other); - install_remote_cargo_link(root.path(), &remote, true).unwrap(); - install_vscomm_links(vec!["cargo".into()], root.path(), &vscomm, native.to_str().unwrap()).unwrap(); + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); + fs::remove_file(root.path().join("build-my-car")).unwrap(); + symlink(&other, root.path().join("build-my-car")).unwrap(); + assert!(synchronize_remote_wrappers(vec!["build-my-car".into()], root.path(), &remote, &vscomm).is_err()); - assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), remote); + assert_eq!(fs::read_link(root.path().join("build-my-car")).unwrap(), other); + assert!(!is_managed_remote_wrapper(root.path(), "build-my-car").unwrap()); } #[test] -fn repeated_install_is_idempotent_and_stale_managed_links_are_replaced() { +fn invalid_and_reserved_wrapper_names_fail_closed() { let root = tempfile::tempdir().unwrap(); let remote = root.path().join("bunkerbox-remote"); let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); executable(&vscomm); - install_remote_make_link(root.path(), &remote, true).unwrap(); - install_remote_make_link(root.path(), &remote, true).unwrap(); - install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, "").unwrap(); - install_vscomm_links(vec!["make".into(), "cargo".into()], root.path(), &vscomm, "").unwrap(); - assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); - assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), vscomm); - - fs::remove_file(root.path().join("make")).unwrap(); - symlink(root.path().join("old/bunkerbox-remote"), root.path().join("make")).unwrap(); - install_remote_make_link(root.path(), &remote, true).unwrap(); - assert_eq!(fs::read_link(root.path().join("make")).unwrap(), remote); - - install_remote_make_link(root.path(), &remote, false).unwrap(); - assert!(fs::symlink_metadata(root.path().join("make")).is_err()); + for name in ["", ".", "..", "../escape", "bunkerbox", "bunkerbox-remote", REMOTE_WRAPPER_STATE_FILE] { + assert!(validate_remote_wrapper_name(name.to_string()).is_err(), "accepted reserved or invalid wrapper: {name}"); + assert!(synchronize_remote_wrappers(vec![name.to_string()], root.path(), &remote, &vscomm).is_err()); + } } #[test] -fn repeated_cargo_install_is_idempotent_and_repairs_stale_managed_links() { +fn no_state_means_existing_remote_link_is_unmanaged() { let root = tempfile::tempdir().unwrap(); let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); executable(&remote); + executable(&vscomm); + symlink(&remote, root.path().join("build-my-car")).unwrap(); - install_remote_cargo_link(root.path(), &remote, true).unwrap(); - install_remote_cargo_link(root.path(), &remote, true).unwrap(); - fs::remove_file(root.path().join("cargo")).unwrap(); - symlink(root.path().join("old/bunkerbox-remote"), root.path().join("cargo")).unwrap(); - install_remote_cargo_link(root.path(), &remote, true).unwrap(); + assert!(synchronize_remote_wrappers(vec!["build-my-car".into()], root.path(), &remote, &vscomm).is_err()); + assert_eq!(fs::read_link(root.path().join("build-my-car")).unwrap(), remote); +} - assert_eq!(fs::read_link(root.path().join("cargo")).unwrap(), remote); +#[test] +fn wrapper_state_is_private_and_symlink_state_is_rejected() { + let root = tempfile::tempdir().unwrap(); + let remote = root.path().join("bunkerbox-remote"); + let vscomm = root.path().join("bunkerbox-vscomm"); + let outside = root.path().join("outside"); + executable(&remote); + executable(&vscomm); + fs::write(&outside, b"outside").unwrap(); + + synchronize(root.path(), &["build-my-car"], &remote, &vscomm); + let state = root.path().join(REMOTE_WRAPPER_STATE_FILE); + let metadata = fs::metadata(&state).unwrap(); + assert!(metadata.file_type().is_file()); + assert_eq!(metadata.permissions().mode() & 0o077, 0); + assert_eq!(metadata.uid(), unsafe { libc::geteuid() }); + + fs::remove_file(&state).unwrap(); + symlink(&outside, &state).unwrap(); + assert!(synchronize_remote_wrappers(Vec::new(), root.path(), &remote, &vscomm).is_err()); + assert_eq!(fs::read_link(state).unwrap(), outside); } diff --git a/src/loopback_ut.rs b/src/loopback_ut.rs index 9989140..de78db9 100644 --- a/src/loopback_ut.rs +++ b/src/loopback_ut.rs @@ -116,6 +116,22 @@ async fn sync_and_build_materialize_a_bound_snapshot() { assert!(fs::read_dir(&session_state.jobs_root).unwrap().next().is_none()); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn arbitrary_logical_tool_uses_its_trusted_target_mapping() { + let (_temp, session, target, session_id) = fixture(); + let mut tools = resolve_fixed_tools(["printf".to_string()]); + let Some(printf) = tools.remove("printf") else { return }; + tools.insert("build-my-car".into(), printf); + let backend = LoopbackBackend::new(session, tools); + let snapshot_id = sync_capability(&backend, target, session_id).await; + let (events, receiver) = mpsc::channel(8); + let request = authorized_build(target, session_id, "build-my-car", vec!["custom $(argument)".into()], Vec::new(), snapshot_id); + assert_eq!(backend.execute(request, events).await, Ok(())); + let events = collect_events(receiver).await; + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Stdout(bytes) if bytes == b"custom $(argument)"))); + assert!(events.iter().any(|event| matches!(event, RemoteBackendEvent::Completed { exit_code: 0 }))); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn diagnostic_sync_releases_capability_and_snapshot_storage() { let (_temp, session, target, session_id) = fixture(); diff --git a/src/main_ut.rs b/src/main_ut.rs index bb656a8..c2b3624 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -59,14 +59,14 @@ fn run_handoff_rejects_zero_session() { } #[test] -fn remote_tool_names_only_enables_the_fixed_make_and_cargo_wrappers() { +fn remote_tool_names_propagates_the_complete_configured_set() { assert_eq!( remote_tool_names(&[ RemoteToolSpec { name: "make".into(), allow_args: true }, RemoteToolSpec { name: "cargo".into(), allow_args: false }, RemoteToolSpec { name: "cmake".into(), allow_args: true }, ]), - vec!["make", "cargo"] + vec!["make", "cargo", "cmake"] ); } diff --git a/src/remote_client_ut.rs b/src/remote_client_ut.rs index 8aa8e50..85aff97 100644 --- a/src/remote_client_ut.rs +++ b/src/remote_client_ut.rs @@ -24,6 +24,17 @@ fn cargo_environment_is_empty_even_when_names_are_requested() { std::env::remove_var("BB_TEST_CARGO_FLAGS"); } +#[test] +fn configured_remote_tools_are_generic_validated_and_deduplicated() { + std::env::set_var("BUNKERBOX_REMOTE_TOOLS", "make,build-my-car,cargo"); + assert_eq!(configured_remote_tools().unwrap(), vec!["build-my-car", "cargo", "make"]); + std::env::set_var("BUNKERBOX_REMOTE_TOOLS", "make,make"); + assert!(configured_remote_tools().is_err()); + std::env::set_var("BUNKERBOX_REMOTE_TOOLS", "bunkerbox-status"); + assert!(configured_remote_tools().is_err()); + std::env::remove_var("BUNKERBOX_REMOTE_TOOLS"); +} + #[test] fn cancel_request_uses_a_distinct_request_and_target_identity() { let request = remote_cancel_request(RequestId([8; 16]), WorkspaceSessionId([2; 16]), RequestId([7; 16])); diff --git a/src/remote_ut.rs b/src/remote_ut.rs index f3f17b8..3e548d7 100644 --- a/src/remote_ut.rs +++ b/src/remote_ut.rs @@ -47,6 +47,7 @@ fn policy_rejects_unapproved_tool() { let policy = policy(vec!["cargo".into()]); assert_eq!(policy.authorize(&context(), request("make")), Err(RemoteAuthorizationError::ToolNotAllowed("make".into()))); + assert_eq!(policy.authorize(&context(), request("build-my-car")), Err(RemoteAuthorizationError::ToolNotAllowed("build-my-car".into()))); } #[test] From 0b35a493b362d0435a219c88f76236ed43b030e0 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 14:45:01 +0200 Subject: [PATCH 41/52] Update worker protocol to change targets --- crates/bunkerbox-worker-protocol/src/lib.rs | 48 ++++++++++++- .../bunkerbox-worker-protocol/src/lib_ut.rs | 25 ++++++- crates/bunkerbox-worker/src/process.rs | 67 ++++++++++++++++--- crates/bunkerbox-worker/src/process_ut.rs | 37 ++++++++++ crates/bunkerbox-worker/src/worker.rs | 4 +- 5 files changed, 167 insertions(+), 14 deletions(-) diff --git a/crates/bunkerbox-worker-protocol/src/lib.rs b/crates/bunkerbox-worker-protocol/src/lib.rs index 89c1b5a..757c7d8 100644 --- a/crates/bunkerbox-worker-protocol/src/lib.rs +++ b/crates/bunkerbox-worker-protocol/src/lib.rs @@ -22,6 +22,8 @@ pub const WORKER_MAGIC: [u8; 4] = WORKER_PROTOCOL_MAGIC; pub const WORKER_PROTOCOL_VERSION: u16 = 1; pub const WORKER_VERSION: u16 = WORKER_PROTOCOL_VERSION; pub const WORKER_ARTIFACT_PROTOCOL_VERSION: u16 = 2; +/// Adds trusted target-PATH command resolution without changing V1/V2 fields. +pub const WORKER_COMMAND_PROTOCOL_VERSION: u16 = 3; pub const WORKER_FRAME_HEADER_LEN: usize = 4 + 2 + 1 + 4; pub const WORKER_FRAME_HEADER_SIZE: usize = WORKER_FRAME_HEADER_LEN; pub const WORKER_ID_LEN: usize = 16; @@ -525,6 +527,7 @@ impl WorkerUploadEntry { pub struct WorkerBuild { pub tool: WorkerTool, pub trusted_executable: WorkerExecutablePath, + pub target_command: Option, pub argv: Vec, pub cwd: WorkerRelativePath, pub guest_env: Vec<(String, String)>, @@ -543,6 +546,28 @@ impl WorkerBuild { let build = Self { tool: WorkerTool::new(tool)?, trusted_executable: WorkerExecutablePath::new(trusted_executable)?, + target_command: None, + argv, + cwd: WorkerRelativePath::new(cwd)?, + guest_env, + target_env, + upload_token, + artifact_paths: Vec::new(), + artifact_max_file_bytes: MAX_WORKER_ARTIFACT_FILE_BYTES, + artifact_max_total_bytes: MAX_WORKER_ARTIFACT_TOTAL_BYTES, + }; + build.validate()?; + Ok(build) + } + + pub fn new_command( + tool: impl Into, command: impl Into, argv: Vec, cwd: impl Into, guest_env: Vec<(String, String)>, + target_env: Vec<(String, String)>, upload_token: WorkerUploadId, + ) -> WorkerResult { + let build = Self { + tool: WorkerTool::new(tool)?, + trusted_executable: WorkerExecutablePath::new("/usr/bin/bunkerbox-command")?, + target_command: Some(WorkerTool::new(command)?), argv, cwd: WorkerRelativePath::new(cwd)?, guest_env, @@ -589,6 +614,10 @@ impl WorkerBuild { self.trusted_executable.as_str() } + pub fn target_command(&self) -> Option<&WorkerTool> { + self.target_command.as_ref() + } + pub fn argv(&self) -> &[String] { &self.argv } @@ -916,6 +945,9 @@ impl WorkerMessage { if version == WORKER_PROTOCOL_VERSION && self.requires_artifact_version() { return Err(invalid("worker artifact message requires artifact-capable protocol version")); } + if version < WORKER_COMMAND_PROTOCOL_VERSION && matches!(self, Self::Build { build, .. } if build.target_command.is_some()) { + return Err(invalid("worker command identity requires command-capable protocol version")); + } self.validate()?; let mut payload = WireWriter::new(); encode_payload(self, &mut payload, version)?; @@ -1358,7 +1390,11 @@ fn decode_artifact_entry(reader: &mut WireReader<'_>) -> WorkerResult WorkerResult<()> { build.validate()?; writer.string(build.tool.as_str(), MAX_WORKER_TOOL_BYTES, "worker tool")?; - writer.string(build.trusted_executable.as_str(), MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?; + if version >= WORKER_COMMAND_PROTOCOL_VERSION { + writer.string(build.target_command.as_ref().map_or(build.tool.as_str(), WorkerTool::as_str), MAX_WORKER_TOOL_BYTES, "worker command")?; + } else { + writer.string(build.trusted_executable.as_str(), MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?; + } writer.count(build.argv.len(), MAX_WORKER_ARG_COUNT, "worker argv")?; for argument in &build.argv { writer.string(argument, MAX_WORKER_ARG_BYTES, "worker argument")?; @@ -1380,7 +1416,12 @@ fn encode_build(writer: &mut WireWriter, build: &WorkerBuild, version: u16) -> W fn decode_build(reader: &mut WireReader<'_>, version: u16) -> WorkerResult { let tool = WorkerTool::new(reader.string(MAX_WORKER_TOOL_BYTES, "worker tool")?)?; - let trusted_executable = WorkerExecutablePath::new(reader.string(MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?)?; + let (trusted_executable, target_command) = if version >= WORKER_COMMAND_PROTOCOL_VERSION { + let command = WorkerTool::new(reader.string(MAX_WORKER_TOOL_BYTES, "worker command")?)?; + (WorkerExecutablePath::new("/usr/bin/bunkerbox-command")?, Some(command)) + } else { + (WorkerExecutablePath::new(reader.string(MAX_WORKER_EXECUTABLE_PATH_BYTES, "worker executable path")?)?, None) + }; let argument_count = reader.count(MAX_WORKER_ARG_COUNT, "worker argv")?; let mut argv = Vec::with_capacity(argument_count); for _ in 0..argument_count { @@ -1403,6 +1444,7 @@ fn decode_build(reader: &mut WireReader<'_>, version: u16) -> WorkerResult) -> WorkerProtocolError { } fn is_supported_version(version: u16) -> bool { - matches!(version, WORKER_PROTOCOL_VERSION | WORKER_ARTIFACT_PROTOCOL_VERSION) + matches!(version, WORKER_PROTOCOL_VERSION | WORKER_ARTIFACT_PROTOCOL_VERSION | WORKER_COMMAND_PROTOCOL_VERSION) } fn validate_version(version: u16) -> WorkerResult<()> { diff --git a/crates/bunkerbox-worker-protocol/src/lib_ut.rs b/crates/bunkerbox-worker-protocol/src/lib_ut.rs index bff8c92..6ea7f74 100644 --- a/crates/bunkerbox-worker-protocol/src/lib_ut.rs +++ b/crates/bunkerbox-worker-protocol/src/lib_ut.rs @@ -70,6 +70,19 @@ fn all_message_variants_round_trip() { } } +#[test] +fn command_identity_round_trips_only_in_protocol_v3() { + let (request_id, session_id, upload_id) = ids(); + let build = WorkerBuild::new_command("make", "gmake", vec!["all".into()], "", Vec::new(), Vec::new(), upload_id).unwrap(); + let message = WorkerMessage::Build { request_id, session_id, build: build.clone() }; + assert!(message.encode_version(WORKER_PROTOCOL_VERSION).is_err()); + + let frame = message.encode_version(WORKER_COMMAND_PROTOCOL_VERSION).unwrap(); + let (version, decoded) = WorkerMessage::decode_versioned(&frame).unwrap(); + assert_eq!(version, WORKER_COMMAND_PROTOCOL_VERSION); + assert_eq!(decoded, message); +} + #[test] fn blocking_helpers_preserve_the_async_wire_encoding() { let message = WorkerMessage::Stdout { request_id: ids().0, session_id: ids().1, data: b"blocking parity".to_vec() }; @@ -299,7 +312,7 @@ fn unknown_nested_kinds_and_flags_are_rejected() { } #[test] -fn artifact_messages_round_trip_only_in_protocol_v2() { +fn artifact_messages_round_trip_in_artifact_capable_protocols() { let (request_id, session_id, _upload_id) = ids(); let artifact_set_id = WorkerArtifactSetId([4; 16]); let artifact = WorkerArtifactEntry::new("dist/result", 0o755, 4, [8; 32]).unwrap(); @@ -319,6 +332,16 @@ fn artifact_messages_round_trip_only_in_protocol_v2() { assert_eq!(version, WORKER_ARTIFACT_PROTOCOL_VERSION); assert_eq!(decoded, message); assert!(message.encode().is_err()); + let v3 = message.encode_version(WORKER_COMMAND_PROTOCOL_VERSION).unwrap(); + let decoded_v3 = WorkerMessage::decode_versioned(&v3).unwrap().1; + match (message, decoded_v3) { + (WorkerMessage::Build { build: original, .. }, WorkerMessage::Build { build: decoded, .. }) => { + assert_eq!(decoded.tool(), original.tool()); + assert_eq!(decoded.target_command(), Some(original.tool())); + assert_eq!(decoded.artifact_paths(), original.artifact_paths()); + } + (original, decoded) => assert_eq!(decoded, original), + } } let v1_build = WorkerMessage::Build { request_id, session_id, build: build() }; diff --git a/crates/bunkerbox-worker/src/process.rs b/crates/bunkerbox-worker/src/process.rs index 08f64e1..483371f 100644 --- a/crates/bunkerbox-worker/src/process.rs +++ b/crates/bunkerbox-worker/src/process.rs @@ -1,11 +1,13 @@ use crate::platform; use crate::storage; use bunkerbox_worker_protocol::{WorkerBuild, WorkerMessage, WorkerRequestId, WorkerSessionId}; +use std::ffi::OsString; use std::fs::{self, File}; use std::io::{self, Read}; use std::os::fd::AsRawFd; use std::os::unix::fs::PermissionsExt; use std::os::unix::process::CommandExt; +use std::path::PathBuf; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::mpsc::{self, RecvTimeoutError}; use std::sync::Arc; @@ -128,17 +130,38 @@ pub fn execute_build_with_limits( job: &JobWorkspace, build: &WorkerBuild, request_id: WorkerRequestId, session_id: WorkerSessionId, sink: &S, disconnected: &dyn Fn() -> bool, build_timeout: Duration, max_output_bytes: u64, ) -> Result { - validate_executable(build.trusted_executable_path())?; + let command_path = build.target_command().map(|command| resolve_target_command(command.as_str())).transpose()?; + if command_path.is_none() { + validate_executable(build.trusted_executable_path())?; + } let cwd = storage::open_relative_directory(job.root(), build.cwd().as_str())?; let cwd_fd = cwd.as_raw_fd(); - let mut command = std::process::Command::new(build.trusted_executable_path()); + let mut command = std::process::Command::new(command_path.as_deref().unwrap_or_else(|| std::path::Path::new(build.trusted_executable_path()))); command.args(build.argv()).stdin(std::process::Stdio::null()).stdout(std::process::Stdio::piped()).stderr(std::process::Stdio::piped()); - command.env_clear(); - for (key, value) in build.guest_env() { - command.env(key, value); - } - for (key, value) in build.target_env() { - command.env(key, value); + if build.target_command().is_some() { + let baseline = worker_environment(); + command.env_clear(); + for (key, value) in baseline { + command.env(key, value); + } + for (key, value) in build.target_env() { + if !protected_environment_name(key) { + command.env(key, value); + } + } + for (key, value) in build.guest_env() { + if !protected_environment_name(key) { + command.env(key, value); + } + } + } else { + command.env_clear(); + for (key, value) in build.guest_env() { + command.env(key, value); + } + for (key, value) in build.target_env() { + command.env(key, value); + } } unsafe { command.pre_exec(move || { @@ -340,6 +363,34 @@ fn spawn_pump( }) } +fn resolve_target_command(command: &str) -> Result { + let path = std::env::var_os("PATH").ok_or_else(|| "worker target PATH is unavailable".to_string())?; + for directory in std::env::split_paths(&path) { + let directory = if directory.as_os_str().is_empty() { PathBuf::from(".") } else { directory }; + let candidate = if directory.is_absolute() { + directory.join(command) + } else { + std::env::current_dir().map_err(|error| format!("resolve worker target PATH: {error}"))?.join(directory).join(command) + }; + if let Ok(metadata) = fs::metadata(&candidate) { + if metadata.file_type().is_file() && metadata.permissions().mode() & 0o111 != 0 { + return Ok(candidate); + } + } + } + Err(format!("worker target command is not executable in PATH: {command}")) +} + +fn protected_environment_name(name: &str) -> bool { + matches!(name, "PATH" | "HOME" | "TMPDIR" | "LANG") || name.starts_with("LC_") +} + +fn worker_environment() -> Vec<(OsString, OsString)> { + std::env::vars_os() + .filter(|(key, _)| key.to_str().is_some_and(|key| matches!(key, "PATH" | "HOME" | "TMPDIR" | "LANG") || key.starts_with("LC_"))) + .collect() +} + fn validate_executable(path: &str) -> Result<(), String> { let metadata = fs::metadata(path).map_err(|error| format!("inspect worker executable: {error}"))?; if !metadata.file_type().is_file() { diff --git a/crates/bunkerbox-worker/src/process_ut.rs b/crates/bunkerbox-worker/src/process_ut.rs index 870aceb..350aa28 100644 --- a/crates/bunkerbox-worker/src/process_ut.rs +++ b/crates/bunkerbox-worker/src/process_ut.rs @@ -100,3 +100,40 @@ fn configured_build_timeout_overrides_worker_default() { .unwrap_err(); assert!(error.contains("timed out")); } + +#[test] +fn command_identity_uses_worker_path_and_protects_baseline_environment() { + let temp = tempdir().unwrap(); + fs::set_permissions(temp.path(), fs::Permissions::from_mode(0o700)).unwrap(); + let jobs = platform::open_root(temp.path()).unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + let build = WorkerBuild::new_command( + "logical-shell", + "sh", + vec!["-c".to_string(), "printf '%s\\n%s\\n%s' \"$PATH\" \"$HOME\" \"$CHECK\"".to_string()], + "", + vec![ + ("PATH".to_string(), "/guest-path".to_string()), + ("HOME".to_string(), "/guest-home".to_string()), + ("CHECK".to_string(), "guest".to_string()), + ], + vec![("PATH".to_string(), "/target-path".to_string()), ("CHECK".to_string(), "target".to_string())], + WorkerUploadId([11; 16]), + ) + .unwrap(); + let writer = FrameWriter::new(Vec::new()); + assert_eq!(execute_build(&job, &build, WorkerRequestId([12; 16]), WorkerSessionId([13; 16]), &writer, &|| false).unwrap(), 0); + + let bytes = writer.into_inner().unwrap(); + let mut reader = Cursor::new(bytes); + let mut output = Vec::new(); + while let Some(message) = WorkerMessage::read_blocking_optional(&mut reader).unwrap() { + if let WorkerMessage::Stdout { data, .. } = message { + output.extend_from_slice(&data); + } + } + let output = String::from_utf8(output).unwrap(); + assert!(!output.contains("/guest-path")); + assert!(!output.contains("/target-path")); + assert!(output.ends_with("guest")); +} diff --git a/crates/bunkerbox-worker/src/worker.rs b/crates/bunkerbox-worker/src/worker.rs index 9d29267..c7be00a 100644 --- a/crates/bunkerbox-worker/src/worker.rs +++ b/crates/bunkerbox-worker/src/worker.rs @@ -2,7 +2,7 @@ use crate::process::{self, JobWorkspace, OutputSink}; use crate::storage::{ArtifactSpool, UploadStore, UploadTransaction}; use bunkerbox_worker_protocol::{ WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadId, MAX_WORKER_CHUNK_BYTES, - MAX_WORKER_ERROR_BYTES, WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, + MAX_WORKER_ERROR_BYTES, WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_COMMAND_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, }; use std::io::{self, Read, Write}; use std::sync::atomic::{AtomicBool, AtomicU16, AtomicUsize, Ordering}; @@ -86,7 +86,7 @@ impl WorkerService { if version != hello_version { return Err("worker Hello version does not match frame version".to_string()); } - if version != WORKER_PROTOCOL_VERSION && version != WORKER_ARTIFACT_PROTOCOL_VERSION { + if version != WORKER_PROTOCOL_VERSION && version != WORKER_ARTIFACT_PROTOCOL_VERSION && version != WORKER_COMMAND_PROTOCOL_VERSION { return Err(format!("unsupported worker protocol version: {version}")); } if session_id.0 == [0; 16] { From cbd0ad4472f9f07b2a62e98384dcac0651fde512 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 14:45:21 +0200 Subject: [PATCH 42/52] Update config docs on configuring remote targets --- docs/config/project.md | 29 ++++++++++++++++++++++++ docs/reference/config-schema.md | 39 +++++++++++++++++++++++++++++++++ 2 files changed, 68 insertions(+) diff --git a/docs/config/project.md b/docs/config/project.md index b15bba9..eef4008 100644 --- a/docs/config/project.md +++ b/docs/config/project.md @@ -184,6 +184,35 @@ If your project has both a `Cargo.toml` and a `Makefile`, both show up. Detection only runs when the list is empty — once you've added a command, you're in full control. +### `remote` + +The optional `remote` section controls transparent remote wrappers. It is +host-owned policy; it does not give the guest permission to choose a target. + +```yaml +project: + remote: + exclude: [target/] + environment: [CC, CXX] + tools: + - name: make + allow-args: true + - name: cargo + allow-args: true + artifacts: [build/app] +``` + +`environment`, `tools`, and `artifacts` are validated allowlists. Artifact +paths are normalized workspace-relative paths and never come from a guest +request. `command` is an optional basename-only target override, such as +`make` to `gmake`; absolute paths and shell strings are rejected. + +Remote targets are declared separately in `.bunkerbox/remote.conf`. If that +file is absent, Bunkerbox starts with the implicit `localhost` target only. +The host TUI starts on `localhost`; use `Ctrl+Alt+B` to select a remote target. +The guest and AI have no target-selection command, and a remote failure never +falls back to local execution. + ## Legacy migration Older versions of Bunkerbox used `.bunkerbox/env.conf` with a flat structure. diff --git a/docs/reference/config-schema.md b/docs/reference/config-schema.md index 462d8c3..9d39a31 100644 --- a/docs/reference/config-schema.md +++ b/docs/reference/config-schema.md @@ -89,6 +89,17 @@ project: - "make *" - "cargo *" + # Optional host-owned transparent remote policy. + remote: + exclude: [target/] + environment: [CC, CXX] + tools: + - name: make + allow-args: true + - name: cargo + allow-args: true + artifacts: [build/app] + # Override shared runtime defaults (optional, uncomment to use): # image: # workspace: direct @@ -99,6 +110,34 @@ project: # - extra.api.example.com ``` +## Project-local remote.conf + +Remote targets are optional and live at `.bunkerbox/remote.conf`. The file +contains only compact target definitions. SSH aliases, identities, host-key +state, and credentials remain in the host OpenSSH configuration. + +```yaml +targets: + netbsd: + ssh: builder@netbsd-builder:2222 + workspace: /var/tmp/bunkerbox + project: + remote: + tools: + - name: make + command: gmake + allow-args: true + resources: + build-timeout-seconds: 3600 + max-active-builds: 1 +``` + +`ssh` accepts `[user@]host[:port]` or a host-side OpenSSH alias. `workspace` +must be an absolute normalized target path. `localhost` is implicit, always +listed first, and selected at startup. Press `Ctrl+Alt+B` in the host TUI to +choose another target; `Ctrl+Alt+H` shows the host controls. Target selection +is frozen per transaction, and remote failures do not retry locally. + ## Sandbox profile Profiles are host-side YAML files selected by the project configuration. From 2851cc2be38f1764d6d928afbc4b0e5ce6c67db2 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 14:45:35 +0200 Subject: [PATCH 43/52] Add TUI to switch remote targets --- src/cfg.rs | 40 +++- src/cfg_ut.rs | 13 +- src/daemon.rs | 335 +++++++++++++++++++++++++-- src/daemon_ut.rs | 14 +- src/main.rs | 160 ++++++------- src/main_ut.rs | 6 +- src/remote.rs | 30 ++- src/remote_target.rs | 485 ++++++++++++++++++++++++++++++++++++++++ src/remote_target_ut.rs | 61 +++++ src/remote_ut.rs | 16 ++ src/ssh.rs | 112 ++++++---- src/ssh_ut.rs | 26 ++- src/tui.rs | 250 +++++++++++++++++++-- src/tui_ut.rs | 26 ++- 14 files changed, 1392 insertions(+), 182 deletions(-) diff --git a/src/cfg.rs b/src/cfg.rs index d1f8199..5f84cf3 100644 --- a/src/cfg.rs +++ b/src/cfg.rs @@ -189,7 +189,7 @@ pub struct AppliedRuntime { pub encrypt: Option>, } -#[derive(Debug, Default, serde::Serialize, serde::Deserialize)] +#[derive(Clone, Debug, Default, serde::Serialize, serde::Deserialize)] pub struct ProjectConfig { #[serde(default)] pub project: ProjectSection, @@ -199,7 +199,7 @@ pub struct ProjectConfig { pub profiles: Vec, } -#[derive(Debug, Default, serde::Serialize, serde::Deserialize)] +#[derive(Clone, Debug, Default, serde::Serialize, serde::Deserialize)] pub struct ProjectSection { #[serde(default)] pub env: EnvMode, @@ -213,7 +213,8 @@ pub struct ProjectSection { pub remote: RemoteSection, } -#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)] +#[derive(Clone, Debug, Default, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] pub struct RemoteSection { #[serde(default)] pub exclude: Vec, @@ -221,16 +222,21 @@ pub struct RemoteSection { pub environment: Vec, #[serde(default)] pub tools: Vec, + #[serde(default)] + pub artifacts: Vec, } -#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)] +#[serde(deny_unknown_fields)] pub struct RemoteToolSpec { pub name: String, + #[serde(default)] + pub command: Option, #[serde(default, rename = "allow-args")] pub allow_args: bool, } -#[derive(Debug, Default, serde::Serialize, serde::Deserialize)] +#[derive(Clone, Debug, Default, serde::Serialize, serde::Deserialize)] pub struct ImageOverrides { #[serde(default)] pub workspace: Option, @@ -296,10 +302,14 @@ impl ProjectConfig { let mut tools = std::collections::BTreeSet::new(); for tool in &self.project.remote.tools { validate_remote_wrapper_name(tool.name.clone())?; + if let Some(command) = &tool.command { + validate_remote_wrapper_name(command.clone())?; + } if !tools.insert(tool.name.clone()) { return Err(format!("duplicate remote tool: {}", tool.name)); } } + crate::artifact::ArtifactPolicy::new(self.project.remote.artifacts.clone())?; Ok(()) } @@ -396,7 +406,11 @@ impl ProjectConfig { } } - if !self.project.remote.exclude.is_empty() || !self.project.remote.environment.is_empty() || !self.project.remote.tools.is_empty() { + if !self.project.remote.exclude.is_empty() + || !self.project.remote.environment.is_empty() + || !self.project.remote.tools.is_empty() + || !self.project.remote.artifacts.is_empty() + { y.push_str(" remote:\n"); y.push_str(" exclude:\n"); if self.project.remote.exclude.is_empty() { @@ -419,7 +433,19 @@ impl ProjectConfig { y.push_str(" []\n"); } else { for tool in &self.project.remote.tools { - y.push_str(&format!(" - name: \"{}\"\n allow-args: {}\n", tool.name, tool.allow_args)); + y.push_str(&format!(" - name: \"{}\"\n", tool.name)); + if let Some(command) = &tool.command { + y.push_str(&format!(" command: \"{command}\"\n")); + } + y.push_str(&format!(" allow-args: {}\n", tool.allow_args)); + } + } + y.push_str(" artifacts:\n"); + if self.project.remote.artifacts.is_empty() { + y.push_str(" []\n"); + } else { + for artifact in &self.project.remote.artifacts { + y.push_str(&format!(" - \"{artifact}\"\n")); } } } diff --git a/src/cfg_ut.rs b/src/cfg_ut.rs index 8e6d8f7..0952df4 100644 --- a/src/cfg_ut.rs +++ b/src/cfg_ut.rs @@ -539,11 +539,20 @@ fn load_or_create_validates_remote_policy_configuration() { ); let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); assert_eq!(cfg.project.remote.environment, vec!["PROJECT_MODE"]); - assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "make".into(), allow_args: false }]); + assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "make".into(), command: None, allow_args: false }]); write_project_conf(root.path(), "project:\n remote:\n tools:\n - name: cargo\n allow-args: true\n"); let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); - assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "cargo".into(), allow_args: true }]); + assert_eq!(cfg.project.remote.tools, vec![RemoteToolSpec { name: "cargo".into(), command: None, allow_args: true }]); + + write_project_conf(root.path(), "project:\n remote:\n tools:\n - name: make\n command: gmake\n allow-args: true\n"); + let cfg = ProjectConfig::load_or_create(root.path()).unwrap(); + assert_eq!(cfg.project.remote.tools[0].command.as_deref(), Some("gmake")); + write_project_conf( + root.path(), + "project:\n remote:\n tools:\n - name: make\n command: /usr/bin/make\n allow-args: true\n", + ); + assert!(ProjectConfig::load_or_create(root.path()).is_err()); } #[test] diff --git a/src/daemon.rs b/src/daemon.rs index c7396b0..9d76216 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -8,11 +8,12 @@ use crate::remote::{ RemoteEnvironmentPolicy, RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteTargetId, RemoteToolPolicy, }; use crate::remote_target::SshTarget; +use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; use crate::sandbox::{resolve_profile, MergedProfile, NetworkMode}; use crate::ssh::SshBackend; use crate::vscomm::{validate_exec_request, validate_process_path, ExecRequest, Frame, FrameType, TOOLCHAIN_PORT}; use crate::workspace::WorkspaceCwd; -use rand::Rng; +use rand::{Rng, RngCore}; use std::collections::BTreeMap; use std::fs::File; use std::io::{BufRead, BufReader}; @@ -212,7 +213,7 @@ impl RemoteExecutionRegistry { for entry in entries.values() { let _ = entry.request_cancel(); } - } + }; } } @@ -438,7 +439,9 @@ struct VsockSession { workspace: PathBuf, merged_profile: Option>, proxy_config: Option>, - remote_broker: Arc, + remote_router: Arc, + local_session_id: crate::remote::WorkspaceSessionId, + local_capabilities: Mutex>, } pub struct VsockDaemon { @@ -448,6 +451,106 @@ pub struct VsockDaemon { sandbox_proxy_dir: Option, } +struct RemoteRouter { + active_target: ActiveBuildTarget, + brokers: BTreeMap>, + snapshots: Mutex>, + requests: Mutex>>, +} + +impl RemoteRouter { + fn new(active_target: ActiveBuildTarget, brokers: BTreeMap>) -> Self { + Self { active_target, brokers, snapshots: Mutex::new(BTreeMap::new()), requests: Mutex::new(BTreeMap::new()) } + } + + fn active_remote_broker(&self) -> Result<(String, Arc), RemoteDispatchError> { + let label = self.active_target.current(); + let broker = self.brokers.get(&label).cloned().ok_or(RemoteDispatchError::Unauthorized(RemoteAuthorizationError::TargetNotAllowed))?; + Ok((label, broker)) + } + + async fn dispatch(&self, request: RemoteRequest, writer: &mut W) -> Result<(), String> { + let (broker, label) = match request.operation() { + crate::remote::RemoteOperation::Sync(_) => { + let (label, broker) = self.active_remote_broker().map_err(|error| format!("remote dispatch failed: {error:?}"))?; + (broker, Some(label)) + } + crate::remote::RemoteOperation::Build(build) => { + let label = self + .snapshots + .lock() + .map_err(|_| "remote target snapshot lock poisoned".to_string())? + .get(&build.snapshot_id()) + .cloned() + .ok_or_else(|| "remote snapshot target binding is unavailable".to_string())?; + let broker = self.brokers.get(&label).cloned().ok_or_else(|| "remote target binding is unavailable".to_string())?; + (broker, Some(label)) + } + crate::remote::RemoteOperation::Cancel { target_request_id } => { + let broker = self + .requests + .lock() + .map_err(|_| "remote target request lock poisoned".to_string())? + .get(target_request_id) + .cloned() + .ok_or_else(|| "remote cancellation target is unavailable".to_string())?; + (broker, None) + } + }; + + let request_id = request.request_id(); + if !matches!(request.operation(), crate::remote::RemoteOperation::Cancel { .. }) { + self.requests.lock().map_err(|_| "remote target request lock poisoned".to_string())?.insert(request_id, broker.clone()); + } + let (event_tx, mut event_rx) = mpsc::channel(64); + let mut dispatch = Box::pin(broker.dispatch(request.clone(), event_tx)); + let result = loop { + tokio::select! { + dispatch_result = &mut dispatch => { + while let Ok(event) = event_rx.try_recv() { + self.forward_event(request_id, label.as_deref(), &request, event, writer).await?; + } + break dispatch_result; + } + event = event_rx.recv() => { + let Some(event) = event else { break Err(RemoteDispatchError::EventSinkClosed) }; + self.forward_event(request_id, label.as_deref(), &request, event, writer).await?; + } + } + }; + if matches!(request.operation(), crate::remote::RemoteOperation::Build(_)) { + if let crate::remote::RemoteOperation::Build(build) = request.operation() { + self.snapshots.lock().map_err(|_| "remote target snapshot lock poisoned".to_string())?.remove(&build.snapshot_id()); + } + } + self.requests.lock().map_err(|_| "remote target request lock poisoned".to_string())?.remove(&request_id); + result.map_err(|error| format!("remote dispatch failed: {error:?}")) + } + + async fn forward_event( + &self, request_id: crate::remote::RequestId, label: Option<&str>, request: &RemoteRequest, event: RemoteBackendEvent, writer: &mut W, + ) -> Result<(), String> { + if let (Some(label), crate::remote::RemoteOperation::Sync(_), RemoteBackendEvent::SyncCompleted { snapshot_id }) = + (label, request.operation(), &event) + { + self.snapshots.lock().map_err(|_| "remote target snapshot lock poisoned".to_string())?.insert(*snapshot_id, label.to_string()); + } + let response = + crate::vscomm::RemoteEvent::from_backend_event(request_id, event).to_frame().map_err(|error| format!("encode remote event: {error}"))?; + write_frame(writer, &response).await + } + + fn cancel_all(&self) { + for broker in self.brokers.values() { + broker.cancel_all(); + } + } + + fn cleanup_timeout(&self) -> std::time::Duration { + self.brokers.values().map(|broker| broker.cleanup_timeout()).max().unwrap_or_else(|| RemoteResourcePolicy::default().cleanup_timeout) + } +} + pub struct RemoteDaemonConfig { session: Arc, allowed_tools: Vec, @@ -533,11 +636,14 @@ impl RemoteDaemonConfig { } struct RemoteComponents { - context: RemoteExecutionContext, - policy: RemoteAuthorizationPolicy, - backend: Arc, + context: Option, + policy: Option, + backend: Option>, admission_limits: RemoteAdmissionLimits, cleanup_timeout: std::time::Duration, + router: Option>, + local_session_id: crate::remote::WorkspaceSessionId, + legacy_remote: bool, } impl VsockDaemon { @@ -565,6 +671,7 @@ impl VsockDaemon { } .with_snapshot_authority(session.clone()); let remote_context = RemoteExecutionContext { target: session.target(), workspace_session_id: session.session_id() }; + let local_session_id = session.session_id(); let backend: Arc = match backend { RemoteBackendSelection::Loopback { tools, target_environment } => Arc::new( LoopbackBackend::new(session, tools) @@ -577,19 +684,74 @@ impl VsockDaemon { } }; let remote_components = RemoteComponents { - context: remote_context, - policy: remote_policy, - backend, + context: Some(remote_context), + policy: Some(remote_policy), + backend: Some(backend), admission_limits, cleanup_timeout: resources.cleanup_timeout, + router: None, + local_session_id, + legacy_remote: true, }; Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, remote_components) } + #[allow(clippy::too_many_arguments)] + pub fn start_with_target_catalog( + passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, + catalog: BuildTargetCatalog, active_target: ActiveBuildTarget, sessions: BTreeMap>, + local_session_id: crate::remote::WorkspaceSessionId, global_active: usize, + ) -> Result { + let mut brokers = BTreeMap::new(); + let mut cleanup_timeout = RemoteResourcePolicy::default().cleanup_timeout; + for target in catalog.remote_targets() { + let label = target.summary().label().to_string(); + let session = sessions.get(&label).cloned().ok_or_else(|| format!("missing remote session for target {label}"))?; + let environment = RemoteEnvironmentPolicy::from_names(target.project().environment.clone())?; + let policy = RemoteAuthorizationPolicy::from_policies(session.target(), session.session_id(), target.tool_policies(), environment)? + .with_snapshot_authority(session.clone()); + let resources = target.target().resources(); + cleanup_timeout = cleanup_timeout.max(resources.cleanup_timeout()); + let backend: Arc = Arc::new( + SshBackend::new(session.clone(), target.target().clone())? + .with_artifacts(target.artifact_policy().clone(), resources.artifact_limits()), + ); + let target_active = resources.max_active_builds(); + let admission = RemoteAdmissionLimits::new(global_active, target_active)?; + let context = RemoteExecutionContext { target: session.target(), workspace_session_id: session.session_id() }; + let broker = Arc::new( + RemoteBroker::new(policy, context, backend).with_admission_limits(admission).with_cleanup_timeout(resources.cleanup_timeout()), + ); + brokers.insert(label, broker); + } + let router = Arc::new(RemoteRouter::new(active_target, brokers)); + let components = RemoteComponents { + context: None, + policy: None, + backend: None, + admission_limits: RemoteAdmissionLimits::default(), + cleanup_timeout, + router: Some(router), + local_session_id, + legacy_remote: false, + }; + Self::start_inner(passthrough, env_mode, workspace, profiles, share_dir, allow, components) + } + fn start_inner( passthrough: Vec, env_mode: EnvMode, workspace: PathBuf, profiles: Vec, share_dir: PathBuf, allow: Vec, remote: RemoteComponents, ) -> Result { + let RemoteComponents { + context, + policy, + backend, + admission_limits, + cleanup_timeout, + router: router_override, + local_session_id, + legacy_remote, + } = remote; let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); let merged_profile = if profiles.is_empty() { @@ -628,18 +790,27 @@ impl VsockDaemon { } let connections = Arc::new(Mutex::new(Vec::new())); - let remote_broker = Arc::new( - RemoteBroker::new(remote.policy, remote.context, remote.backend) - .with_admission_limits(remote.admission_limits) - .with_cleanup_timeout(remote.cleanup_timeout), - ); + let remote_router = router_override.unwrap_or_else(|| { + let context = context.expect("single remote context is present"); + let policy = policy.expect("single remote policy is present"); + let backend = backend.expect("single remote backend is present"); + let remote_broker = + Arc::new(RemoteBroker::new(policy, context, backend).with_admission_limits(admission_limits).with_cleanup_timeout(cleanup_timeout)); + let label = if legacy_remote { "legacy-remote" } else { "localhost" }; + let mut brokers = BTreeMap::new(); + brokers.insert(label.to_string(), remote_broker); + let active = if legacy_remote { ActiveBuildTarget::with_label(label) } else { ActiveBuildTarget::new() }; + Arc::new(RemoteRouter::new(active, brokers)) + }); let session = Arc::new(VsockSession { passthrough: Arc::new(passthrough), env_mode, workspace, merged_profile, proxy_config: proxy_config.map(Arc::new), - remote_broker, + remote_router, + local_session_id, + local_capabilities: Mutex::new(BTreeMap::new()), }); let listener = tokio_vsock::VsockListener::bind(tokio_vsock::VsockAddr::new(libc::VMADDR_CID_ANY, TOOLCHAIN_PORT)).map_err(|error| { @@ -704,9 +875,9 @@ async fn daemon_loop( } } - session.remote_broker.cancel_all(); + session.remote_router.cancel_all(); let tasks = connections.lock().map_err(|_| "connection task lock poisoned".to_string())?.drain(..).collect::>(); - let cleanup_deadline = session.remote_broker.cleanup_timeout(); + let cleanup_deadline = session.remote_router.cleanup_timeout(); let mut tasks = tasks; if tokio::time::timeout(cleanup_deadline, async { for task in &mut tasks { @@ -730,7 +901,7 @@ async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSess let frame = Frame::read_async(&mut reader).await.map_err(|e| format!("read frame: {e}"))?; if matches!(frame.frame_type, FrameType::RemoteRequest) { - return dispatch_remote_frame(frame, &session.remote_broker, &mut writer).await; + return dispatch_remote_frame_for_session(frame, session, &mut writer).await; } if !matches!(frame.frame_type, FrameType::ExecReq) { return Err(format!("expected ExecReq or RemoteRequest, got {:?}", frame.frame_type as u16)); @@ -759,6 +930,134 @@ async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSess Ok(()) } +async fn dispatch_remote_frame_for_session(frame: Frame, session: &VsockSession, writer: &mut W) -> Result<(), String> { + let request = crate::vscomm::RemoteRequest::from_frame(frame).map_err(|err| format!("decode remote request: {err}"))?.into_domain()?; + let local_snapshot = match request.operation() { + crate::remote::RemoteOperation::Build(build) => { + session.local_capabilities.lock().map_err(|_| "local target capability lock poisoned".to_string())?.contains_key(&build.snapshot_id()) + } + _ => false, + }; + match request.operation() { + crate::remote::RemoteOperation::Sync(_) if session.remote_router.active_target.current() == "localhost" => { + dispatch_local_remote_request(request, session, writer).await + } + crate::remote::RemoteOperation::Build(_) if local_snapshot => dispatch_local_remote_request(request, session, writer).await, + _ => session.remote_router.dispatch(request, writer).await, + } +} + +async fn dispatch_local_remote_request( + request: RemoteRequest, session: &VsockSession, writer: &mut W, +) -> Result<(), String> { + let request_id = request.request_id(); + if request.workspace_session_id() != session.local_session_id { + return write_local_remote_event( + writer, + request_id, + RemoteBackendEvent::Error { message: "local remote session does not match this sandbox".to_string() }, + ) + .await; + } + match request.operation() { + crate::remote::RemoteOperation::Sync(sync) => { + if sync.retain_capability() { + let snapshot_id = loop { + let mut bytes = [0u8; 16]; + rand::thread_rng().fill_bytes(&mut bytes); + let candidate = crate::remote::RemoteSnapshotId::from_bytes(bytes); + if !candidate.is_zero() + && !session + .local_capabilities + .lock() + .map_err(|_| "local target capability lock poisoned".to_string())? + .contains_key(&candidate) + { + break candidate; + } + }; + session.local_capabilities.lock().map_err(|_| "local target capability lock poisoned".to_string())?.insert(snapshot_id, ()); + write_local_remote_event(writer, request_id, RemoteBackendEvent::SyncCompleted { snapshot_id }).await + } else { + write_local_remote_event(writer, request_id, RemoteBackendEvent::Completed { exit_code: 0 }).await + } + } + crate::remote::RemoteOperation::Build(build) => { + let snapshot_id = build.snapshot_id(); + let available = + session.local_capabilities.lock().map_err(|_| "local target capability lock poisoned".to_string())?.remove(&snapshot_id).is_some(); + if !available { + return write_local_remote_event( + writer, + request_id, + RemoteBackendEvent::Error { message: "remote snapshot capability is unavailable".to_string() }, + ) + .await; + } + if !is_allowed(&session.passthrough, build.tool().as_str(), build.argv()) { + return write_local_remote_event( + writer, + request_id, + RemoteBackendEvent::Error { message: "local passthrough authorization rejected".to_string() }, + ) + .await; + } + let exec_request = ExecRequest { + cwd: build.cwd().as_str().to_string(), + command: build.tool().as_str().to_string(), + args: build.argv().to_vec(), + env: build.env().to_vec(), + }; + let (mut frame_reader, mut frame_writer) = tokio::io::duplex(64 * 1024); + let mut execute = Box::pin(execute_request(&mut frame_writer, session, &exec_request)); + loop { + tokio::select! { + result = &mut execute => { + if let Err(error) = result { + write_local_remote_event(writer, request_id, RemoteBackendEvent::Error { message: error }).await?; + } + break; + } + frame = Frame::read_async(&mut frame_reader) => { + let frame = frame.map_err(|error| format!("read local execution frame: {error}"))?; + match frame.frame_type { + FrameType::Stdout => write_local_remote_event(writer, request_id, RemoteBackendEvent::Stdout(frame.payload)).await?, + FrameType::Stderr => write_local_remote_event(writer, request_id, RemoteBackendEvent::Stderr(frame.payload)).await?, + FrameType::Exit => { + if frame.payload.len() != 4 { + return Err("local execution returned malformed exit frame".to_string()); + } + let exit_code = i32::from_le_bytes([frame.payload[0], frame.payload[1], frame.payload[2], frame.payload[3]]); + write_local_remote_event(writer, request_id, RemoteBackendEvent::Completed { exit_code }).await?; + break; + } + _ => return Err("local execution returned an unexpected frame".to_string()), + } + } + } + } + Ok(()) + } + crate::remote::RemoteOperation::Cancel { .. } => { + write_local_remote_event( + writer, + request_id, + RemoteBackendEvent::Error { message: "local target cancellation is unavailable".to_string() }, + ) + .await + } + } +} + +async fn write_local_remote_event( + writer: &mut W, request_id: crate::remote::RequestId, event: RemoteBackendEvent, +) -> Result<(), String> { + let response = crate::vscomm::RemoteEvent::from_backend_event(request_id, event) + .to_frame() + .map_err(|error| format!("encode local remote event: {error}"))?; + write_frame(writer, &response).await +} + pub async fn dispatch_remote_frame(frame: Frame, broker: &RemoteBroker, writer: &mut W) -> Result<(), String> { let request = crate::vscomm::RemoteRequest::from_frame(frame).map_err(|err| format!("decode remote request: {err}"))?.into_domain()?; let request_id = request.request_id(); diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index df7f6cc..127461c 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -1,4 +1,4 @@ -use super::{dispatch_remote_frame, is_allowed, RemoteBroker, RemoteDispatchError}; +use super::{dispatch_remote_frame, is_allowed, RemoteBroker, RemoteDispatchError, RemoteRouter}; use super::{monitor_bwrap_status, ChildEvent}; use crate::cfg::EnvMode; use crate::remote::{ @@ -6,10 +6,12 @@ use crate::remote::{ RemoteExecutionContext, RemoteExecutionControl, RemoteFuture, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, RequestId, WorkspaceRelativePath, WorkspaceSessionId, }; +use crate::remote_target::ActiveBuildTarget; use crate::vscomm::{ Frame, RemoteBuild as WireRemoteBuild, RemoteRequest as WireRemoteRequest, RemoteTool as WireRemoteTool, RequestId as WireRequestId, WorkspaceRelativePath as WireWorkspaceRelativePath, WorkspaceSessionId as WireWorkspaceSessionId, }; +use std::collections::BTreeMap; use std::io::Write; use std::path::Path; use std::sync::{Arc, Mutex}; @@ -448,7 +450,15 @@ fn local_exec_request_still_builds_on_the_local_path() { workspace: workspace.path().to_path_buf(), merged_profile: None, proxy_config: None, - remote_broker: Arc::new(remote_broker(Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }))), + remote_router: Arc::new(RemoteRouter::new( + ActiveBuildTarget::new(), + BTreeMap::from([( + "localhost".to_string(), + Arc::new(remote_broker(Arc::new(RecordingBackend { calls: Mutex::new(Vec::new()), emit: Vec::new(), result: None }))), + )]), + )), + local_session_id: WorkspaceSessionId([2; 16]), + local_capabilities: Mutex::new(BTreeMap::new()), }; let request = crate::vscomm::ExecRequest { cwd: "/workspace".into(), command: "true".into(), args: Vec::new(), env: Vec::new() }; diff --git a/src/main.rs b/src/main.rs index e058d7e..0bd189b 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,5 +1,6 @@ -use bunkerbox::cfg::{ProjectConfig, RemoteToolSpec, WorkspaceMode}; -use bunkerbox::remote::{RemoteAdmissionLimits, RemoteEnvironmentPolicy, RemoteToolPolicy}; +#[cfg(test)] +use bunkerbox::cfg::RemoteToolSpec; +use bunkerbox::cfg::{ProjectConfig, WorkspaceMode}; use bunkerbox::{cfg, cfgsetup, clidef, cmdrun, daemon, kata, logging, loopback, overlay, remote_target, snapshot, tui, vscomm, workspace}; use rand::RngCore; use std::ffi::OsString; @@ -14,6 +15,9 @@ const WORKSPACE_HANDOFF_MAGIC: &[u8; 4] = b"WS01"; const STARTUP_READY_MAGIC: &[u8; 4] = b"RDY1"; const MAX_WORKSPACE_HANDOFF_BYTES: usize = 64 * 1024; +type RemoteSessions = std::collections::BTreeMap>; +type SetupState = (workspace::WorkspaceHandle, RemoteSessions, daemon::VsockDaemon); + fn main() { if let Err(err) = run() { eprintln!("bunkerbox: {err}"); @@ -190,30 +194,23 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option (catalog, None), + Ok(None) => (remote_target::BuildTargetCatalog::localhost_only(repo_root.clone(), env.clone())?, None), + Err(error) => { + logging::diagnostic(&format!("remote.conf disabled for this run: {error}")); + (remote_target::BuildTargetCatalog::localhost_only(repo_root.clone(), env.clone())?, Some(error)) } - } + }; + let active_target = remote_target::ActiveBuildTarget::new(); let merged_allow: Vec = config.allow.clone().unwrap_or_default().into_iter().chain(env.image.allow.clone().unwrap_or_default()).collect(); let passthrough = env.project.passthrough.clone(); let env_mode = env.project.env; let profiles = env.profiles.clone(); - let remote_environment = RemoteEnvironmentPolicy::from_names(env.project.remote.environment.clone())?; - let remote_environment_names = remote_environment.allowed_names().map(str::to_string).collect::>(); - let remote_tool_policies = - env.project.remote.tools.iter().map(|tool| (tool.name.clone(), RemoteToolPolicy::new(tool.allow_args))).collect::>(); - let configured_remote_tool_names = env.project.remote.tools.iter().map(|tool| tool.name.clone()).collect::>(); - let remote_tool_names = remote_tool_names(&env.project.remote.tools); + let remote_environment_names = target_catalog.environment_names(); + let remote_tool_names = target_catalog.wrapper_names(); let share_dir_owned = share_dir.to_path_buf(); let mut sock_fds = [-1i32, -1]; @@ -230,7 +227,7 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option> = Arc::new(Mutex::new(tui::OverlayState::new())); + if let Some(error) = &remote_config_error { + tui::show_host_error(&overlay, "remote.conf", error); + } let status_listener = match start_status_listener(overlay.clone()) { Ok(listener) => listener, Err(error) => { @@ -320,75 +320,76 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option Result<(workspace::WorkspaceHandle, Arc, daemon::VsockDaemon), String> { - let mut setup_parent = unsafe { File::from_raw_fd(setup_parent_fd) }; - let startup_status = unsafe { File::from_raw_fd(startup_status_write) }; - logging::set_status_fd(startup_status.as_raw_fd()); - let _runtime_guard = setup_handle.enter(); - - if let Err(error) = read_startup_ready(&mut setup_parent) { - logging::log(&format!("Startup failed: {error}")); - return Err(error); - } + let target_catalog_for_setup = target_catalog.clone(); + let active_target_for_setup = active_target.clone(); + let setup_thread = std::thread::spawn(move || -> Result { + let mut setup_parent = unsafe { File::from_raw_fd(setup_parent_fd) }; + let startup_status = unsafe { File::from_raw_fd(startup_status_write) }; + logging::set_status_fd(startup_status.as_raw_fd()); + let _runtime_guard = setup_handle.enter(); + + if let Err(error) = read_startup_ready(&mut setup_parent) { + logging::log(&format!("Startup failed: {error}")); + return Err(error); + } - let setup_result = (|| -> Result<(workspace::WorkspaceHandle, Arc, daemon::VsockDaemon), String> { - logging::log("Preparing workspace..."); - let workspace = workspace::resolve(workspace_mode, quota, exclude.as_deref(), &name)?; - let remote_session = new_session_id(); - let target = new_target_id(); - logging::log("Preparing remote session..."); - loopback::cleanup_stale_roots(&std::env::temp_dir())?; - let snapshot_root = std::env::temp_dir().join(format!("bunkerbox-snapshots-{}-{}", std::process::id(), remote_session.to_hex())); - let jobs_root = std::env::temp_dir().join(format!("bunkerbox-loopback-{}-{}", std::process::id(), remote_session.to_hex())); + let setup_result = (|| -> Result { + logging::log("Preparing workspace..."); + let workspace = workspace::resolve(workspace_mode, quota, exclude.as_deref(), &name)?; + let remote_session = new_session_id(); + logging::log("Preparing remote session..."); + loopback::cleanup_stale_roots(&std::env::temp_dir())?; + let mut sessions = RemoteSessions::new(); + for remote_target in target_catalog_for_setup.remote_targets() { + let target_label = remote_target.summary().label(); + let snapshot_root = + std::env::temp_dir().join(format!("bunkerbox-snapshots-{}-{}-{target_label}", std::process::id(), remote_session.to_hex())); + let jobs_root = + std::env::temp_dir().join(format!("bunkerbox-remote-{}-{}-{target_label}", std::process::id(), remote_session.to_hex())); let snapshot_store = snapshot::SnapshotStore::new(&snapshot_root); - let exclusion_policy = snapshot::SnapshotExclusionPolicy::from_remote_config(&env, exclude.as_deref())?; + let mut target_project = target_catalog_for_setup.base_project().clone(); + target_project.project.remote = remote_target.project().clone(); + let exclusion_policy = snapshot::SnapshotExclusionPolicy::from_remote_config(&target_project, exclude.as_deref())?; let snapshot_builder = snapshot::SnapshotBuilder::new(snapshot_store.clone(), snapshot::SnapshotLimits::default(), exclusion_policy); let session = Arc::new(loopback::RunRemoteSession::new( bunkerbox::remote::WorkspaceSessionId(remote_session.0), - target, + new_target_id(), workspace.path().to_path_buf(), snapshot_store, snapshot_builder, jobs_root, )?); - let tools = loopback::resolve_fixed_tools(configured_remote_tool_names.clone()); - logging::log("Starting remote daemon..."); - let global_active = config.remote_max_active_builds()?; - let target_active = remote_backend.target().map(|target| target.resources().max_active_builds()).unwrap_or(1); - let remote_config = match remote_backend.backend() { - remote_target::BackendMode::Loopback => daemon::RemoteDaemonConfig::loopback(session.clone(), Vec::new(), tools), - remote_target::BackendMode::Ssh => { - let target = remote_backend.target().cloned().ok_or_else(|| "SSH backend selection has no target".to_string())?; - daemon::RemoteDaemonConfig::ssh(session.clone(), target)? - } - }; - let daemon = daemon::VsockDaemon::start_with_remote( - passthrough, - env_mode, - workspace.path().to_path_buf(), - profiles, - share_dir_owned, - merged_allow, - remote_config - .with_artifacts(artifact_policy, artifact_limits) - .with_policy(remote_tool_policies, remote_environment) - .with_admission_limits(RemoteAdmissionLimits::new(global_active, target_active)?), - )?; - if let Err(error) = write_run_handoff(&mut setup_parent, workspace.path(), remote_session) { - tokio::runtime::Handle::current().block_on(daemon.shutdown()); - return Err(error); - } - Ok((workspace, session, daemon)) - })(); - - if let Err(error) = &setup_result { - logging::log(&format!("Startup failed: {error}")); + sessions.insert(target_label.to_string(), session); } - setup_result - }); + logging::log("Starting remote daemon..."); + let global_active = config.remote_max_active_builds()?; + let daemon = daemon::VsockDaemon::start_with_target_catalog( + passthrough, + env_mode, + workspace.path().to_path_buf(), + profiles, + share_dir_owned, + merged_allow, + target_catalog_for_setup, + active_target_for_setup, + sessions.clone(), + bunkerbox::remote::WorkspaceSessionId(remote_session.0), + global_active, + )?; + if let Err(error) = write_run_handoff(&mut setup_parent, workspace.path(), remote_session) { + tokio::runtime::Handle::current().block_on(daemon.shutdown()); + return Err(error); + } + Ok((workspace, sessions, daemon)) + })(); - let tui_result = tui::event_loop(master, rows, cols, parent_fd, startup_status_read, overlay); + if let Err(error) = &setup_result { + logging::log(&format!("Startup failed: {error}")); + } + setup_result + }); + + let tui_result = tui::event_loop(master, rows, cols, parent_fd, startup_status_read, overlay, target_catalog, active_target); let tui_error = tui_result.err(); if tui_error.is_some() { unsafe { libc::kill(pid, libc::SIGTERM) }; @@ -405,9 +406,9 @@ fn run_packaged_runtime(config: cfg::RuntimeConfig, workspace_override: Option { + Ok((workspace, sessions, daemon)) => { tokio::runtime::Handle::current().block_on(daemon.shutdown()); - drop(remote_session); + drop(sessions); drop(workspace); (true, None) } @@ -450,6 +451,7 @@ fn new_target_id() -> bunkerbox::remote::RemoteTargetId { } } +#[cfg(test)] fn remote_tool_names(entries: &[RemoteToolSpec]) -> Vec { entries.iter().map(|tool| tool.name.clone()).collect() } diff --git a/src/main_ut.rs b/src/main_ut.rs index c2b3624..5d3136b 100644 --- a/src/main_ut.rs +++ b/src/main_ut.rs @@ -62,9 +62,9 @@ fn run_handoff_rejects_zero_session() { fn remote_tool_names_propagates_the_complete_configured_set() { assert_eq!( remote_tool_names(&[ - RemoteToolSpec { name: "make".into(), allow_args: true }, - RemoteToolSpec { name: "cargo".into(), allow_args: false }, - RemoteToolSpec { name: "cmake".into(), allow_args: true }, + RemoteToolSpec { name: "make".into(), command: None, allow_args: true }, + RemoteToolSpec { name: "cargo".into(), command: None, allow_args: false }, + RemoteToolSpec { name: "cmake".into(), command: None, allow_args: true }, ]), vec!["make", "cargo", "cmake"] ); diff --git a/src/remote.rs b/src/remote.rs index 31aed9e..79891d5 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -232,19 +232,29 @@ impl RemoteEnvironmentPolicy { } } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq)] pub struct RemoteToolPolicy { allow_arbitrary_argv: bool, + command: Option, } impl RemoteToolPolicy { pub fn new(allow_arbitrary_argv: bool) -> Self { - Self { allow_arbitrary_argv } + Self { allow_arbitrary_argv, command: None } + } + + pub fn with_command(mut self, command: impl Into) -> Self { + self.command = Some(command.into()); + self } - pub fn allows_arbitrary_argv(self) -> bool { + pub fn allows_arbitrary_argv(&self) -> bool { self.allow_arbitrary_argv } + + pub fn command(&self) -> Option<&str> { + self.command.as_deref() + } } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -473,6 +483,7 @@ pub struct RemoteExecutionContext { pub struct AuthorizedRemoteRequest { request: RemoteRequest, target: RemoteTargetId, + target_command: Option, } impl AuthorizedRemoteRequest { @@ -487,6 +498,10 @@ impl AuthorizedRemoteRequest { pub fn request(&self) -> &RemoteRequest { &self.request } + + pub fn target_command(&self) -> Option<&str> { + self.target_command.as_deref() + } } #[derive(Debug, Clone, PartialEq, Eq)] @@ -538,6 +553,9 @@ impl RemoteAuthorizationPolicy { let mut allowed_tools = BTreeMap::new(); for (tool, policy) in tools { RemoteTool::new(tool.clone())?; + if let Some(command) = policy.command() { + RemoteTool::new(command.to_string())?; + } if allowed_tools.insert(tool.clone(), policy).is_some() { return Err(format!("duplicate remote tool policy: {tool}")); } @@ -565,6 +583,7 @@ impl RemoteAuthorizationPolicy { return Err(RemoteAuthorizationError::TargetNotAllowed); } + let mut target_command = None; let request = match request.operation() { RemoteOperation::Sync(_) => request, RemoteOperation::Cancel { target_request_id } => { @@ -577,7 +596,7 @@ impl RemoteAuthorizationPolicy { if self.snapshot_authority.as_ref().is_none_or(|authority| !authority.snapshot_available(self.allowed_session, build.snapshot_id())) { return Err(RemoteAuthorizationError::SnapshotNotAllowed); } - let Some(tool_policy) = self.allowed_tools.get(build.tool().as_str()).copied() else { + let Some(tool_policy) = self.allowed_tools.get(build.tool().as_str()).cloned() else { return Err(RemoteAuthorizationError::ToolNotAllowed(build.tool().as_str().to_string())); }; if !tool_policy.allows_arbitrary_argv() && !build.argv().is_empty() { @@ -591,13 +610,14 @@ impl RemoteAuthorizationPolicy { } else { self.environment.filter(build.env())? }; + target_command = tool_policy.command().map(str::to_string); let filtered = RemoteBuild::new(build.cwd.clone(), build.tool.clone(), build.argv.clone(), environment, build.snapshot_id()) .map_err(RemoteAuthorizationError::InvalidEnvironment)?; RemoteRequest::build(request.request_id, request.workspace_session_id, filtered) } }; - Ok(AuthorizedRemoteRequest { request, target: context.target }) + Ok(AuthorizedRemoteRequest { request, target: context.target, target_command }) } } diff --git a/src/remote_target.rs b/src/remote_target.rs index d39b54a..576435e 100644 --- a/src/remote_target.rs +++ b/src/remote_target.rs @@ -1,4 +1,6 @@ use crate::artifact::{ArtifactLimits, ArtifactPolicy}; +use crate::cfg::{ProjectConfig, RemoteSection}; +use crate::remote::{RemoteEnvironmentPolicy, RemoteToolPolicy}; use serde::de::{self, MapAccess, Visitor}; use serde::Deserialize; use std::collections::BTreeMap; @@ -9,6 +11,7 @@ use std::marker::PhantomData; use std::net::IpAddr; use std::path::{Component, Path, PathBuf}; use std::str::FromStr; +use std::sync::{Arc, Mutex}; use std::time::Duration; #[cfg(unix)] @@ -17,6 +20,8 @@ use std::os::unix::fs::PermissionsExt; pub const CONFIG_VERSION: u64 = 1; pub const CONFIG_FILE_NAME: &str = "remote-targets.yaml"; pub const CONFIG_DIRECTORY_NAME: &str = "bunkerbox"; +pub const REMOTE_PROJECT_CONFIG_FILE_NAME: &str = "remote.conf"; +pub const FIXED_WORKER_PATH: &str = "/usr/local/libexec/bunkerbox-worker"; const MAX_CONFIG_PATH_BYTES: usize = 4096; const MAX_REMOTE_PATH_BYTES: usize = 4096; @@ -163,6 +168,8 @@ pub struct SshTarget { tools: BTreeMap, environment: BTreeMap, resources: ResourceLimits, + compact: bool, + port_explicit: bool, } /// Alias emphasizing that an `SshTarget` can only be obtained after validation. @@ -231,6 +238,34 @@ impl SshTarget { pub fn resources(&self) -> ResourceLimits { self.resources } + + pub fn compact_destination(&self) -> bool { + self.compact + } + + pub fn port_explicit(&self) -> bool { + self.port_explicit + } + + pub fn from_compact(name: String, destination: &str, workspace: String, resources: ResourceLimits) -> Result { + let (user, host, port, port_explicit) = parse_ssh_destination(destination)?; + validate_remote_path("workspace", &workspace, true)?; + Ok(Self { + name, + host, + port, + user, + identity_file: PathBuf::new(), + known_hosts_file: PathBuf::new(), + worker_path: FIXED_WORKER_PATH.to_string(), + workspace_root: workspace, + tools: BTreeMap::new(), + environment: BTreeMap::new(), + resources, + compact: true, + port_explicit, + }) + } } #[derive(Clone, Debug, PartialEq, Eq)] @@ -482,6 +517,7 @@ pub fn default_config_path_with(xdg_config_home: Option<&Path>, home: Option<&Pa ConfigPathHelper::from_paths(xdg_config_home, home).config_path() } +#[derive(Debug)] struct UniqueMap(BTreeMap); impl Default for UniqueMap { @@ -634,6 +670,7 @@ struct RawResources { #[derive(Deserialize)] #[serde(untagged)] +#[derive(Debug)] enum RawQuantity { Integer(u64), Text(String), @@ -691,6 +728,8 @@ fn validate_target(name: String, raw: RawTarget) -> Result { tools, environment, resources, + compact: false, + port_explicit: true, }) } @@ -1011,6 +1050,452 @@ fn nonempty_path(path: Option) -> Option { path.filter(|path| !path.as_os_str().is_empty()) } +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct BuildTargetSummary { + label: String, + workspace: String, + local: bool, +} + +impl BuildTargetSummary { + pub fn label(&self) -> &str { + &self.label + } + + pub fn workspace(&self) -> &str { + &self.workspace + } + + pub fn is_local(&self) -> bool { + self.local + } +} + +#[derive(Clone, Debug)] +pub struct RemoteBuildTarget { + summary: BuildTargetSummary, + target: SshTarget, + project: RemoteSection, + artifact_policy: ArtifactPolicy, +} + +impl RemoteBuildTarget { + pub fn summary(&self) -> &BuildTargetSummary { + &self.summary + } + + pub fn target(&self) -> &SshTarget { + &self.target + } + + pub fn project(&self) -> &RemoteSection { + &self.project + } + + pub fn artifact_policy(&self) -> &ArtifactPolicy { + &self.artifact_policy + } + + pub fn tool_policies(&self) -> Vec<(String, RemoteToolPolicy)> { + self.project + .tools + .iter() + .map(|tool| { + let command = tool.command.clone().unwrap_or_else(|| tool.name.clone()); + (tool.name.clone(), RemoteToolPolicy::new(tool.allow_args).with_command(command)) + }) + .collect() + } +} + +#[derive(Clone, Debug)] +pub struct BuildTargetCatalog { + project_root: PathBuf, + base_project: ProjectConfig, + summaries: Vec, + remotes: BTreeMap, +} + +impl BuildTargetCatalog { + pub fn localhost_only(project_root: PathBuf, base_project: ProjectConfig) -> Result { + validate_remote_section(&base_project.project.remote)?; + Ok(Self { + summaries: vec![BuildTargetSummary { label: "localhost".to_string(), workspace: project_root.display().to_string(), local: true }], + project_root, + base_project, + remotes: BTreeMap::new(), + }) + } + + pub fn load_optional(project_root: &Path, base_project: &ProjectConfig) -> Result, String> { + let path = project_root.join(".bunkerbox").join(REMOTE_PROJECT_CONFIG_FILE_NAME); + if !path.exists() { + return Ok(None); + } + let contents = fs::read_to_string(&path).map_err(|error| format!("failed to read {}: {error}", path.display()))?; + let raw: RawRemoteProjectConfig = serde_yaml::from_str(&contents).map_err(|error| format!("failed to parse {}: {error}", path.display()))?; + Self::from_raw(project_root.to_path_buf(), base_project.clone(), raw).map(Some) + } + + fn from_raw(project_root: PathBuf, base_project: ProjectConfig, raw: RawRemoteProjectConfig) -> Result { + validate_remote_section(&base_project.project.remote)?; + let mut remotes = BTreeMap::new(); + let mut summaries = vec![BuildTargetSummary { label: "localhost".to_string(), workspace: project_root.display().to_string(), local: true }]; + for (label, target) in raw.targets.0 { + validate_name("target label", &label)?; + if label == "localhost" { + return Err("remote target label 'localhost' is reserved".to_string()); + } + let project = merge_remote_overlay(&base_project.project.remote, target.project.as_ref())?; + validate_remote_section(&project)?; + let resources = validate_compact_resources(target.resources)?; + let ssh_target = SshTarget::from_compact(label.clone(), &target.ssh, target.workspace, resources)?; + let artifact_policy = ArtifactPolicy::new(project.artifacts.clone())?; + artifact_policy.validate_limits(resources.artifact_limits())?; + let summary = BuildTargetSummary { label: label.clone(), workspace: ssh_target.workspace_root().to_string(), local: false }; + let remote = RemoteBuildTarget { summary: summary.clone(), target: ssh_target, project, artifact_policy }; + if remotes.insert(label.clone(), remote).is_some() { + return Err(format!("duplicate target label: {label}")); + } + summaries.push(summary); + } + Ok(Self { project_root, base_project, summaries, remotes }) + } + + pub fn project_root(&self) -> &Path { + &self.project_root + } + + pub fn base_project(&self) -> &ProjectConfig { + &self.base_project + } + + pub fn summaries(&self) -> &[BuildTargetSummary] { + &self.summaries + } + + pub fn remote(&self, label: &str) -> Option<&RemoteBuildTarget> { + self.remotes.get(label) + } + + pub fn remote_targets(&self) -> impl Iterator { + self.remotes.values() + } + + pub fn wrapper_names(&self) -> Vec { + let mut names = BTreeMap::new(); + for tool in &self.base_project.project.remote.tools { + names.insert(tool.name.clone(), ()); + } + for target in self.remotes.values() { + for tool in &target.project.tools { + names.insert(tool.name.clone(), ()); + } + } + names.into_keys().collect() + } + + pub fn environment_names(&self) -> Vec { + let mut names = BTreeMap::new(); + for name in &self.base_project.project.remote.environment { + names.insert(name.clone(), ()); + } + for target in self.remotes.values() { + for name in &target.project.environment { + names.insert(name.clone(), ()); + } + } + names.into_keys().collect() + } +} + +#[derive(Clone)] +pub struct ActiveBuildTarget { + selected: Arc>, +} + +impl ActiveBuildTarget { + pub fn new() -> Self { + Self::with_label("localhost") + } + + pub(crate) fn with_label(label: impl Into) -> Self { + Self { selected: Arc::new(Mutex::new(label.into())) } + } + + pub fn current(&self) -> String { + self.selected.lock().map(|value| value.clone()).unwrap_or_else(|_| "localhost".to_string()) + } + + pub fn select(&self, catalog: &BuildTargetCatalog, label: &str) -> Result<(), String> { + if !catalog.summaries.iter().any(|summary| summary.label == label) { + return Err(format!("unknown build target: {label}")); + } + let mut selected = self.selected.lock().map_err(|_| "active build target lock poisoned".to_string())?; + *selected = label.to_string(); + Ok(()) + } +} + +impl Default for ActiveBuildTarget { + fn default() -> Self { + Self::new() + } +} + +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct RawRemoteProjectConfig { + #[serde(default)] + targets: UniqueMap, +} + +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct RawCompactTarget { + ssh: String, + workspace: String, + #[serde(default)] + project: Option, + #[serde(default)] + resources: Option, +} + +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct RawTargetProject { + #[serde(default)] + remote: Option, +} + +#[derive(Debug, Deserialize)] +#[serde(deny_unknown_fields)] +struct RawRemoteOverlay { + #[serde(default)] + exclude: Option>, + #[serde(default)] + environment: Option>, + #[serde(default)] + tools: Option>, + #[serde(default)] + artifacts: Option>, +} + +#[derive(Debug, Default, Deserialize)] +#[serde(deny_unknown_fields)] +struct RawTargetResources { + #[serde( + default, + rename = "connect-timeout-seconds", + alias = "connect-timeout", + alias = "connect_timeout_seconds", + alias = "connect_timeout", + alias = "connect" + )] + connect_timeout: Option, + #[serde( + default, + rename = "sync-timeout-seconds", + alias = "sync-timeout", + alias = "sync_timeout_seconds", + alias = "sync_timeout", + alias = "sync" + )] + sync_timeout: Option, + #[serde( + default, + rename = "build-timeout-seconds", + alias = "build-timeout", + alias = "build_timeout_seconds", + alias = "build_timeout", + alias = "build" + )] + build_timeout: Option, + #[serde(default, rename = "max-output-bytes", alias = "max-output", alias = "max_output_bytes", alias = "max_output")] + max_output: Option, + #[serde(default, rename = "idle-output-timeout-seconds", alias = "idle-output-timeout", alias = "idle_output_timeout")] + idle_output_timeout: Option, + #[serde(default, rename = "cleanup-timeout-seconds", alias = "cleanup-timeout", alias = "cleanup_timeout")] + cleanup_timeout: Option, + #[serde(default, rename = "artifact-timeout-seconds", alias = "artifact-timeout", alias = "artifact_timeout")] + artifact_timeout: Option, + #[serde(default, rename = "max-artifact-bytes", alias = "max-artifact-bytes-per-file", alias = "max_artifact_bytes")] + max_artifact_bytes: Option, + #[serde(default, rename = "max-artifact-total-bytes", alias = "max_artifact_total_bytes")] + max_artifact_total_bytes: Option, + #[serde(default, rename = "max-artifact-entries", alias = "max_artifact_entries")] + max_artifact_entries: Option, + #[serde(default, rename = "max-worker-uploads", alias = "max_worker_uploads")] + max_worker_uploads: Option, + #[serde(default, rename = "max-worker-upload-bytes", alias = "max_worker_upload_bytes")] + max_worker_upload_bytes: Option, + #[serde(default, rename = "max-worker-jobs", alias = "max_worker_jobs")] + max_worker_jobs: Option, + #[serde(default, rename = "max-worker-job-bytes", alias = "max_worker_job_bytes")] + max_worker_job_bytes: Option, + #[serde(default, rename = "max-worker-artifact-spools", alias = "max_worker_artifact_spools")] + max_worker_artifact_spools: Option, + #[serde(default, rename = "max-worker-artifact-spool-bytes", alias = "max_worker_artifact_spool_bytes")] + max_worker_artifact_spool_bytes: Option, + #[serde(default, rename = "max-worker-state-entries", alias = "max_worker_state_entries")] + max_worker_state_entries: Option, + #[serde(default, rename = "max-active-builds", alias = "max_active_builds")] + max_active_builds: Option, +} + +fn merge_remote_overlay(base: &RemoteSection, target: Option<&RawTargetProject>) -> Result { + let Some(target) = target.and_then(|project| project.remote.as_ref()) else { return Ok(base.clone()) }; + let mut merged = base.clone(); + if let Some(exclude) = &target.exclude { + merged.exclude = exclude.clone(); + } + if let Some(environment) = &target.environment { + merged.environment = environment.clone(); + } + if let Some(tools) = &target.tools { + merged.tools = tools.clone(); + } + if let Some(artifacts) = &target.artifacts { + merged.artifacts = artifacts.clone(); + } + Ok(merged) +} + +fn validate_remote_section(section: &RemoteSection) -> Result<(), String> { + RemoteEnvironmentPolicy::from_names(section.environment.clone())?; + crate::snapshot::SnapshotExclusionPolicy::from_patterns(section.exclude.clone())?; + ArtifactPolicy::new(section.artifacts.clone())?; + let mut names = BTreeMap::new(); + for tool in §ion.tools { + crate::remote::validate_remote_wrapper_name(tool.name.clone())?; + if let Some(command) = &tool.command { + crate::remote::validate_remote_wrapper_name(command.clone())?; + } + if names.insert(tool.name.clone(), ()).is_some() { + return Err(format!("duplicate remote tool: {}", tool.name)); + } + } + Ok(()) +} + +fn validate_compact_resources(raw: Option) -> Result { + let defaults = crate::remote::RemoteResourcePolicy::default(); + let artifact_defaults = ArtifactLimits::default(); + let worker_defaults = WorkerStateLimits::default(); + let raw = raw.unwrap_or_default(); + let connect_timeout = raw.connect_timeout.map_or(Ok(defaults.sync_timeout), |value| parse_duration("connect-timeout", value))?; + let sync_timeout = raw.sync_timeout.map_or(Ok(defaults.sync_timeout), |value| parse_duration("sync-timeout", value))?; + let build_timeout = raw.build_timeout.map_or(Ok(defaults.build_timeout), |value| parse_duration("build-timeout", value))?; + let idle_output_timeout = + raw.idle_output_timeout.map_or(Ok(defaults.idle_output_timeout), |value| parse_duration("idle-output-timeout", value))?; + let cleanup_timeout = raw.cleanup_timeout.map_or(Ok(defaults.cleanup_timeout), |value| parse_duration("cleanup-timeout", value))?; + validate_lifecycle_duration("connect-timeout", connect_timeout)?; + validate_lifecycle_duration("sync-timeout", sync_timeout)?; + validate_lifecycle_duration("build-timeout", build_timeout)?; + validate_lifecycle_duration("idle-output-timeout", idle_output_timeout)?; + validate_lifecycle_duration("cleanup-timeout", cleanup_timeout)?; + let max_output_bytes = raw.max_output.map_or(Ok(defaults.max_output_bytes), |value| parse_size("max-output", value))?; + let max_active_builds = parse_count("max-active-builds", raw.max_active_builds, 1)?; + if max_active_builds == 0 || max_active_builds > 64 { + return Err("max-active-builds must be between 1 and 64".to_string()); + } + let artifact_timeout = raw.artifact_timeout.map_or(Ok(artifact_defaults.timeout), |value| parse_duration("artifact-timeout", value))?; + let max_artifact_bytes = raw.max_artifact_bytes.map_or(Ok(artifact_defaults.max_file_bytes), |value| parse_size("max-artifact-bytes", value))?; + let max_artifact_total_bytes = + raw.max_artifact_total_bytes.map_or(Ok(artifact_defaults.max_total_bytes), |value| parse_size("max-artifact-total-bytes", value))?; + let max_artifact_entries = raw + .max_artifact_entries + .map_or(Ok(artifact_defaults.max_entries), |value| usize::try_from(value).map_err(|_| "max-artifact-entries is too large".to_string()))?; + let artifact = ArtifactLimits::new(artifact_timeout, max_artifact_entries, max_artifact_bytes, max_artifact_total_bytes)?; + let worker = WorkerStateLimits::new( + parse_count("max-worker-uploads", raw.max_worker_uploads, worker_defaults.max_uploads)?, + raw.max_worker_upload_bytes.map_or(Ok(worker_defaults.max_upload_bytes), |value| parse_size("max-worker-upload-bytes", value))?, + parse_count("max-worker-jobs", raw.max_worker_jobs, worker_defaults.max_jobs)?, + raw.max_worker_job_bytes.map_or(Ok(worker_defaults.max_job_bytes), |value| parse_size("max-worker-job-bytes", value))?, + parse_count("max-worker-artifact-spools", raw.max_worker_artifact_spools, worker_defaults.max_artifact_spools)?, + raw.max_worker_artifact_spool_bytes + .map_or(Ok(worker_defaults.max_artifact_spool_bytes), |value| parse_size("max-worker-artifact-spool-bytes", value))?, + parse_count("max-worker-state-entries", raw.max_worker_state_entries, worker_defaults.max_state_entries)?, + )?; + Ok(ResourceLimits { + connect_timeout, + sync_timeout, + build_timeout, + idle_output_timeout, + cleanup_timeout, + max_output_bytes, + max_active_builds, + artifact, + worker, + }) +} + +fn parse_ssh_destination(value: &str) -> Result<(String, String, u16, bool), String> { + if value.is_empty() + || value.len() > MAX_HOST_BYTES + || !value.is_ascii() + || value.chars().any(char::is_control) + || value.chars().any(char::is_whitespace) + { + return Err("SSH destination has invalid syntax".to_string()); + } + if value.starts_with('-') || value.contains('/') || value.matches('@').count() > 1 { + return Err("SSH destination has invalid syntax".to_string()); + } + let (user, authority) = value.split_once('@').map_or((String::new(), value), |(user, authority)| (user.to_string(), authority)); + if !user.is_empty() { + validate_username(&user)?; + } + let (host, port, port_explicit) = if let Some(rest) = authority.strip_prefix('[') { + let close = rest.find(']').ok_or_else(|| "SSH destination has invalid bracketed host".to_string())?; + let host = &rest[..close]; + let suffix = &rest[close + 1..]; + if suffix.is_empty() { + (host.to_string(), 22, false) + } else { + let port = suffix + .strip_prefix(':') + .ok_or_else(|| "SSH destination has invalid port".to_string())? + .parse::() + .map_err(|_| "SSH destination has invalid port".to_string())?; + if port == 0 { + return Err("SSH destination port must be positive".to_string()); + } + (host.to_string(), port, true) + } + } else if authority.matches(':').count() == 1 { + let (host, port) = authority.split_once(':').unwrap(); + let port = port.parse::().map_err(|_| "SSH destination has invalid port".to_string())?; + if port == 0 { + return Err("SSH destination port must be positive".to_string()); + } + (host.to_string(), port, true) + } else if authority.contains(':') { + return Err("SSH destination must use a bracketed IPv6 host".to_string()); + } else { + (authority.to_string(), 22, false) + }; + validate_destination_host(&host)?; + Ok((user, host, port, port_explicit)) +} + +fn validate_destination_host(host: &str) -> Result<(), String> { + if host.is_empty() || host.len() > MAX_HOST_BYTES || host.starts_with('-') || host.ends_with('-') || host.chars().any(char::is_control) { + return Err("SSH destination host has invalid syntax".to_string()); + } + if IpAddr::from_str(host).is_ok() { + return Ok(()); + } + let bytes = host.as_bytes(); + if !bytes[0].is_ascii_alphanumeric() + || !bytes[bytes.len() - 1].is_ascii_alphanumeric() + || !bytes.iter().all(|byte| byte.is_ascii_alphanumeric() || matches!(*byte, b'.' | b'-' | b'_' | b'+')) + { + return Err("SSH destination host has invalid syntax".to_string()); + } + Ok(()) +} + struct RedactedEnvironment(usize); impl fmt::Debug for RedactedEnvironment { diff --git a/src/remote_target_ut.rs b/src/remote_target_ut.rs index 5a08f33..30ee1c5 100644 --- a/src/remote_target_ut.rs +++ b/src/remote_target_ut.rs @@ -1,4 +1,5 @@ use super::*; +use crate::cfg::{ProjectConfig, RemoteToolSpec}; use std::fs; use std::path::{Path, PathBuf}; use std::time::Duration; @@ -85,6 +86,66 @@ fn write_config(fixture: &Fixture, yaml: &str) { fs::write(&fixture.config, yaml).unwrap(); } +#[test] +fn compact_project_catalog_is_localhost_first_and_freezes_target_overlay() { + let fixture = Fixture::new(); + let bunkerbox = fixture.project.join(".bunkerbox"); + fs::create_dir(&bunkerbox).unwrap(); + fs::write( + bunkerbox.join(REMOTE_PROJECT_CONFIG_FILE_NAME), + r#"targets: + netbsd: + ssh: builder@build.example.test:2222 + workspace: /var/tmp/bunkerbox + project: + remote: + environment: [CC] + tools: + - name: make + command: gmake + allow-args: true + artifacts: [build/output] +"#, + ) + .unwrap(); + + let mut base = ProjectConfig::default(); + base.project.remote.tools = vec![RemoteToolSpec { name: "cargo".into(), command: None, allow_args: false }]; + let catalog = BuildTargetCatalog::load_optional(&fixture.project, &base).unwrap().unwrap(); + assert_eq!(catalog.summaries().iter().map(BuildTargetSummary::label).collect::>(), vec!["localhost", "netbsd"]); + let target = catalog.remote("netbsd").unwrap(); + assert!(target.target().compact_destination()); + assert_eq!(target.target().worker_path(), FIXED_WORKER_PATH); + assert_eq!(target.project().tools.len(), 1); + assert_eq!(target.project().tools[0].name, "make"); + assert_eq!(target.project().tools[0].command.as_deref(), Some("gmake")); + assert_eq!(target.tool_policies()[0].1.command(), Some("gmake")); + assert_eq!(target.artifact_policy().paths(), &["build/output".to_string()]); + assert_eq!(catalog.wrapper_names(), vec!["cargo", "make"]); + assert_eq!(catalog.environment_names(), vec!["CC"]); +} + +#[test] +fn compact_target_overlay_rejects_non_remote_project_fields() { + let fixture = Fixture::new(); + let bunkerbox = fixture.project.join(".bunkerbox"); + fs::create_dir(&bunkerbox).unwrap(); + fs::write( + bunkerbox.join(REMOTE_PROJECT_CONFIG_FILE_NAME), + "targets:\n netbsd:\n ssh: build.example.test\n workspace: /var/tmp/bunkerbox\n project:\n image:\n session-mb: 1\n", + ) + .unwrap(); + assert!(BuildTargetCatalog::load_optional(&fixture.project, &ProjectConfig::default()).is_err()); +} + +#[test] +fn compact_ssh_destination_accepts_aliases_users_ports_and_ipv6() { + assert_eq!(parse_ssh_destination("builder@build.example.test:2200").unwrap(), ("builder".into(), "build.example.test".into(), 2200, true)); + assert_eq!(parse_ssh_destination("[2001:db8::1]").unwrap(), ("".into(), "2001:db8::1".into(), 22, false)); + assert!(parse_ssh_destination("-oProxyCommand=bad").is_err()); + assert!(parse_ssh_destination("build.example.test/path").is_err()); +} + #[test] fn valid_config_loads_and_validates_an_ssh_target() { let fixture = Fixture::new(); diff --git a/src/remote_ut.rs b/src/remote_ut.rs index 3e548d7..1c4961f 100644 --- a/src/remote_ut.rs +++ b/src/remote_ut.rs @@ -168,6 +168,22 @@ fn command_policy_rejects_unapproved_arguments() { assert_eq!(cargo_policy.authorize(&context(), request("cargo")), Err(RemoteAuthorizationError::ToolArgumentsNotAllowed("cargo".into()))); } +#[test] +fn command_policy_rewrites_only_the_trusted_target_command() { + let policy = RemoteAuthorizationPolicy::from_policies( + RemoteTargetId([3; 16]), + WorkspaceSessionId([2; 16]), + [("make".into(), RemoteToolPolicy::new(true).with_command("gmake"))], + RemoteEnvironmentPolicy::default(), + ) + .unwrap() + .with_snapshot_authority(std::sync::Arc::new(TestSnapshotAuthority)); + let authorized = policy.authorize(&context(), request("make")).unwrap(); + let RemoteOperation::Build(build) = authorized.request().operation() else { panic!("expected build") }; + assert_eq!(build.tool().as_str(), "make"); + assert_eq!(authorized.target_command(), Some("gmake")); +} + #[test] fn request_and_cancel_ids_must_be_nonzero() { let policy = policy(vec!["make".into()]); diff --git a/src/ssh.rs b/src/ssh.rs index 57dce03..3a12147 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -9,7 +9,7 @@ use crate::snapshot::SnapshotEntryKind; use crate::worker_protocol::{ self, WorkerArtifactEntry, WorkerArtifactPath, WorkerArtifactSetId, WorkerBuild, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRelativePath, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, WorkerUploadId, MAX_WORKER_CHUNK_BYTES, MAX_WORKER_FILE_BYTES, - WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, + WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_COMMAND_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, }; use rand::RngCore; use std::collections::BTreeMap; @@ -27,6 +27,16 @@ const SSH_PROGRAM: &str = "/usr/bin/ssh"; const MAX_SSH_DIAGNOSTIC_BYTES: usize = 16 * 1024; const PROCESS_REAP_TIMEOUT: Duration = Duration::from_secs(2); +fn build_protocol_version(target: &SshTarget, artifacts: bool) -> u16 { + if target.compact_destination() { + WORKER_COMMAND_PROTOCOL_VERSION + } else if artifacts { + WORKER_ARTIFACT_PROTOCOL_VERSION + } else { + WORKER_PROTOCOL_VERSION + } +} + type WorkerReader = Box; type WorkerWriter = Box; @@ -39,7 +49,7 @@ pub struct SshLaunchSpec { impl SshLaunchSpec { pub fn from_target(target: &SshTarget) -> Result { - if target.host().is_empty() || target.user().is_empty() || target.port() == 0 { + if target.host().is_empty() || target.port() == 0 || (!target.compact_destination() && target.user().is_empty()) { return Err("SSH target is not fully validated".to_string()); } @@ -59,22 +69,12 @@ impl SshLaunchSpec { worker.max_state_entries, ); let connect_timeout = target.resources().connect_timeout().as_secs().max(1).to_string(); - let args = vec![ - "-F".to_string(), - "/dev/null".to_string(), + let mut args = vec![ "-o".to_string(), "BatchMode=yes".to_string(), "-o".to_string(), "StrictHostKeyChecking=yes".to_string(), "-o".to_string(), - format!("UserKnownHostsFile={}", target.known_hosts_file().display()), - "-o".to_string(), - "GlobalKnownHostsFile=/dev/null".to_string(), - "-o".to_string(), - "IdentitiesOnly=yes".to_string(), - "-o".to_string(), - "IdentityAgent=none".to_string(), - "-o".to_string(), "ForwardAgent=no".to_string(), "-o".to_string(), "ClearAllForwardings=yes".to_string(), @@ -90,17 +90,32 @@ impl SshLaunchSpec { "EscapeChar=none".to_string(), "-o".to_string(), format!("ConnectTimeout={connect_timeout}"), - "-p".to_string(), - target.port().to_string(), - "-i".to_string(), - target.identity_file().display().to_string(), - "-l".to_string(), - target.user().to_string(), - "--".to_string(), - target.host().to_string(), - remote_command.clone(), ]; + if !target.compact_destination() { + args.splice(0..0, ["-F".to_string(), "/dev/null".to_string()]); + args.splice( + 4..4, + [ + "-o".to_string(), + format!("UserKnownHostsFile={}", target.known_hosts_file().display()), + "-o".to_string(), + "GlobalKnownHostsFile=/dev/null".to_string(), + "-o".to_string(), + "IdentitiesOnly=yes".to_string(), + "-o".to_string(), + "IdentityAgent=none".to_string(), + ], + ); + args.extend(["-p".to_string(), target.port().to_string(), "-i".to_string(), target.identity_file().display().to_string()]); + } else if target.port_explicit() { + args.extend(["-p".to_string(), target.port().to_string()]); + } + if !target.user().is_empty() { + args.extend(["-l".to_string(), target.user().to_string()]); + } + args.extend(["--".to_string(), target.host().to_string(), remote_command.clone()]); + Ok(Self { program: PathBuf::from(SSH_PROGRAM), args, remote_command }) } @@ -263,6 +278,7 @@ impl RemoteBackend for SshBackend { &'a self, request: AuthorizedRemoteRequest, control: RemoteExecutionControl, events: tokio::sync::mpsc::Sender, ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { let operation = request.request().operation().clone(); + let target_command = request.target_command().map(str::to_string); let request_id = request.request_id(); let session = self.session.clone(); let target = self.target.clone(); @@ -271,7 +287,8 @@ impl RemoteBackend for SshBackend { let artifact_policy = self.artifact_policy.clone(); let artifact_limits = self.artifact_limits; Box::pin(async move { - let backend = SshExecution { session, target, factory, uploads, artifact_policy, artifact_limits, control: control.clone() }; + let backend = + SshExecution { session, target, factory, uploads, artifact_policy, artifact_limits, target_command, control: control.clone() }; match operation { RemoteOperation::Sync(sync) => backend.execute_sync(request_id.0, sync.retain_capability(), events).await, RemoteOperation::Build(build) => backend.execute_build(request_id.0, &build, events).await, @@ -296,6 +313,7 @@ struct SshExecution { uploads: Arc>>, artifact_policy: ArtifactPolicy, artifact_limits: ArtifactLimits, + target_command: Option, control: RemoteExecutionControl, } @@ -448,28 +466,32 @@ impl SshExecution { return Err(RemoteBackendError::Failed(error)); } }; - let executable = match self.target.tools().get(build.tool().as_str()) { - Some(executable) => executable.clone(), - None => { - drop(claim); - self.cleanup_upload(request_id, upload_id).await; - return Err(RemoteBackendError::Transport { - class: RemoteFailureClass::WorkerUnavailable, - message: format!("remote tool is not configured: {}", build.tool().as_str()), - }); - } - }; let guest_env = build.env().to_vec(); let target_env = self.target.environment().iter().map(|(key, value)| (key.clone(), value.clone())).collect::>(); - let worker_build = match WorkerBuild::new( - build.tool().as_str(), - executable, - build.argv().to_vec(), - build.cwd().as_str(), - guest_env, - target_env, - upload_id, - ) { + let worker_build = match if self.target.compact_destination() { + WorkerBuild::new_command( + build.tool().as_str(), + self.target_command.as_deref().unwrap_or(build.tool().as_str()), + build.argv().to_vec(), + build.cwd().as_str(), + guest_env, + target_env, + upload_id, + ) + } else { + let executable = match self.target.tools().get(build.tool().as_str()) { + Some(executable) => executable.clone(), + None => { + drop(claim); + self.cleanup_upload(request_id, upload_id).await; + return Err(RemoteBackendError::Transport { + class: RemoteFailureClass::WorkerUnavailable, + message: format!("remote tool is not configured: {}", build.tool().as_str()), + }); + } + }; + WorkerBuild::new(build.tool().as_str(), executable, build.argv().to_vec(), build.cwd().as_str(), guest_env, target_env, upload_id) + } { Ok(worker_build) => worker_build, Err(error) => { drop(claim); @@ -506,7 +528,7 @@ impl SshExecution { .handshake( WorkerRequestId(request_id), WorkerSessionId(self.session.session_id().0), - if self.artifact_policy.is_enabled() { WORKER_ARTIFACT_PROTOCOL_VERSION } else { WORKER_PROTOCOL_VERSION }, + build_protocol_version(&self.target, self.artifact_policy.is_enabled()), ) .await?; Ok(connection) @@ -537,7 +559,7 @@ impl SshExecution { artifact_policy: self.artifact_policy.clone(), artifact_limits: self.artifact_limits, workspace_root: self.session.workspace_root().to_path_buf(), - protocol_version: if self.artifact_policy.is_enabled() { WORKER_ARTIFACT_PROTOCOL_VERSION } else { WORKER_PROTOCOL_VERSION }, + protocol_version: build_protocol_version(&self.target, self.artifact_policy.is_enabled()), control: &self.control, handshake: false, }, diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index 14bc92f..1a77428 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -1,10 +1,11 @@ use super::*; use crate::artifact::{ArtifactLimits, ArtifactPolicy}; +use crate::cfg::ProjectConfig; use crate::remote::{ RemoteAuthorizationPolicy, RemoteBuild, RemoteExecutionContext, RemoteRequest, RemoteSnapshotId, RemoteTool, RequestId, WorkspaceRelativePath, WorkspaceSessionId, }; -use crate::remote_target::RemoteTargetConfig; +use crate::remote_target::{BuildTargetCatalog, RemoteTargetConfig}; use crate::snapshot::{SnapshotBuilder, SnapshotExclusionPolicy, SnapshotLimits, SnapshotStore}; use crate::worker_protocol::{ self, WorkerArtifactEntry, WorkerArtifactSetId, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerSessionId, WorkerUploadEntry, @@ -22,6 +23,29 @@ use tokio::sync::oneshot; #[cfg(unix)] use std::os::unix::fs::PermissionsExt; +#[test] +fn compact_target_launch_uses_host_ssh_defaults_and_fixed_worker() { + let temp = tempdir().unwrap(); + let project = temp.path().join("project"); + fs::create_dir(&project).unwrap(); + let bunkerbox = project.join(".bunkerbox"); + fs::create_dir(&bunkerbox).unwrap(); + fs::write( + bunkerbox.join(crate::remote_target::REMOTE_PROJECT_CONFIG_FILE_NAME), + "targets:\n build:\n ssh: builder@build.example.test:2200\n workspace: /var/tmp/bunkerbox\n", + ) + .unwrap(); + let catalog = BuildTargetCatalog::load_optional(&project, &ProjectConfig::default()).unwrap().unwrap(); + let target = catalog.remote("build").unwrap().target(); + let spec = SshLaunchSpec::from_target(target).unwrap(); + assert_eq!(spec.program(), Path::new("/usr/bin/ssh")); + assert!(spec.args().windows(2).all(|pair| pair != ["-F", "/dev/null"])); + assert!(!spec.args().iter().any(|arg| arg == "-i" || arg == "UserKnownHostsFile=/dev/null")); + assert!(spec.args().windows(2).any(|pair| pair == ["-p", "2200"])); + assert!(spec.remote_command().contains("/usr/local/libexec/bunkerbox-worker")); + assert!(spec.remote_command().contains("--stdio")); +} + #[derive(Clone, Copy)] enum ScriptMode { Success, diff --git a/src/tui.rs b/src/tui.rs index 930b574..1b8489b 100644 --- a/src/tui.rs +++ b/src/tui.rs @@ -1,5 +1,7 @@ +use std::ffi::CString; use std::io; use std::os::fd::AsRawFd; +use std::os::unix::ffi::OsStrExt; use std::os::unix::io::RawFd; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; @@ -13,10 +15,12 @@ use crossterm::terminal::{self, EnterAlternateScreen, LeaveAlternateScreen}; use crossterm::ExecutableCommand; use ratatui::backend::CrosstermBackend; use ratatui::prelude::*; +use ratatui::text::{Line, Span}; use ratatui::widgets::{Block, BorderType, Borders, Clear, Padding, Paragraph, Widget, Wrap}; use ratatui::Terminal; +use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; use crate::vscomm::{self, parse_triggers, Trigger}; mod palette; @@ -87,6 +91,99 @@ pub struct OverlayState { pub last_content_scan: Instant, } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum HostPopup { + None, + Targets, + Help, +} + +struct HostUiState { + popup: HostPopup, + target_index: usize, + confirmation: Option<(String, Instant)>, + free_bytes: Option, + last_free_refresh: Instant, +} + +impl HostUiState { + fn new(catalog: &BuildTargetCatalog) -> Self { + Self { + popup: HostPopup::None, + target_index: catalog.summaries().iter().position(|target| target.label() == "localhost").unwrap_or(0), + confirmation: None, + free_bytes: None, + last_free_refresh: Instant::now() - Duration::from_secs(10), + } + } + + fn refresh_free_space(&mut self, catalog: &BuildTargetCatalog) { + if self.last_free_refresh.elapsed() < Duration::from_secs(5) { + return; + } + self.last_free_refresh = Instant::now(); + self.free_bytes = local_free_space(catalog.project_root()); + } + + fn handle_key(&mut self, key: KeyEvent, catalog: &BuildTargetCatalog, active: &ActiveBuildTarget) -> bool { + match self.popup { + HostPopup::Targets => { + match key.code { + KeyCode::Up => { + self.target_index = self.target_index.saturating_sub(1); + } + KeyCode::Down => { + self.target_index = (self.target_index + 1).min(catalog.summaries().len().saturating_sub(1)); + } + KeyCode::Enter => { + if let Some(target) = catalog.summaries().get(self.target_index) { + if active.select(catalog, target.label()).is_ok() { + self.confirmation = Some((format!("Build target: {}", target.label()), Instant::now())); + } + } + self.popup = HostPopup::None; + } + KeyCode::Esc => { + self.popup = HostPopup::None; + } + _ => {} + } + return true; + } + HostPopup::Help => { + if key.code == KeyCode::Esc { + self.popup = HostPopup::None; + } + return true; + } + HostPopup::None => {} + } + + if key.code == KeyCode::Char('b') && key.modifiers == (KeyModifiers::CONTROL | KeyModifiers::ALT) { + let current = active.current(); + self.target_index = catalog.summaries().iter().position(|target| target.label() == current).unwrap_or(0); + self.popup = HostPopup::Targets; + return true; + } + if key.code == KeyCode::Char('h') && key.modifiers == (KeyModifiers::CONTROL | KeyModifiers::ALT) { + self.popup = HostPopup::Help; + return true; + } + false + } +} + +pub fn show_host_error(overlay: &Arc>, title: &str, message: &str) { + if let Ok(mut state) = overlay.lock() { + state.error_toast = + Some(ErrorToast { title: title.chars().take(80).collect(), message: message.chars().take(512).collect(), shown_at: Instant::now() }); + } +} + +pub fn guest_rows(physical_rows: u16) -> u16 { + physical_rows.saturating_sub(1).max(1) +} + impl Default for OverlayState { fn default() -> Self { Self::new() @@ -490,7 +587,10 @@ fn push_dec_special_graphic(out: &mut Vec, byte: u8) { out.extend_from_slice(mapped.encode_utf8(&mut buf).as_bytes()); } -fn hold_error_popup(overlay: &Arc>, term: &mut Terminal>, screen: &vt100::Screen) { +fn hold_error_popup( + overlay: &Arc>, term: &mut Terminal>, screen: &vt100::Screen, host: &HostUiState, + catalog: &BuildTargetCatalog, active: &ActiveBuildTarget, +) { let is_error = overlay.lock().unwrap().has_error; if !is_error { return; @@ -500,7 +600,7 @@ fn hold_error_popup(overlay: &Arc>, term: &mut Terminal bool { /// /// `overlay` is shared with the VSOCK status listener so VM-originated /// UI commands can update popups, progress bars, and status text. +#[allow(clippy::too_many_arguments)] pub fn event_loop( master_fd: RawFd, rows: u16, cols: u16, status_fd: RawFd, startup_status_fd: RawFd, overlay: Arc>, + catalog: BuildTargetCatalog, active: ActiveBuildTarget, ) -> Result<(), String> { let stdin_fd = io::stdin().as_raw_fd(); @@ -560,7 +662,7 @@ pub fn event_loop( stdout.execute(EnterAlternateScreen).map_err(|e| format!("alt screen: {e}"))?; stdout.execute(cursor::Show).map_err(|e| format!("cursor: {e}"))?; - let mut term = Term::new(rows, cols); + let mut term = Term::new(guest_rows(rows), cols); let backend = CrosstermBackend::new(stdout); let mut terminal = Terminal::new(backend).map_err(|e| format!("terminal: {e}"))?; @@ -569,6 +671,7 @@ pub fn event_loop( let mut last_rows = rows; let mut last_cols = cols; + let mut host = HostUiState::new(&catalog); let mut status_buf = Vec::new(); let mut startup_status_buf = Vec::new(); @@ -585,8 +688,8 @@ pub fn event_loop( if new_cols != last_cols || new_rows != last_rows { last_cols = new_cols; last_rows = new_rows; - term.set_size(new_rows, new_cols); - let ws = libc::winsize { ws_row: new_rows, ws_col: new_cols, ws_xpixel: 0, ws_ypixel: 0 }; + term.set_size(guest_rows(new_rows), new_cols); + let ws = libc::winsize { ws_row: guest_rows(new_rows), ws_col: new_cols, ws_xpixel: 0, ws_ypixel: 0 }; unsafe { libc::ioctl(master_fd, libc::TIOCSWINSZ, &ws); } @@ -638,7 +741,7 @@ pub fn event_loop( } pty_output = true; } else { - hold_error_popup(&overlay, &mut terminal, term.screen()); + hold_error_popup(&overlay, &mut terminal, term.screen(), &host, &catalog, &active); break; } } @@ -664,6 +767,8 @@ pub fn event_loop( let is_password = overlay.lock().is_ok_and(|s| matches!(s.popup.content, popup::PopupContent::Password { .. })); if is_password { handle_password_key(&overlay, status_fd, key); + } else if host.handle_key(key, &catalog, &active) { + continue; } else if let Some(bytes) = key_to_bytes(&key, term.application_cursor_keys()) { unsafe { libc::write(master_fd, bytes.as_ptr() as *const libc::c_void, bytes.len()); @@ -671,6 +776,9 @@ pub fn event_loop( } } Event::Mouse(mouse) => { + if mouse.row >= guest_rows(last_rows) { + continue; + } if let Some(bytes) = mouse_to_bytes(mouse, term.mouse_tracking(), term.mouse_encoding()) { unsafe { libc::write(master_fd, bytes.as_ptr() as *const libc::c_void, bytes.len()); @@ -685,6 +793,7 @@ pub fn event_loop( { let mut state = overlay.lock().unwrap(); let now = Instant::now(); + host.refresh_free_space(&catalog); if pty_output { for action in &mut state.pending { @@ -737,7 +846,7 @@ pub fn event_loop( } } - if let Err(err) = terminal.draw(|f| render_frame(f, term.screen(), &state)) { + if let Err(err) = terminal.draw(|f| render_frame(f, term.screen(), &state, &host, &catalog, &active)) { cleanup_terminal(&mut terminal, mouse_capture_enabled); return Err(format!("draw: {err}")); } @@ -935,21 +1044,124 @@ fn to_ratatui_color(c: vt100::Color) -> Color { } } -/// Renders one frame: writes every vt100 screen cell to the ratatui buffer -/// with full color and attributes, then draws overlay widgets (popup, -/// progress bar, status box) on top. -fn render_frame(f: &mut Frame, screen: &vt100::Screen, overlay: &OverlayState) { +fn local_free_space(path: &std::path::Path) -> Option { + let path = CString::new(path.as_os_str().as_bytes()).ok()?; + let mut stats = unsafe { std::mem::zeroed::() }; + let result = unsafe { libc::statvfs(path.as_ptr(), &mut stats) }; + if result != 0 { + return None; + } + stats.f_bavail.checked_mul(stats.f_frsize) +} + +fn format_bytes(bytes: u64) -> String { + const UNITS: [&str; 4] = ["B", "MiB", "GiB", "TiB"]; + let mut value = bytes as f64; + let mut unit = 0; + while value >= 1024.0 && unit < UNITS.len() - 1 { + value /= 1024.0; + unit += 1; + } + if unit == 0 { + format!("{} {}", bytes, UNITS[unit]) + } else { + format!("{value:.1} {}", UNITS[unit]) + } +} + +fn centered_rect(area: Rect, width: u16, height: u16) -> Rect { + let width = width.min(area.width.saturating_sub(2)); + let height = height.min(area.height.saturating_sub(2)); + Rect { x: area.x + area.width.saturating_sub(width) / 2, y: area.y + area.height.saturating_sub(height) / 2, width, height } +} + +fn render_status_bar(area: Rect, buf: &mut Buffer, host: &HostUiState, active: &ActiveBuildTarget) { + if area.height == 0 || area.width == 0 { + return; + } + let target = active.current(); + let free = host.free_bytes.map(format_bytes).unwrap_or_else(|| "unknown".to_string()); + let target_segment = host + .confirmation + .as_ref() + .filter(|(_, shown_at)| shown_at.elapsed() < Duration::from_secs(3)) + .map_or_else(|| format!("Target: {target}"), |(message, _)| format!("Target: {target} ({message})")); + let mut segments = vec![target_segment, "Ctrl+Alt+B Targets".to_string(), "Ctrl+Alt+H Help".to_string(), format!("Local workspace free: {free}")]; + while segments.len() > 1 { + let text = format!(" {}", segments.join(" | ")); + if text.chars().count() <= usize::from(area.width) { + break; + } + segments.pop(); + } + let mut text = format!(" {}", segments.join(" | ")); + if text.chars().count() > usize::from(area.width) { + text = text.chars().take(usize::from(area.width)).collect(); + } + Paragraph::new(text).style(Style::default().fg(palette::FG).bg(palette::BG_1)).render(area, buf); +} + +fn render_host_popup(area: Rect, buf: &mut Buffer, host: &HostUiState, catalog: &BuildTargetCatalog) { + let (title, lines, height) = match host.popup { + HostPopup::None => return, + HostPopup::Help => ( + "Bunkerbox Help", + vec![ + Line::from(Span::styled("Ctrl-Alt-B", Style::default().fg(palette::ACCENT))), + Line::from("Select Build Target"), + Line::from(Span::styled("Ctrl-Alt-H", Style::default().fg(palette::ACCENT))), + Line::from("Show this help"), + Line::from(Span::styled("Esc", Style::default().fg(palette::ACCENT))), + Line::from("Close host popup"), + ], + 10, + ), + HostPopup::Targets => { + let mut lines = Vec::new(); + for (index, target) in catalog.summaries().iter().enumerate() { + let marker = if index == host.target_index { "> " } else { " " }; + lines.push(Line::from(vec![ + Span::styled(marker, Style::default().fg(palette::ACCENT)), + Span::styled(target.label(), Style::default().fg(palette::FG)), + Span::styled(format!(": {}", target.workspace()), Style::default().fg(palette::MUTED)), + ])); + } + let height = (lines.len() as u16 + 4).max(5); + ("Build Targets", lines, height) + } + }; + let popup_area = centered_rect(area, 64, height); + if popup_area.width < 4 || popup_area.height < 3 { + return; + } + Clear.render(popup_area, buf); + let block = Block::default() + .title(title) + .borders(Borders::ALL) + .border_type(BorderType::Rounded) + .border_style(Style::default().fg(palette::ACCENT)) + .style(Style::default().fg(palette::FG).bg(palette::POPUP_BG)) + .padding(Padding::horizontal(1)); + Paragraph::new(lines).block(block).render(popup_area, buf); +} + +/// Renders one frame: writes the guest vt100 screen into the reduced guest +/// viewport, then draws host-owned status and popup controls. +fn render_frame( + f: &mut Frame, screen: &vt100::Screen, overlay: &OverlayState, host: &HostUiState, catalog: &BuildTargetCatalog, active: &ActiveBuildTarget, +) { let area = f.area(); + let guest_area = Rect { height: area.height.saturating_sub(1), ..area }; let (rows, cols) = screen.size(); - let max_rows = area.height.min(rows); - let max_cols = area.width.min(cols); + let max_rows = guest_area.height.min(rows); + let max_cols = guest_area.width.min(cols); let buf = f.buffer_mut(); for row in 0..max_rows { let mut col: u16 = 0; while col < max_cols { - let x = area.x + col; - let y = area.y + row; + let x = guest_area.x + col; + let y = guest_area.y + row; if let Some(cell) = screen.cell(row, col) { if cell.is_wide_continuation() { @@ -1006,13 +1218,15 @@ fn render_frame(f: &mut Frame, screen: &vt100::Screen, overlay: &OverlayState) { { let buf = f.buffer_mut(); - overlay.popup.render(area, buf); - render_error_toast(area, buf, overlay.error_toast.as_ref()); + overlay.popup.render(guest_area, buf); + render_error_toast(guest_area, buf, overlay.error_toast.as_ref()); + render_status_bar(Rect { y: area.bottom().saturating_sub(1), height: area.height.min(1), ..area }, buf, host, active); + render_host_popup(area, buf, host, catalog); } let (cursor_row, cursor_col) = screen.cursor_position(); if cursor_row < max_rows && cursor_col < max_cols { - f.set_cursor_position((area.x + cursor_col, area.y + cursor_row)); + f.set_cursor_position((guest_area.x + cursor_col, guest_area.y + cursor_row)); } } diff --git a/src/tui_ut.rs b/src/tui_ut.rs index 22c7e4f..82b0ef2 100644 --- a/src/tui_ut.rs +++ b/src/tui_ut.rs @@ -1,7 +1,29 @@ -use super::{dispatch_ui_command, mouse_to_bytes, process_status_bytes, MouseEncoding, MouseTracking, OverlayState, Term}; -use crossterm::event::{KeyModifiers, MouseButton, MouseEvent, MouseEventKind}; +use super::{ + dispatch_ui_command, guest_rows, mouse_to_bytes, process_status_bytes, HostPopup, HostUiState, MouseEncoding, MouseTracking, OverlayState, Term, +}; +use crate::cfg::ProjectConfig; +use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; +use crossterm::event::{KeyCode, KeyEvent, KeyModifiers, MouseButton, MouseEvent, MouseEventKind}; +use std::path::PathBuf; use std::sync::{Arc, Mutex}; +#[test] +fn guest_viewport_reserves_exactly_one_physical_row() { + assert_eq!(guest_rows(24), 23); + assert_eq!(guest_rows(1), 1); + assert_eq!(guest_rows(0), 1); +} + +#[test] +fn host_target_shortcut_is_consumed_before_guest_bytes() { + let catalog = BuildTargetCatalog::localhost_only(PathBuf::from("/tmp/project"), ProjectConfig::default()).unwrap(); + let active = ActiveBuildTarget::new(); + let mut host = HostUiState::new(&catalog); + let key = KeyEvent::new(KeyCode::Char('b'), KeyModifiers::CONTROL | KeyModifiers::ALT); + assert!(host.handle_key(key, &catalog, &active)); + assert_eq!(host.popup, HostPopup::Targets); +} + #[test] fn internal_error_creates_a_non_modal_toast() { let mut state = OverlayState::new(); From eecb69d8308e5375dc94c4fe9cfd7acce0aa90c7 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 15:26:33 +0200 Subject: [PATCH 44/52] Update config docs --- docs/config/project.md | 8 +++++++- docs/reference/config-schema.md | 12 ++++++++++-- 2 files changed, 17 insertions(+), 3 deletions(-) diff --git a/docs/config/project.md b/docs/config/project.md index eef4008..05d1802 100644 --- a/docs/config/project.md +++ b/docs/config/project.md @@ -209,7 +209,13 @@ request. `command` is an optional basename-only target override, such as Remote targets are declared separately in `.bunkerbox/remote.conf`. If that file is absent, Bunkerbox starts with the implicit `localhost` target only. -The host TUI starts on `localhost`; use `Ctrl+Alt+B` to select a remote target. +The host TUI starts on `localhost`; use `Ctrl+Alt+B` to select a remote target +or `Ctrl+Alt+S` to open the host-owned Remote Setup forms. Setup can edit the +target label, compact SSH destination, workspace, remote tools, environment +names, snapshot exclusions, artifact paths, and resource limits. Save writes +`.bunkerbox/remote.conf` atomically and applies changes on the next Bunkerbox +run only. It does not probe SSH or test a connection. A malformed existing +configuration is shown as an error and is never silently overwritten. The guest and AI have no target-selection command, and a remote failure never falls back to local execution. diff --git a/docs/reference/config-schema.md b/docs/reference/config-schema.md index 9d39a31..ff53608 100644 --- a/docs/reference/config-schema.md +++ b/docs/reference/config-schema.md @@ -135,8 +135,16 @@ targets: `ssh` accepts `[user@]host[:port]` or a host-side OpenSSH alias. `workspace` must be an absolute normalized target path. `localhost` is implicit, always listed first, and selected at startup. Press `Ctrl+Alt+B` in the host TUI to -choose another target; `Ctrl+Alt+H` shows the host controls. Target selection -is frozen per transaction, and remote failures do not retry locally. +choose another target; `Ctrl+Alt+S` opens Remote Setup; and `Ctrl+Alt+H` shows +the host controls. Remote Setup edits only the compact target fields and the +`project.remote` overlay: tool allowlists, environment names, snapshot +exclusions, artifact paths, and optional resource limits. Blank resource fields +remain unset, and omitted overlay fields continue inheriting `project.conf`. +Save validates the complete draft and atomically writes the file for the next +Bunkerbox run; it does not change the current catalog, target, or backend and +does not perform SSH probing. Existing malformed or unsafe configurations are +reported rather than replaced. Target selection is frozen per transaction, and +remote failures do not retry locally. ## Sandbox profile From 503d3b614dae3db75ad30d8dc6abccbd90093767 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 15:26:45 +0200 Subject: [PATCH 45/52] Add TUI for remote setup --- src/cfg.rs | 2 +- src/remote_target.rs | 389 ++++++++++- src/tui.rs | 1513 +++++++++++++++++++++++++++++++++++++++++- 3 files changed, 1863 insertions(+), 41 deletions(-) diff --git a/src/cfg.rs b/src/cfg.rs index 5f84cf3..15a6fc5 100644 --- a/src/cfg.rs +++ b/src/cfg.rs @@ -230,7 +230,7 @@ pub struct RemoteSection { #[serde(deny_unknown_fields)] pub struct RemoteToolSpec { pub name: String, - #[serde(default)] + #[serde(default, skip_serializing_if = "Option::is_none")] pub command: Option, #[serde(default, rename = "allow-args")] pub allow_args: bool, diff --git a/src/remote_target.rs b/src/remote_target.rs index 576435e..cdb8bb2 100644 --- a/src/remote_target.rs +++ b/src/remote_target.rs @@ -1,12 +1,13 @@ use crate::artifact::{ArtifactLimits, ArtifactPolicy}; -use crate::cfg::{ProjectConfig, RemoteSection}; +use crate::cfg::{ProjectConfig, RemoteSection, RemoteToolSpec}; use crate::remote::{RemoteEnvironmentPolicy, RemoteToolPolicy}; use serde::de::{self, MapAccess, Visitor}; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use std::collections::BTreeMap; use std::env; use std::fmt; -use std::fs::{self, File}; +use std::fs::{self, File, OpenOptions}; +use std::io::Write; use std::marker::PhantomData; use std::net::IpAddr; use std::path::{Component, Path, PathBuf}; @@ -14,8 +15,10 @@ use std::str::FromStr; use std::sync::{Arc, Mutex}; use std::time::Duration; +use std::sync::atomic::{AtomicU64, Ordering}; + #[cfg(unix)] -use std::os::unix::fs::PermissionsExt; +use std::os::unix::fs::{OpenOptionsExt, PermissionsExt}; pub const CONFIG_VERSION: u64 = 1; pub const CONFIG_FILE_NAME: &str = "remote-targets.yaml"; @@ -31,6 +34,7 @@ const MAX_NAME_BYTES: usize = 64; const MAX_TOOL_PATH_BYTES: usize = 4096; const MAX_ENV_NAME_BYTES: usize = 256; const MAX_ENV_VALUE_BYTES: usize = 16 * 1024; +static NEXT_CONFIG_TEMP: AtomicU64 = AtomicU64::new(1); /// Selects the host-side facility used for a project. #[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)] @@ -517,7 +521,7 @@ pub fn default_config_path_with(xdg_config_home: Option<&Path>, home: Option<&Pa ConfigPathHelper::from_paths(xdg_config_home, home).config_path() } -#[derive(Debug)] +#[derive(Clone, Debug)] struct UniqueMap(BTreeMap); impl Default for UniqueMap { @@ -668,14 +672,15 @@ struct RawResources { max_active_builds: Option, } -#[derive(Deserialize)] +#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] #[serde(untagged)] -#[derive(Debug)] -enum RawQuantity { +pub enum RemoteQuantity { Integer(u64), Text(String), } +type RawQuantity = RemoteQuantity; + fn validate_target(name: String, raw: RawTarget) -> Result { if raw.transport != "ssh" { return Err("remote target transport must be 'ssh'".to_string()); @@ -1243,14 +1248,14 @@ impl Default for ActiveBuildTarget { } } -#[derive(Debug, Deserialize)] +#[derive(Clone, Debug, Deserialize)] #[serde(deny_unknown_fields)] struct RawRemoteProjectConfig { #[serde(default)] targets: UniqueMap, } -#[derive(Debug, Deserialize)] +#[derive(Clone, Debug, Deserialize)] #[serde(deny_unknown_fields)] struct RawCompactTarget { ssh: String, @@ -1261,14 +1266,14 @@ struct RawCompactTarget { resources: Option, } -#[derive(Debug, Deserialize)] +#[derive(Clone, Debug, Deserialize)] #[serde(deny_unknown_fields)] struct RawTargetProject { #[serde(default)] remote: Option, } -#[derive(Debug, Deserialize)] +#[derive(Clone, Debug, Deserialize)] #[serde(deny_unknown_fields)] struct RawRemoteOverlay { #[serde(default)] @@ -1281,7 +1286,7 @@ struct RawRemoteOverlay { artifacts: Option>, } -#[derive(Debug, Default, Deserialize)] +#[derive(Clone, Debug, Default, Deserialize)] #[serde(deny_unknown_fields)] struct RawTargetResources { #[serde( @@ -1343,6 +1348,318 @@ struct RawTargetResources { max_active_builds: Option, } +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct RemoteConfigDraft { + #[serde(default)] + pub targets: BTreeMap, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct RemoteTargetDraft { + pub ssh: String, + pub workspace: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub project: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub resources: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct RemoteProjectOverlayDraft { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub remote: Option, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct RemoteOverlayDraft { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub exclude: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub environment: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub artifacts: Option>, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct RemoteResourceOverridesDraft { + #[serde(default, rename = "connect-timeout-seconds", skip_serializing_if = "Option::is_none")] + pub connect_timeout: Option, + #[serde(default, rename = "sync-timeout-seconds", skip_serializing_if = "Option::is_none")] + pub sync_timeout: Option, + #[serde(default, rename = "build-timeout-seconds", skip_serializing_if = "Option::is_none")] + pub build_timeout: Option, + #[serde(default, rename = "max-output-bytes", skip_serializing_if = "Option::is_none")] + pub max_output: Option, + #[serde(default, rename = "idle-output-timeout-seconds", skip_serializing_if = "Option::is_none")] + pub idle_output_timeout: Option, + #[serde(default, rename = "cleanup-timeout-seconds", skip_serializing_if = "Option::is_none")] + pub cleanup_timeout: Option, + #[serde(default, rename = "artifact-timeout-seconds", skip_serializing_if = "Option::is_none")] + pub artifact_timeout: Option, + #[serde(default, rename = "max-artifact-bytes", skip_serializing_if = "Option::is_none")] + pub max_artifact_bytes: Option, + #[serde(default, rename = "max-artifact-total-bytes", skip_serializing_if = "Option::is_none")] + pub max_artifact_total_bytes: Option, + #[serde(default, rename = "max-artifact-entries", skip_serializing_if = "Option::is_none")] + pub max_artifact_entries: Option, + #[serde(default, rename = "max-worker-uploads", skip_serializing_if = "Option::is_none")] + pub max_worker_uploads: Option, + #[serde(default, rename = "max-worker-upload-bytes", skip_serializing_if = "Option::is_none")] + pub max_worker_upload_bytes: Option, + #[serde(default, rename = "max-worker-jobs", skip_serializing_if = "Option::is_none")] + pub max_worker_jobs: Option, + #[serde(default, rename = "max-worker-job-bytes", skip_serializing_if = "Option::is_none")] + pub max_worker_job_bytes: Option, + #[serde(default, rename = "max-worker-artifact-spools", skip_serializing_if = "Option::is_none")] + pub max_worker_artifact_spools: Option, + #[serde(default, rename = "max-worker-artifact-spool-bytes", skip_serializing_if = "Option::is_none")] + pub max_worker_artifact_spool_bytes: Option, + #[serde(default, rename = "max-worker-state-entries", skip_serializing_if = "Option::is_none")] + pub max_worker_state_entries: Option, + #[serde(default, rename = "max-active-builds", skip_serializing_if = "Option::is_none")] + pub max_active_builds: Option, +} + +impl RemoteResourceOverridesDraft { + pub fn is_empty(&self) -> bool { + self == &Self::default() + } +} + +impl From for RemoteConfigDraft { + fn from(raw: RawRemoteProjectConfig) -> Self { + Self { targets: raw.targets.0.into_iter().map(|(label, target)| (label, target.into())).collect() } + } +} + +impl From for RemoteTargetDraft { + fn from(raw: RawCompactTarget) -> Self { + Self { ssh: raw.ssh, workspace: raw.workspace, project: raw.project.map(Into::into), resources: raw.resources.map(Into::into) } + } +} + +impl From for RemoteProjectOverlayDraft { + fn from(raw: RawTargetProject) -> Self { + Self { remote: raw.remote.map(Into::into) } + } +} + +impl From for RemoteOverlayDraft { + fn from(raw: RawRemoteOverlay) -> Self { + Self { exclude: raw.exclude, environment: raw.environment, tools: raw.tools, artifacts: raw.artifacts } + } +} + +impl From for RemoteResourceOverridesDraft { + fn from(raw: RawTargetResources) -> Self { + Self { + connect_timeout: raw.connect_timeout, + sync_timeout: raw.sync_timeout, + build_timeout: raw.build_timeout, + max_output: raw.max_output, + idle_output_timeout: raw.idle_output_timeout, + cleanup_timeout: raw.cleanup_timeout, + artifact_timeout: raw.artifact_timeout, + max_artifact_bytes: raw.max_artifact_bytes, + max_artifact_total_bytes: raw.max_artifact_total_bytes, + max_artifact_entries: raw.max_artifact_entries, + max_worker_uploads: raw.max_worker_uploads, + max_worker_upload_bytes: raw.max_worker_upload_bytes, + max_worker_jobs: raw.max_worker_jobs, + max_worker_job_bytes: raw.max_worker_job_bytes, + max_worker_artifact_spools: raw.max_worker_artifact_spools, + max_worker_artifact_spool_bytes: raw.max_worker_artifact_spool_bytes, + max_worker_state_entries: raw.max_worker_state_entries, + max_active_builds: raw.max_active_builds, + } + } +} + +impl RemoteTargetDraft { + fn raw_project(&self) -> Option { + self.project.as_ref().map(|project| RawTargetProject { + remote: project.remote.as_ref().map(|remote| RawRemoteOverlay { + exclude: remote.exclude.clone(), + environment: remote.environment.clone(), + tools: remote.tools.clone(), + artifacts: remote.artifacts.clone(), + }), + }) + } + + fn raw_resources(&self) -> Option { + self.resources.as_ref().map(|resources| RawTargetResources { + connect_timeout: resources.connect_timeout.clone(), + sync_timeout: resources.sync_timeout.clone(), + build_timeout: resources.build_timeout.clone(), + max_output: resources.max_output.clone(), + idle_output_timeout: resources.idle_output_timeout.clone(), + cleanup_timeout: resources.cleanup_timeout.clone(), + artifact_timeout: resources.artifact_timeout.clone(), + max_artifact_bytes: resources.max_artifact_bytes.clone(), + max_artifact_total_bytes: resources.max_artifact_total_bytes.clone(), + max_artifact_entries: resources.max_artifact_entries, + max_worker_uploads: resources.max_worker_uploads, + max_worker_upload_bytes: resources.max_worker_upload_bytes.clone(), + max_worker_jobs: resources.max_worker_jobs, + max_worker_job_bytes: resources.max_worker_job_bytes.clone(), + max_worker_artifact_spools: resources.max_worker_artifact_spools, + max_worker_artifact_spool_bytes: resources.max_worker_artifact_spool_bytes.clone(), + max_worker_state_entries: resources.max_worker_state_entries, + max_active_builds: resources.max_active_builds, + }) + } +} + +impl RemoteConfigDraft { + pub fn load_optional(project_root: &Path, base_project: &ProjectConfig) -> Result, String> { + if !project_root.is_absolute() { + return Err("project root must be absolute".to_string()); + } + let config_dir = project_root.join(".bunkerbox"); + match fs::symlink_metadata(&config_dir) { + Ok(metadata) if metadata.file_type().is_symlink() || !metadata.file_type().is_dir() => { + return Err(format!("{} must be a real directory", config_dir.display())) + } + Ok(_) => {} + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(format!("inspect {}: {error}", config_dir.display())), + } + let path = config_dir.join(REMOTE_PROJECT_CONFIG_FILE_NAME); + let metadata = match fs::symlink_metadata(&path) { + Ok(metadata) => metadata, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None), + Err(error) => return Err(format!("failed to inspect {}: {error}", path.display())), + }; + if metadata.file_type().is_symlink() || !metadata.file_type().is_file() { + return Err(format!("{} must be a regular file", path.display())); + } + let contents = fs::read_to_string(&path).map_err(|error| format!("failed to read {}: {error}", path.display()))?; + let raw: RawRemoteProjectConfig = serde_yaml::from_str(&contents).map_err(|error| format!("failed to parse {}: {error}", path.display()))?; + let draft = Self::from(raw); + draft.validate(base_project)?; + Ok(Some(draft)) + } + + pub fn validate(&self, base_project: &ProjectConfig) -> Result<(), String> { + for (label, target) in &self.targets { + self.validate_target(label, target, base_project)?; + } + Ok(()) + } + + pub fn validate_target(&self, label: &str, target: &RemoteTargetDraft, base_project: &ProjectConfig) -> Result<(), String> { + validate_name("target label", label)?; + if label == "localhost" { + return Err("remote target label 'localhost' is reserved".to_string()); + } + let project = merge_remote_overlay(&base_project.project.remote, target.raw_project().as_ref())?; + validate_remote_section(&project)?; + let resources = validate_compact_resources(target.raw_resources())?; + let _ssh_target = SshTarget::from_compact(label.to_string(), &target.ssh, target.workspace.clone(), resources)?; + let artifact_policy = ArtifactPolicy::new(project.artifacts)?; + artifact_policy.validate_limits(resources.artifact_limits())?; + Ok(()) + } + + pub fn to_yaml(&self) -> Result { + let mut serializable = self.clone(); + for target in serializable.targets.values_mut() { + if target.resources.as_ref().is_some_and(RemoteResourceOverridesDraft::is_empty) { + target.resources = None; + } + } + serde_yaml::to_string(&serializable).map_err(|error| format!("serialize remote configuration: {error}")) + } + + pub fn write_atomic(&self, project_root: &Path, base_project: &ProjectConfig) -> Result<(), String> { + self.validate(base_project)?; + if !project_root.is_absolute() { + return Err("project root must be absolute".to_string()); + } + let config_dir = project_root.join(".bunkerbox"); + ensure_real_directory(&config_dir)?; + + let destination = config_dir.join(REMOTE_PROJECT_CONFIG_FILE_NAME); + ensure_regular_destination(&destination)?; + + let contents = self.to_yaml()?.into_bytes(); + let mut temporary = None; + let mut file = None; + for _ in 0..32 { + let sequence = NEXT_CONFIG_TEMP.fetch_add(1, Ordering::Relaxed); + let path = config_dir.join(format!(".{REMOTE_PROJECT_CONFIG_FILE_NAME}.tmp-{}-{sequence}", std::process::id())); + let mut options = OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + options.mode(0o600).custom_flags(libc::O_NOFOLLOW); + match options.open(&path) { + Ok(value) => { + temporary = Some(path); + file = Some(value); + break; + } + Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {} + Err(error) => return Err(format!("create remote configuration temporary file: {error}")), + } + } + let temporary = temporary.ok_or_else(|| "could not create remote configuration temporary file".to_string())?; + let mut file = file.expect("remote configuration temporary file exists with its path"); + let result = (|| { + file.write_all(&contents).map_err(|error| format!("write remote configuration: {error}"))?; + file.sync_all().map_err(|error| format!("sync remote configuration: {error}"))?; + drop(file); + ensure_regular_destination(&destination)?; + fs::rename(&temporary, &destination).map_err(|error| format!("publish remote configuration: {error}"))?; + File::open(&config_dir) + .and_then(|directory| directory.sync_all()) + .map_err(|error| format!("sync remote configuration directory: {error}")) + })(); + if result.is_err() { + let _ = fs::remove_file(&temporary); + } + result + } +} + +fn ensure_real_directory(path: &Path) -> Result<(), String> { + match fs::symlink_metadata(path) { + Ok(metadata) => { + if metadata.file_type().is_symlink() || !metadata.file_type().is_dir() { + Err(format!("{} must be a real directory", path.display())) + } else { + Ok(()) + } + } + Err(error) if error.kind() == std::io::ErrorKind::NotFound => { + match fs::create_dir(path) { + Ok(()) => {} + Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => {} + Err(error) => return Err(format!("create {}: {error}", path.display())), + } + let metadata = fs::symlink_metadata(path).map_err(|error| format!("inspect {}: {error}", path.display()))?; + if metadata.file_type().is_symlink() || !metadata.file_type().is_dir() { + return Err(format!("{} must be a real directory", path.display())); + } + Ok(()) + } + Err(error) => Err(format!("inspect {}: {error}", path.display())), + } +} + +fn ensure_regular_destination(path: &Path) -> Result<(), String> { + match fs::symlink_metadata(path) { + Ok(metadata) if metadata.file_type().is_symlink() || !metadata.file_type().is_file() => { + Err(format!("{} must be a regular file", path.display())) + } + Ok(_) => Ok(()), + Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()), + Err(error) => Err(format!("inspect {}: {error}", path.display())), + } +} + fn merge_remote_overlay(base: &RemoteSection, target: Option<&RawTargetProject>) -> Result { let Some(target) = target.and_then(|project| project.remote.as_ref()) else { return Ok(base.clone()) }; let mut merged = base.clone(); @@ -1378,6 +1695,52 @@ fn validate_remote_section(section: &RemoteSection) -> Result<(), String> { Ok(()) } +pub fn validate_remote_target_label(label: &str) -> Result<(), String> { + validate_name("target label", label)?; + if label == "localhost" { + return Err("remote target label 'localhost' is reserved".to_string()); + } + Ok(()) +} + +pub fn validate_ssh_destination(value: &str) -> Result<(), String> { + parse_ssh_destination(value).map(|_| ()) +} + +pub fn validate_workspace_path(value: &str) -> Result<(), String> { + validate_remote_path("workspace", value, true).map(|_| ()) +} + +pub fn validate_remote_tool_spec(tool: &RemoteToolSpec) -> Result<(), String> { + crate::remote::validate_remote_wrapper_name(tool.name.clone())?; + if let Some(command) = &tool.command { + crate::remote::validate_remote_wrapper_name(command.clone())?; + } + Ok(()) +} + +pub fn validate_remote_environment_names(names: &[String]) -> Result<(), String> { + RemoteEnvironmentPolicy::from_names(names.to_vec()).map(|_| ()) +} + +pub fn validate_remote_exclusion_entries(entries: &[String]) -> Result<(), String> { + crate::snapshot::SnapshotExclusionPolicy::from_patterns(entries.to_vec()).map(|_| ()) +} + +pub fn validate_remote_artifact_paths(entries: &[String]) -> Result<(), String> { + ArtifactPolicy::new(entries.to_vec()).map(|_| ()) +} + +pub fn validate_remote_resource_overrides(resources: &RemoteResourceOverridesDraft) -> Result { + let target = RemoteTargetDraft { + ssh: "builder@localhost".to_string(), + workspace: "/workspace".to_string(), + project: None, + resources: Some(resources.clone()), + }; + validate_compact_resources(target.raw_resources()) +} + fn validate_compact_resources(raw: Option) -> Result { let defaults = crate::remote::RemoteResourcePolicy::default(); let artifact_defaults = ArtifactLimits::default(); diff --git a/src/tui.rs b/src/tui.rs index 1b8489b..0990560 100644 --- a/src/tui.rs +++ b/src/tui.rs @@ -20,7 +20,13 @@ use ratatui::widgets::{Block, BorderType, Borders, Clear, Padding, Paragraph, Wi use ratatui::Terminal; -use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; +use crate::cfg::RemoteToolSpec; +use crate::remote_target::{ + validate_remote_artifact_paths, validate_remote_environment_names, validate_remote_exclusion_entries, validate_remote_resource_overrides, + validate_remote_target_label, validate_remote_tool_spec, validate_ssh_destination, validate_workspace_path, ActiveBuildTarget, + BuildTargetCatalog, RemoteConfigDraft, RemoteOverlayDraft, RemoteProjectOverlayDraft, RemoteQuantity, RemoteResourceOverridesDraft, + RemoteTargetDraft, +}; use crate::vscomm::{self, parse_triggers, Trigger}; mod palette; @@ -96,6 +102,509 @@ enum HostPopup { None, Targets, Help, + Setup, + ConfigError, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum SetupScreen { + List, + TargetForm, + Overrides, + Tools, + ToolForm, + Environment, + Exclusions, + Artifacts, + Resources, + AdvancedResources, + ConfirmDelete, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum TargetFormFocus { + Label, + Ssh, + Workspace, + Overrides, + Resources, + Save, + Cancel, +} + +impl TargetFormFocus { + fn next(self, reverse: bool) -> Self { + let fields = [Self::Label, Self::Ssh, Self::Workspace, Self::Overrides, Self::Resources, Self::Save, Self::Cancel]; + let index = fields.iter().position(|field| *field == self).unwrap_or(0); + let next = if reverse { (index + fields.len() - 1) % fields.len() } else { (index + 1) % fields.len() }; + fields[next] + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum OverrideFocus { + Tools, + Environment, + Exclusions, + Artifacts, + Done, +} + +impl OverrideFocus { + fn next(self, reverse: bool) -> Self { + let fields = [Self::Tools, Self::Environment, Self::Exclusions, Self::Artifacts, Self::Done]; + let index = fields.iter().position(|field| *field == self).unwrap_or(0); + let next = if reverse { (index + fields.len() - 1) % fields.len() } else { (index + 1) % fields.len() }; + fields[next] + } +} + +#[derive(Clone, Debug)] +struct TextField { + value: String, + cursor: usize, +} + +impl TextField { + fn new(value: impl Into) -> Self { + let value = value.into(); + let cursor = value.chars().count(); + Self { value, cursor } + } + + fn handle_key(&mut self, key: KeyEvent) -> bool { + if key.modifiers.contains(KeyModifiers::CONTROL) || key.modifiers.contains(KeyModifiers::ALT) { + return false; + } + match key.code { + KeyCode::Char(character) => { + self.insert(character); + true + } + KeyCode::Backspace => { + if self.cursor > 0 { + let start = self.byte_index(self.cursor - 1); + let end = self.byte_index(self.cursor); + self.value.replace_range(start..end, ""); + self.cursor -= 1; + } + true + } + KeyCode::Delete => { + if self.cursor < self.value.chars().count() { + let start = self.byte_index(self.cursor); + let end = self.byte_index(self.cursor + 1); + self.value.replace_range(start..end, ""); + } + true + } + KeyCode::Left => { + self.cursor = self.cursor.saturating_sub(1); + true + } + KeyCode::Right => { + self.cursor = (self.cursor + 1).min(self.value.chars().count()); + true + } + KeyCode::Home => { + self.cursor = 0; + true + } + KeyCode::End => { + self.cursor = self.value.chars().count(); + true + } + _ => false, + } + } + + fn insert(&mut self, character: char) { + let index = self.byte_index(self.cursor); + self.value.insert(index, character); + self.cursor += 1; + } + + fn byte_index(&self, character_index: usize) -> usize { + self.value.char_indices().nth(character_index).map_or(self.value.len(), |(index, _)| index) + } +} + +#[derive(Clone, Debug)] +struct TargetFormState { + original_label: Option, + draft: RemoteTargetDraft, + label: TextField, + ssh: TextField, + workspace: TextField, + focus: TargetFormFocus, +} + +impl TargetFormState { + fn new(original_label: Option, draft: RemoteTargetDraft) -> Self { + Self { + original_label: original_label.clone(), + label: TextField::new(original_label.clone().unwrap_or_default()), + ssh: TextField::new(draft.ssh.clone()), + workspace: TextField::new(draft.workspace.clone()), + draft, + focus: TargetFormFocus::Label, + } + } + + fn candidate(&self) -> RemoteTargetDraft { + let mut draft = self.draft.clone(); + draft.ssh = self.ssh.value.clone(); + draft.workspace = self.workspace.value.clone(); + draft + } +} + +#[derive(Clone, Debug)] +struct ToolFormState { + original_index: Option, + name: TextField, + command: TextField, + allow_args: bool, + focus: usize, +} + +impl ToolFormState { + fn new(original_index: Option, tool: Option<&RemoteToolSpec>) -> Self { + Self { + original_index, + name: TextField::new(tool.map_or("", |tool| tool.name.as_str())), + command: TextField::new(tool.and_then(|tool| tool.command.as_deref()).unwrap_or_default()), + allow_args: tool.is_some_and(|tool| tool.allow_args), + focus: 0, + } + } + + fn next_focus(&mut self, reverse: bool) { + self.focus = if reverse { (self.focus + 4) % 5 } else { (self.focus + 1) % 5 }; + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum StringListKind { + Environment, + Exclusions, + Artifacts, +} + +impl StringListKind { + fn title(self) -> &'static str { + match self { + Self::Environment => "Remote Environment Names", + Self::Exclusions => "Remote Exclusions", + Self::Artifacts => "Remote Artifacts", + } + } +} + +#[derive(Clone, Debug)] +struct StringListState { + kind: StringListKind, + entries: Vec, + selected: usize, + changed: bool, + editing: Option, + editing_index: Option, + error: Option, +} + +impl StringListState { + fn new(kind: StringListKind, entries: Option>) -> Self { + Self { kind, entries: entries.unwrap_or_default(), selected: 0, changed: false, editing: None, editing_index: None, error: None } + } +} + +#[derive(Clone, Debug)] +struct ToolListState { + entries: Vec, + selected: usize, + changed: bool, + error: Option, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ResourceFieldKind { + Quantity, + Count, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ResourceFieldKey { + ConnectTimeout, + SyncTimeout, + BuildTimeout, + MaxOutput, + IdleOutputTimeout, + CleanupTimeout, + ArtifactTimeout, + MaxArtifactBytes, + MaxArtifactTotalBytes, + MaxArtifactEntries, + MaxWorkerUploads, + MaxWorkerUploadBytes, + MaxWorkerJobs, + MaxWorkerJobBytes, + MaxWorkerArtifactSpools, + MaxWorkerArtifactSpoolBytes, + MaxWorkerStateEntries, + MaxActiveBuilds, +} + +#[derive(Clone, Debug)] +struct ResourceInput { + key: ResourceFieldKey, + label: &'static str, + kind: ResourceFieldKind, + value: TextField, +} + +#[derive(Clone, Debug)] +struct ResourceFormState { + inputs: Vec, + use_defaults: bool, + focus: usize, + advanced: bool, +} + +impl ResourceFormState { + fn new(resources: Option<&RemoteResourceOverridesDraft>) -> Self { + let resources = resources.cloned().unwrap_or_default(); + let inputs = vec![ + resource_input( + ResourceFieldKey::ConnectTimeout, + "Connect timeout", + ResourceFieldKind::Quantity, + quantity_text(resources.connect_timeout.as_ref()), + ), + resource_input( + ResourceFieldKey::SyncTimeout, + "Sync timeout", + ResourceFieldKind::Quantity, + quantity_text(resources.sync_timeout.as_ref()), + ), + resource_input( + ResourceFieldKey::BuildTimeout, + "Build timeout", + ResourceFieldKind::Quantity, + quantity_text(resources.build_timeout.as_ref()), + ), + resource_input( + ResourceFieldKey::MaxOutput, + "Max output bytes", + ResourceFieldKind::Quantity, + quantity_text(resources.max_output.as_ref()), + ), + resource_input( + ResourceFieldKey::IdleOutputTimeout, + "Idle output timeout", + ResourceFieldKind::Quantity, + quantity_text(resources.idle_output_timeout.as_ref()), + ), + resource_input( + ResourceFieldKey::CleanupTimeout, + "Cleanup timeout", + ResourceFieldKind::Quantity, + quantity_text(resources.cleanup_timeout.as_ref()), + ), + resource_input( + ResourceFieldKey::ArtifactTimeout, + "Artifact timeout", + ResourceFieldKind::Quantity, + quantity_text(resources.artifact_timeout.as_ref()), + ), + resource_input( + ResourceFieldKey::MaxArtifactBytes, + "Max artifact bytes", + ResourceFieldKind::Quantity, + quantity_text(resources.max_artifact_bytes.as_ref()), + ), + resource_input( + ResourceFieldKey::MaxArtifactTotalBytes, + "Max artifact total bytes", + ResourceFieldKind::Quantity, + quantity_text(resources.max_artifact_total_bytes.as_ref()), + ), + resource_input( + ResourceFieldKey::MaxArtifactEntries, + "Max artifact entries", + ResourceFieldKind::Count, + count_text(resources.max_artifact_entries), + ), + resource_input( + ResourceFieldKey::MaxWorkerUploads, + "Max worker uploads", + ResourceFieldKind::Count, + count_text(resources.max_worker_uploads), + ), + resource_input( + ResourceFieldKey::MaxWorkerUploadBytes, + "Max worker upload bytes", + ResourceFieldKind::Quantity, + quantity_text(resources.max_worker_upload_bytes.as_ref()), + ), + resource_input(ResourceFieldKey::MaxWorkerJobs, "Max worker jobs", ResourceFieldKind::Count, count_text(resources.max_worker_jobs)), + resource_input( + ResourceFieldKey::MaxWorkerJobBytes, + "Max worker job bytes", + ResourceFieldKind::Quantity, + quantity_text(resources.max_worker_job_bytes.as_ref()), + ), + resource_input( + ResourceFieldKey::MaxWorkerArtifactSpools, + "Max worker artifact spools", + ResourceFieldKind::Count, + count_text(resources.max_worker_artifact_spools), + ), + resource_input( + ResourceFieldKey::MaxWorkerArtifactSpoolBytes, + "Max worker artifact spool bytes", + ResourceFieldKind::Quantity, + quantity_text(resources.max_worker_artifact_spool_bytes.as_ref()), + ), + resource_input( + ResourceFieldKey::MaxWorkerStateEntries, + "Max worker state entries", + ResourceFieldKind::Count, + count_text(resources.max_worker_state_entries), + ), + resource_input(ResourceFieldKey::MaxActiveBuilds, "Max active builds", ResourceFieldKind::Count, count_text(resources.max_active_builds)), + ]; + Self { inputs, use_defaults: false, focus: 0, advanced: false } + } + + fn visible_indices(&self) -> Vec { + self.inputs + .iter() + .enumerate() + .filter_map(|(index, input)| { + let advanced = matches!( + input.key, + ResourceFieldKey::ArtifactTimeout + | ResourceFieldKey::MaxArtifactBytes + | ResourceFieldKey::MaxArtifactTotalBytes + | ResourceFieldKey::MaxArtifactEntries + | ResourceFieldKey::MaxWorkerUploads + | ResourceFieldKey::MaxWorkerUploadBytes + | ResourceFieldKey::MaxWorkerJobs + | ResourceFieldKey::MaxWorkerJobBytes + | ResourceFieldKey::MaxWorkerArtifactSpools + | ResourceFieldKey::MaxWorkerArtifactSpoolBytes + | ResourceFieldKey::MaxWorkerStateEntries + ); + (advanced == self.advanced).then_some(index) + }) + .collect() + } + + fn action_index(&self, action: usize) -> usize { + self.visible_indices().len() + action + } + + fn draft(&self) -> Result, String> { + let mut resources = RemoteResourceOverridesDraft::default(); + for input in &self.inputs { + let value = input.value.value.trim(); + if value.is_empty() { + continue; + } + let _field_kind = input.kind; + let count = || value.parse::().map_err(|_| format!("{} must be an integer", input.label)); + match input.key { + ResourceFieldKey::ConnectTimeout => resources.connect_timeout = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::SyncTimeout => resources.sync_timeout = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::BuildTimeout => resources.build_timeout = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::MaxOutput => resources.max_output = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::IdleOutputTimeout => resources.idle_output_timeout = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::CleanupTimeout => resources.cleanup_timeout = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::ArtifactTimeout => resources.artifact_timeout = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::MaxArtifactBytes => resources.max_artifact_bytes = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::MaxArtifactTotalBytes => resources.max_artifact_total_bytes = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::MaxArtifactEntries => resources.max_artifact_entries = Some(count()?), + ResourceFieldKey::MaxWorkerUploads => resources.max_worker_uploads = Some(count()?), + ResourceFieldKey::MaxWorkerUploadBytes => resources.max_worker_upload_bytes = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::MaxWorkerJobs => resources.max_worker_jobs = Some(count()?), + ResourceFieldKey::MaxWorkerJobBytes => resources.max_worker_job_bytes = Some(RemoteQuantity::Text(value.to_string())), + ResourceFieldKey::MaxWorkerArtifactSpools => resources.max_worker_artifact_spools = Some(count()?), + ResourceFieldKey::MaxWorkerArtifactSpoolBytes => { + resources.max_worker_artifact_spool_bytes = Some(RemoteQuantity::Text(value.to_string())) + } + ResourceFieldKey::MaxWorkerStateEntries => resources.max_worker_state_entries = Some(count()?), + ResourceFieldKey::MaxActiveBuilds => resources.max_active_builds = Some(count()?), + } + } + Ok((!resources.is_empty()).then_some(resources)) + } + + fn clear(&mut self) { + for input in &mut self.inputs { + input.value = TextField::new(""); + } + self.use_defaults = true; + self.focus = 0; + } +} + +fn resource_input(key: ResourceFieldKey, label: &'static str, kind: ResourceFieldKind, value: String) -> ResourceInput { + ResourceInput { key, label, kind, value: TextField::new(value) } +} + +fn quantity_text(value: Option<&RemoteQuantity>) -> String { + match value { + Some(RemoteQuantity::Integer(value)) => value.to_string(), + Some(RemoteQuantity::Text(value)) => value.clone(), + None => String::new(), + } +} + +fn count_text(value: Option) -> String { + value.map_or_else(String::new, |value| value.to_string()) +} + +#[derive(Clone, Debug)] +struct RemoteSetupState { + draft: RemoteConfigDraft, + selected: usize, + screen: SetupScreen, + override_focus: OverrideFocus, + target_form: Option, + tool_form: Option, + tools: Option, + strings: Option, + resources: Option, + delete_label: Option, + error: Option, +} + +impl RemoteSetupState { + fn new(draft: RemoteConfigDraft) -> Self { + Self { + draft, + selected: 0, + screen: SetupScreen::List, + override_focus: OverrideFocus::Tools, + target_form: None, + tool_form: None, + tools: None, + strings: None, + resources: None, + delete_label: None, + error: None, + } + } + + fn labels(&self) -> Vec { + self.draft.targets.keys().cloned().collect() + } +} + +#[derive(Clone, Debug)] +struct ConfigErrorState { + message: String, + view: bool, } struct HostUiState { @@ -104,6 +613,8 @@ struct HostUiState { confirmation: Option<(String, Instant)>, free_bytes: Option, last_free_refresh: Instant, + setup: Option, + config_error: Option, } impl HostUiState { @@ -114,6 +625,8 @@ impl HostUiState { confirmation: None, free_bytes: None, last_free_refresh: Instant::now() - Duration::from_secs(10), + setup: None, + config_error: None, } } @@ -127,32 +640,22 @@ impl HostUiState { fn handle_key(&mut self, key: KeyEvent, catalog: &BuildTargetCatalog, active: &ActiveBuildTarget) -> bool { match self.popup { - HostPopup::Targets => { - match key.code { - KeyCode::Up => { - self.target_index = self.target_index.saturating_sub(1); - } - KeyCode::Down => { - self.target_index = (self.target_index + 1).min(catalog.summaries().len().saturating_sub(1)); - } - KeyCode::Enter => { - if let Some(target) = catalog.summaries().get(self.target_index) { - if active.select(catalog, target.label()).is_ok() { - self.confirmation = Some((format!("Build target: {}", target.label()), Instant::now())); - } - } - self.popup = HostPopup::None; - } - KeyCode::Esc => { - self.popup = HostPopup::None; - } - _ => {} + HostPopup::Targets => return self.handle_target_popup_key(key, catalog, active), + HostPopup::Help => { + if key.code == KeyCode::Esc { + self.popup = HostPopup::None; } return true; } - HostPopup::Help => { - if key.code == KeyCode::Esc { + HostPopup::Setup => return self.handle_setup_key(key, catalog), + HostPopup::ConfigError => { + if key.code == KeyCode::Char('v') && !key.modifiers.intersects(KeyModifiers::CONTROL | KeyModifiers::ALT) { + if let Some(error) = &mut self.config_error { + error.view = true; + } + } else if key.code == KeyCode::Esc { self.popup = HostPopup::None; + self.config_error = None; } return true; } @@ -169,8 +672,771 @@ impl HostUiState { self.popup = HostPopup::Help; return true; } + if key.code == KeyCode::Char('s') && key.modifiers == (KeyModifiers::CONTROL | KeyModifiers::ALT) { + self.open_setup(catalog); + return true; + } false } + + fn handle_target_popup_key(&mut self, key: KeyEvent, catalog: &BuildTargetCatalog, active: &ActiveBuildTarget) -> bool { + match key.code { + KeyCode::Up => self.target_index = self.target_index.saturating_sub(1), + KeyCode::Down => self.target_index = (self.target_index + 1).min(catalog.summaries().len().saturating_sub(1)), + KeyCode::Enter => { + if let Some(target) = catalog.summaries().get(self.target_index) { + if active.select(catalog, target.label()).is_ok() { + self.confirmation = Some((format!("Build target: {}", target.label()), Instant::now())); + } + } + self.popup = HostPopup::None; + } + KeyCode::Esc => self.popup = HostPopup::None, + _ => {} + } + true + } + + fn open_setup(&mut self, catalog: &BuildTargetCatalog) { + match RemoteConfigDraft::load_optional(catalog.project_root(), catalog.base_project()) { + Ok(Some(draft)) => { + self.setup = Some(RemoteSetupState::new(draft)); + self.config_error = None; + self.popup = HostPopup::Setup; + } + Ok(None) => { + self.setup = Some(RemoteSetupState::new(RemoteConfigDraft::default())); + self.config_error = None; + self.popup = HostPopup::Setup; + } + Err(error) => { + self.setup = None; + self.config_error = Some(ConfigErrorState { message: error, view: false }); + self.popup = HostPopup::ConfigError; + } + } + } + + fn handle_setup_key(&mut self, key: KeyEvent, catalog: &BuildTargetCatalog) -> bool { + let screen = self.setup.as_ref().map_or(SetupScreen::List, |setup| setup.screen); + match screen { + SetupScreen::List => self.handle_setup_list_key(key, catalog), + SetupScreen::TargetForm => self.handle_target_form_key(key, catalog), + SetupScreen::Overrides => self.handle_overrides_key(key, catalog), + SetupScreen::Tools => self.handle_tools_key(key), + SetupScreen::ToolForm => self.handle_tool_form_key(key), + SetupScreen::Environment | SetupScreen::Exclusions | SetupScreen::Artifacts => self.handle_string_list_key(key), + SetupScreen::Resources | SetupScreen::AdvancedResources => self.handle_resources_key(key, catalog), + SetupScreen::ConfirmDelete => self.handle_delete_key(key), + } + } + + fn handle_setup_list_key(&mut self, key: KeyEvent, catalog: &BuildTargetCatalog) -> bool { + match key.code { + KeyCode::Up => { + if let Some(setup) = &mut self.setup { + setup.selected = setup.selected.saturating_sub(1); + } + } + KeyCode::Down => { + if let Some(setup) = &mut self.setup { + setup.selected = (setup.selected + 1).min(setup.draft.targets.len().saturating_sub(1)); + } + } + KeyCode::Char('a') | KeyCode::Char('A') => self.begin_target_form(None), + KeyCode::Char('e') | KeyCode::Char('E') | KeyCode::Enter => { + let label = self.setup.as_ref().and_then(|setup| setup.labels().get(setup.selected).cloned()); + if let Some(label) = label { + self.begin_target_form(Some(label)); + } + } + KeyCode::Char('d') | KeyCode::Char('D') => { + let label = self.setup.as_ref().and_then(|setup| setup.labels().get(setup.selected).cloned()); + if let Some(setup) = &mut self.setup { + setup.delete_label = label; + if setup.delete_label.is_some() { + setup.screen = SetupScreen::ConfirmDelete; + } + } + } + KeyCode::Char('s') | KeyCode::Char('S') => self.save_setup(catalog), + KeyCode::Esc => { + self.setup = None; + self.popup = HostPopup::None; + } + _ => {} + } + true + } + + fn begin_target_form(&mut self, label: Option) { + let draft = label.as_ref().and_then(|label| self.setup.as_ref()?.draft.targets.get(label).cloned()).unwrap_or_else(|| RemoteTargetDraft { + ssh: String::new(), + workspace: String::new(), + project: None, + resources: None, + }); + if let Some(setup) = &mut self.setup { + setup.error = None; + setup.target_form = Some(TargetFormState::new(label, draft)); + setup.screen = SetupScreen::TargetForm; + } + } + + fn handle_target_form_key(&mut self, key: KeyEvent, catalog: &BuildTargetCatalog) -> bool { + let consumed = if let Some(setup) = &mut self.setup { + if let Some(form) = &mut setup.target_form { + match form.focus { + TargetFormFocus::Label => form.label.handle_key(key), + TargetFormFocus::Ssh => form.ssh.handle_key(key), + TargetFormFocus::Workspace => form.workspace.handle_key(key), + _ => false, + } + } else { + false + } + } else { + false + }; + if consumed { + return true; + } + + if key.code == KeyCode::Tab { + if let Some(form) = self.setup.as_mut().and_then(|setup| setup.target_form.as_mut()) { + form.focus = form.focus.next(key.modifiers.contains(KeyModifiers::SHIFT)); + } + return true; + } + match key.code { + KeyCode::Enter => { + let focus = self.setup.as_ref().and_then(|setup| setup.target_form.as_ref()).map_or(TargetFormFocus::Cancel, |form| form.focus); + match focus { + TargetFormFocus::Overrides => self.open_overrides(), + TargetFormFocus::Resources => self.open_resources(), + TargetFormFocus::Save => self.commit_target_form(catalog), + TargetFormFocus::Cancel => self.cancel_target_form(), + _ => { + if let Some(form) = self.setup.as_mut().and_then(|setup| setup.target_form.as_mut()) { + form.focus = form.focus.next(false); + } + } + } + } + KeyCode::Esc => self.cancel_target_form(), + _ => {} + } + true + } + + fn open_overrides(&mut self) { + if let Some(setup) = &mut self.setup { + setup.error = None; + setup.override_focus = OverrideFocus::Tools; + setup.screen = SetupScreen::Overrides; + } + } + + fn open_resources(&mut self) { + let resources = + self.setup.as_ref().and_then(|setup| setup.target_form.as_ref()).map(|form| ResourceFormState::new(form.draft.resources.as_ref())); + if let Some(setup) = &mut self.setup { + setup.resources = resources; + setup.error = None; + setup.screen = SetupScreen::Resources; + } + } + + fn cancel_target_form(&mut self) { + if let Some(setup) = &mut self.setup { + setup.target_form = None; + setup.tools = None; + setup.tool_form = None; + setup.strings = None; + setup.resources = None; + setup.error = None; + setup.screen = SetupScreen::List; + } + } + + fn commit_target_form(&mut self, catalog: &BuildTargetCatalog) { + let Some(setup) = &mut self.setup else { return }; + let Some(form) = setup.target_form.take() else { return }; + let candidate_label = form.label.value.clone(); + let candidate = normalize_target_draft(form.candidate()); + let result = (|| { + validate_remote_target_label(&candidate_label)?; + validate_ssh_destination(&candidate.ssh)?; + validate_workspace_path(&candidate.workspace)?; + if form.original_label.as_deref() != Some(candidate_label.as_str()) && setup.draft.targets.contains_key(&candidate_label) { + return Err(format!("duplicate target label: {candidate_label}")); + } + let mut draft = setup.draft.clone(); + if let Some(original) = &form.original_label { + draft.targets.remove(original); + } + draft.targets.insert(candidate_label.clone(), candidate); + draft.validate(catalog.base_project())?; + Ok::(draft) + })(); + + match result { + Ok(draft) => { + setup.draft = draft; + setup.selected = setup.draft.targets.keys().position(|label| label == &candidate_label).unwrap_or(0); + setup.error = None; + setup.screen = SetupScreen::List; + } + Err(error) => { + setup.error = Some(error); + setup.target_form = Some(form); + } + } + } + + fn handle_overrides_key(&mut self, key: KeyEvent, _catalog: &BuildTargetCatalog) -> bool { + let focus = self.setup.as_ref().map_or(OverrideFocus::Done, |setup| setup.override_focus); + if key.code == KeyCode::Tab { + if let Some(setup) = &mut self.setup { + setup.override_focus = setup.override_focus.next(key.modifiers.contains(KeyModifiers::SHIFT)); + } + return true; + } + match key.code { + KeyCode::Up => { + if let Some(setup) = &mut self.setup { + setup.override_focus = focus.next(true); + } + } + KeyCode::Down => { + if let Some(setup) = &mut self.setup { + setup.override_focus = focus.next(false); + } + } + KeyCode::Enter => match focus { + OverrideFocus::Tools => self.open_tools(), + OverrideFocus::Environment => self.open_string_list(StringListKind::Environment), + OverrideFocus::Exclusions => self.open_string_list(StringListKind::Exclusions), + OverrideFocus::Artifacts => self.open_string_list(StringListKind::Artifacts), + OverrideFocus::Done => { + if let Some(setup) = &mut self.setup { + setup.screen = SetupScreen::TargetForm; + } + } + }, + KeyCode::Esc => { + if let Some(setup) = &mut self.setup { + setup.screen = SetupScreen::TargetForm; + setup.error = None; + } + } + _ => {} + } + true + } + + fn open_tools(&mut self) { + let entries = self + .setup + .as_ref() + .and_then(|setup| setup.target_form.as_ref()) + .and_then(|form| form.draft.project.as_ref()) + .and_then(|project| project.remote.as_ref()) + .and_then(|remote| remote.tools.clone()) + .unwrap_or_default(); + if let Some(setup) = &mut self.setup { + setup.tools = Some(ToolListState { entries, selected: 0, changed: false, error: None }); + setup.error = None; + setup.screen = SetupScreen::Tools; + } + } + + fn open_string_list(&mut self, kind: StringListKind) { + let entries = self + .setup + .as_ref() + .and_then(|setup| setup.target_form.as_ref()) + .and_then(|form| form.draft.project.as_ref()) + .and_then(|project| project.remote.as_ref()) + .and_then(|remote| match kind { + StringListKind::Environment => remote.environment.clone(), + StringListKind::Exclusions => remote.exclude.clone(), + StringListKind::Artifacts => remote.artifacts.clone(), + }); + if let Some(setup) = &mut self.setup { + setup.strings = Some(StringListState::new(kind, entries)); + setup.error = None; + setup.screen = match kind { + StringListKind::Environment => SetupScreen::Environment, + StringListKind::Exclusions => SetupScreen::Exclusions, + StringListKind::Artifacts => SetupScreen::Artifacts, + }; + } + } + + fn save_setup(&mut self, catalog: &BuildTargetCatalog) { + let Some(draft) = self.setup.as_ref().map(|setup| setup.draft.clone()) else { return }; + match draft.write_atomic(catalog.project_root(), catalog.base_project()) { + Ok(()) => { + self.confirmation = Some(("Saved .bunkerbox/remote.conf; changes apply next run".to_string(), Instant::now())); + self.setup = None; + self.popup = HostPopup::None; + } + Err(error) => { + if let Some(setup) = &mut self.setup { + setup.error = Some(error); + } + } + } + } + + fn handle_tools_key(&mut self, key: KeyEvent) -> bool { + match key.code { + KeyCode::Up => { + if let Some(tools) = self.setup.as_mut().and_then(|setup| setup.tools.as_mut()) { + tools.selected = tools.selected.saturating_sub(1); + } + } + KeyCode::Down => { + if let Some(tools) = self.setup.as_mut().and_then(|setup| setup.tools.as_mut()) { + tools.selected = (tools.selected + 1).min(tools.entries.len().saturating_sub(1)); + } + } + KeyCode::Char('a') | KeyCode::Char('A') => self.begin_tool_form(None), + KeyCode::Char('e') | KeyCode::Char('E') | KeyCode::Enter => { + let selected = self.setup.as_ref().and_then(|setup| setup.tools.as_ref()).map(|tools| tools.selected); + if let Some(index) = selected { + self.begin_tool_form(Some(index)); + } + } + KeyCode::Char('d') | KeyCode::Char('D') => { + if let Some(setup) = &mut self.setup { + if let Some(tools) = &mut setup.tools { + if tools.selected < tools.entries.len() { + tools.entries.remove(tools.selected); + tools.selected = tools.selected.min(tools.entries.len().saturating_sub(1)); + tools.changed = true; + tools.error = None; + } + } + } + } + KeyCode::Esc | KeyCode::Char('q') | KeyCode::Char('Q') => self.finish_tools(), + _ => {} + } + true + } + + fn begin_tool_form(&mut self, index: Option) { + let tool = index.and_then(|index| self.setup.as_ref()?.tools.as_ref()?.entries.get(index).cloned()); + if let Some(setup) = &mut self.setup { + setup.tool_form = Some(ToolFormState::new(index, tool.as_ref())); + setup.error = None; + setup.screen = SetupScreen::ToolForm; + } + } + + fn handle_tool_form_key(&mut self, key: KeyEvent) -> bool { + let focus = self.setup.as_ref().and_then(|setup| setup.tool_form.as_ref()).map_or(4, |form| form.focus); + let consumed = if let Some(form) = self.setup.as_mut().and_then(|setup| setup.tool_form.as_mut()) { + match form.focus { + 0 => form.name.handle_key(key), + 1 => form.command.handle_key(key), + _ => false, + } + } else { + false + }; + if consumed { + return true; + } + if key.code == KeyCode::Tab { + if let Some(form) = self.setup.as_mut().and_then(|setup| setup.tool_form.as_mut()) { + form.next_focus(key.modifiers.contains(KeyModifiers::SHIFT)); + } + return true; + } + match key.code { + KeyCode::Char(' ') if focus == 2 => { + if let Some(form) = self.setup.as_mut().and_then(|setup| setup.tool_form.as_mut()) { + form.allow_args = !form.allow_args; + } + } + KeyCode::Enter if focus == 3 => self.commit_tool_form(), + KeyCode::Enter if focus == 4 => self.cancel_tool_form(), + KeyCode::Enter => { + if let Some(form) = self.setup.as_mut().and_then(|setup| setup.tool_form.as_mut()) { + form.next_focus(false); + } + } + KeyCode::Esc => self.cancel_tool_form(), + _ => {} + } + true + } + + fn commit_tool_form(&mut self) { + let Some(setup) = &mut self.setup else { return }; + let Some(form) = setup.tool_form.take() else { return }; + let name = form.name.value.trim().to_string(); + let command = match form.command.value.trim() { + "" => None, + value => Some(value.to_string()), + }; + let candidate = RemoteToolSpec { name, command, allow_args: form.allow_args }; + let result = (|| { + validate_remote_tool_spec(&candidate)?; + let tools = setup.tools.as_ref().ok_or_else(|| "tool list is unavailable".to_string())?; + if tools.entries.iter().enumerate().any(|(index, tool)| Some(index) != form.original_index && tool.name == candidate.name) { + return Err(format!("duplicate remote tool: {}", candidate.name)); + } + Ok::<(), String>(()) + })(); + match result { + Ok(()) => { + if let Some(tools) = &mut setup.tools { + if let Some(index) = form.original_index { + if index < tools.entries.len() { + tools.entries[index] = candidate; + tools.selected = index; + } + } else { + tools.entries.push(candidate); + tools.selected = tools.entries.len().saturating_sub(1); + } + tools.changed = true; + tools.error = None; + } + setup.screen = SetupScreen::Tools; + setup.error = None; + } + Err(error) => { + setup.error = Some(error.clone()); + let mut restored = form; + restored.name = TextField::new(candidate.name); + restored.command = TextField::new(candidate.command.unwrap_or_default()); + setup.tool_form = Some(restored); + } + } + } + + fn cancel_tool_form(&mut self) { + if let Some(setup) = &mut self.setup { + setup.tool_form = None; + setup.error = None; + setup.screen = SetupScreen::Tools; + } + } + + fn finish_tools(&mut self) { + let Some(setup) = &mut self.setup else { return }; + let Some(tools) = setup.tools.take() else { return }; + if tools.changed { + let mut names = std::collections::BTreeSet::new(); + let result = tools.entries.iter().try_for_each(|tool| { + validate_remote_tool_spec(tool)?; + if !names.insert(tool.name.clone()) { + return Err(format!("duplicate remote tool: {}", tool.name)); + } + Ok::<(), String>(()) + }); + if let Err(error) = result { + setup.error = Some(error); + setup.tools = Some(tools); + return; + } + if let Some(form) = &mut setup.target_form { + remote_overlay_mut(&mut form.draft).tools = Some(tools.entries); + } + } + setup.tool_form = None; + setup.error = None; + setup.screen = SetupScreen::Overrides; + } + + fn handle_string_list_key(&mut self, key: KeyEvent) -> bool { + let editing = self.setup.as_ref().and_then(|setup| setup.strings.as_ref()).is_some_and(|state| state.editing.is_some()); + if editing { + let consumed = if let Some(state) = self.setup.as_mut().and_then(|setup| setup.strings.as_mut()) { + state.editing.as_mut().is_some_and(|field| field.handle_key(key)) + } else { + false + }; + if consumed { + return true; + } + match key.code { + KeyCode::Enter => self.commit_string_edit(), + KeyCode::Esc => { + if let Some(state) = self.setup.as_mut().and_then(|setup| setup.strings.as_mut()) { + state.editing = None; + state.editing_index = None; + } + } + _ => {} + } + return true; + } + + match key.code { + KeyCode::Up => { + if let Some(state) = self.setup.as_mut().and_then(|setup| setup.strings.as_mut()) { + state.selected = state.selected.saturating_sub(1); + } + } + KeyCode::Down => { + if let Some(state) = self.setup.as_mut().and_then(|setup| setup.strings.as_mut()) { + state.selected = (state.selected + 1).min(state.entries.len().saturating_sub(1)); + } + } + KeyCode::Char('a') | KeyCode::Char('A') => self.begin_string_edit(None), + KeyCode::Char('e') | KeyCode::Char('E') | KeyCode::Enter => { + let selected = self.setup.as_ref().and_then(|setup| setup.strings.as_ref()).map(|state| state.selected); + if let Some(index) = selected { + self.begin_string_edit(Some(index)); + } + } + KeyCode::Char('d') | KeyCode::Char('D') => { + if let Some(state) = self.setup.as_mut().and_then(|setup| setup.strings.as_mut()) { + if state.selected < state.entries.len() { + state.entries.remove(state.selected); + state.selected = state.selected.min(state.entries.len().saturating_sub(1)); + state.changed = true; + } + } + } + KeyCode::Esc | KeyCode::Char('q') | KeyCode::Char('Q') => self.finish_string_list(), + _ => {} + } + true + } + + fn begin_string_edit(&mut self, index: Option) { + let value = index.and_then(|index| self.setup.as_ref()?.strings.as_ref()?.entries.get(index).cloned()).unwrap_or_default(); + if let Some(state) = self.setup.as_mut().and_then(|setup| setup.strings.as_mut()) { + state.editing = Some(TextField::new(value)); + state.editing_index = index; + state.error = None; + } + } + + fn commit_string_edit(&mut self) { + let Some(setup) = &mut self.setup else { return }; + let Some(state) = &mut setup.strings else { return }; + let Some(editing) = state.editing.take() else { return }; + let value = editing.value.trim().to_string(); + if value.is_empty() { + state.error = Some("value must not be empty".to_string()); + state.editing = Some(TextField::new(value)); + return; + } + let mut candidate = state.entries.clone(); + if let Some(index) = state.editing_index { + if index < candidate.len() { + candidate[index] = value; + } + } else { + candidate.push(value); + } + let result = validate_string_entries(state.kind, &candidate); + if let Err(error) = result { + state.error = Some(error); + state.editing = Some(editing); + return; + } + state.entries = candidate; + state.selected = state.editing_index.unwrap_or_else(|| state.entries.len().saturating_sub(1)); + state.editing_index = None; + state.changed = true; + state.error = None; + } + + fn finish_string_list(&mut self) { + let Some(setup) = &mut self.setup else { return }; + let Some(state) = setup.strings.take() else { return }; + if state.changed { + if let Err(error) = validate_string_entries(state.kind, &state.entries) { + setup.error = Some(error); + setup.strings = Some(state); + return; + } + if let Some(form) = &mut setup.target_form { + let remote = remote_overlay_mut(&mut form.draft); + match state.kind { + StringListKind::Environment => remote.environment = Some(state.entries), + StringListKind::Exclusions => remote.exclude = Some(state.entries), + StringListKind::Artifacts => remote.artifacts = Some(state.entries), + } + } + } + setup.error = None; + setup.screen = SetupScreen::Overrides; + } + + fn handle_resources_key(&mut self, key: KeyEvent, _catalog: &BuildTargetCatalog) -> bool { + let consumed = if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + resources + .visible_indices() + .get(resources.focus) + .and_then(|index| resources.inputs.get_mut(*index)) + .is_some_and(|input| input.value.handle_key(key)) + } else { + false + }; + if consumed { + return true; + } + if key.code == KeyCode::Tab { + if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + let count = resources.visible_indices().len() + 3; + resources.focus = + if key.modifiers.contains(KeyModifiers::SHIFT) { (resources.focus + count - 1) % count } else { (resources.focus + 1) % count }; + } + return true; + } + match key.code { + KeyCode::Up => { + if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + let count = resources.visible_indices().len() + 3; + resources.focus = (resources.focus + count - 1) % count; + } + } + KeyCode::Down => { + if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + let count = resources.visible_indices().len() + 3; + resources.focus = (resources.focus + 1) % count; + } + } + KeyCode::PageUp | KeyCode::PageDown | KeyCode::Char('a') | KeyCode::Char('A') => { + let advanced = if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + resources.advanced = !resources.advanced; + resources.focus = 0; + resources.advanced + } else { + false + }; + if let Some(setup) = &mut self.setup { + setup.screen = if advanced { SetupScreen::AdvancedResources } else { SetupScreen::Resources }; + } + } + KeyCode::Enter => { + let action = self.setup.as_ref().and_then(|setup| setup.resources.as_ref()).map(|resources| { + let fields = resources.visible_indices().len(); + match resources.focus { + focus if focus < fields => 0, + focus if focus == resources.action_index(0) => 1, + focus if focus == resources.action_index(1) => 2, + _ => 3, + } + }); + match action { + Some(1) => { + if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + resources.clear(); + } + } + Some(2) => self.commit_resources(), + Some(3) => self.cancel_resources(), + _ => { + if let Some(resources) = self.setup.as_mut().and_then(|setup| setup.resources.as_mut()) { + let count = resources.visible_indices().len() + 3; + resources.focus = (resources.focus + 1) % count; + } + } + } + } + KeyCode::Esc => self.cancel_resources(), + _ => {} + } + true + } + + fn commit_resources(&mut self) { + let Some(setup) = &mut self.setup else { return }; + let Some(resources) = setup.resources.take() else { return }; + let draft = match resources.draft() { + Ok(draft) => draft, + Err(error) => { + setup.error = Some(error); + setup.resources = Some(resources); + return; + } + }; + let result = draft.as_ref().map_or(Ok(()), |resources| validate_remote_resource_overrides(resources).map(|_| ())); + match result { + Ok(()) => { + if let Some(form) = &mut setup.target_form { + form.draft.resources = draft; + } + setup.error = None; + setup.screen = SetupScreen::TargetForm; + } + Err(error) => { + setup.error = Some(error); + setup.resources = Some(resources); + } + } + } + + fn cancel_resources(&mut self) { + if let Some(setup) = &mut self.setup { + setup.resources = None; + setup.error = None; + setup.screen = SetupScreen::TargetForm; + } + } + + fn handle_delete_key(&mut self, key: KeyEvent) -> bool { + match key.code { + KeyCode::Enter => { + if let Some(setup) = &mut self.setup { + if let Some(label) = setup.delete_label.take() { + setup.draft.targets.remove(&label); + setup.selected = setup.selected.min(setup.draft.targets.len().saturating_sub(1)); + } + setup.screen = SetupScreen::List; + } + } + KeyCode::Esc => { + if let Some(setup) = &mut self.setup { + setup.delete_label = None; + setup.screen = SetupScreen::List; + } + } + _ => {} + } + true + } +} + +fn validate_string_entries(kind: StringListKind, entries: &[String]) -> Result<(), String> { + match kind { + StringListKind::Environment => validate_remote_environment_names(entries), + StringListKind::Exclusions => validate_remote_exclusion_entries(entries), + StringListKind::Artifacts => validate_remote_artifact_paths(entries), + } +} + +fn remote_overlay_mut(target: &mut RemoteTargetDraft) -> &mut RemoteOverlayDraft { + target.project.get_or_insert_with(RemoteProjectOverlayDraft::default).remote.get_or_insert_with(RemoteOverlayDraft::default) +} + +fn normalize_target_draft(mut target: RemoteTargetDraft) -> RemoteTargetDraft { + if target.project.as_ref().is_some_and(|project| { + project + .remote + .as_ref() + .is_some_and(|remote| remote.exclude.is_none() && remote.environment.is_none() && remote.tools.is_none() && remote.artifacts.is_none()) + }) { + target.project = None; + } + if target.project.as_ref().is_some_and(|project| project.remote.is_none()) { + target.project = None; + } + if target.resources.as_ref().is_some_and(RemoteResourceOverridesDraft::is_empty) { + target.resources = None; + } + target } pub fn show_host_error(overlay: &Arc>, title: &str, message: &str) { @@ -1086,7 +2352,13 @@ fn render_status_bar(area: Rect, buf: &mut Buffer, host: &HostUiState, active: & .as_ref() .filter(|(_, shown_at)| shown_at.elapsed() < Duration::from_secs(3)) .map_or_else(|| format!("Target: {target}"), |(message, _)| format!("Target: {target} ({message})")); - let mut segments = vec![target_segment, "Ctrl+Alt+B Targets".to_string(), "Ctrl+Alt+H Help".to_string(), format!("Local workspace free: {free}")]; + let mut segments = vec![ + target_segment, + "Ctrl+Alt+B Targets".to_string(), + "Ctrl+Alt+S Setup".to_string(), + "Ctrl+Alt+H Help".to_string(), + format!("Local workspace free: {free}"), + ]; while segments.len() > 1 { let text = format!(" {}", segments.join(" | ")); if text.chars().count() <= usize::from(area.width) { @@ -1105,16 +2377,18 @@ fn render_host_popup(area: Rect, buf: &mut Buffer, host: &HostUiState, catalog: let (title, lines, height) = match host.popup { HostPopup::None => return, HostPopup::Help => ( - "Bunkerbox Help", + "Bunkerbox Help".to_string(), vec![ Line::from(Span::styled("Ctrl-Alt-B", Style::default().fg(palette::ACCENT))), Line::from("Select Build Target"), + Line::from(Span::styled("Ctrl-Alt-S", Style::default().fg(palette::ACCENT))), + Line::from("Remote Setup (next run)"), Line::from(Span::styled("Ctrl-Alt-H", Style::default().fg(palette::ACCENT))), Line::from("Show this help"), Line::from(Span::styled("Esc", Style::default().fg(palette::ACCENT))), Line::from("Close host popup"), ], - 10, + 12, ), HostPopup::Targets => { let mut lines = Vec::new(); @@ -1127,8 +2401,10 @@ fn render_host_popup(area: Rect, buf: &mut Buffer, host: &HostUiState, catalog: ])); } let height = (lines.len() as u16 + 4).max(5); - ("Build Targets", lines, height) + ("Build Targets".to_string(), lines, height) } + HostPopup::Setup => render_setup_popup(host), + HostPopup::ConfigError => render_config_error_popup(host), }; let popup_area = centered_rect(area, 64, height); if popup_area.width < 4 || popup_area.height < 3 { @@ -1145,6 +2421,189 @@ fn render_host_popup(area: Rect, buf: &mut Buffer, host: &HostUiState, catalog: Paragraph::new(lines).block(block).render(popup_area, buf); } +fn render_config_error_popup(host: &HostUiState) -> (String, Vec>, u16) { + let Some(error) = &host.config_error else { + return ("Remote Configuration Error".to_string(), vec![Line::from("No configuration error available")], 7); + }; + if error.view { + ( + "Remote Configuration Error".to_string(), + vec![Line::from("The existing remote.conf was not changed."), Line::from(error.message.clone()), Line::from("Esc Cancel")], + 9, + ) + } else { + ( + "Remote Configuration Error".to_string(), + vec![Line::from("Remote Setup cannot edit this file."), Line::from("Press V to view the error."), Line::from("Esc Cancel")], + 8, + ) + } +} + +fn render_setup_popup(host: &HostUiState) -> (String, Vec>, u16) { + let Some(setup) = &host.setup else { + return ("Remote Setup".to_string(), vec![Line::from("No setup state available")], 7); + }; + let mut lines = Vec::new(); + match setup.screen { + SetupScreen::List => { + lines.push(Line::from("A Add Enter/E Edit D Delete S Save Esc Cancel")); + lines.push(Line::from("Targets are applied on the next Bunkerbox run.")); + for (index, label) in setup.labels().iter().enumerate() { + let marker = if index == setup.selected { "> " } else { " " }; + lines.push(Line::from(vec![ + Span::styled(marker, Style::default().fg(palette::ACCENT)), + Span::styled(label.clone(), Style::default().fg(palette::FG)), + ])); + } + if setup.draft.targets.is_empty() { + lines.push(Line::from(Span::styled("(empty; press A to add a target)", Style::default().fg(palette::MUTED)))); + } + setup_error_line(setup, &mut lines); + let height = (lines.len() as u16 + 4).max(8); + ("Remote Setup".to_string(), lines, height) + } + SetupScreen::TargetForm => { + if let Some(form) = &setup.target_form { + lines.push(text_field_line("Label", &form.label, form.focus == TargetFormFocus::Label)); + lines.push(text_field_line("SSH", &form.ssh, form.focus == TargetFormFocus::Ssh)); + lines.push(text_field_line("Workspace", &form.workspace, form.focus == TargetFormFocus::Workspace)); + lines.push(action_line("Project Overrides...", form.focus == TargetFormFocus::Overrides)); + lines.push(action_line("Resources...", form.focus == TargetFormFocus::Resources)); + lines.push(action_line("Save", form.focus == TargetFormFocus::Save)); + lines.push(action_line("Cancel", form.focus == TargetFormFocus::Cancel)); + } + setup_error_line(setup, &mut lines); + let height = (lines.len() as u16 + 4).max(10); + ("Remote Target".to_string(), lines, height) + } + SetupScreen::Overrides => { + lines.push(Line::from("Tab/Up/Down select Enter edit Esc back")); + for (focus, label) in [ + (OverrideFocus::Tools, "Tools..."), + (OverrideFocus::Environment, "Environment names..."), + (OverrideFocus::Exclusions, "Snapshot exclusions..."), + (OverrideFocus::Artifacts, "Artifact paths..."), + (OverrideFocus::Done, "Done"), + ] { + lines.push(action_line(label, setup.override_focus == focus)); + } + setup_error_line(setup, &mut lines); + ("Project Remote Overrides".to_string(), lines, 12) + } + SetupScreen::Tools => { + lines.push(Line::from("A Add Enter/E Edit D Delete Esc Done")); + if let Some(tools) = &setup.tools { + for (index, tool) in tools.entries.iter().enumerate() { + let marker = if index == tools.selected { "> " } else { " " }; + let command = tool.command.as_deref().unwrap_or("(logical name)"); + lines.push(Line::from(format!("{marker}{} -> {}{}", tool.name, command, if tool.allow_args { " [args]" } else { "" }))); + } + if let Some(error) = &tools.error { + lines.push(Line::from(Span::styled(error.clone(), Style::default().fg(palette::ERROR)))); + } + } + setup_error_line(setup, &mut lines); + let height = (lines.len() as u16 + 4).max(8); + ("Remote Tools".to_string(), lines, height) + } + SetupScreen::ToolForm => { + if let Some(form) = &setup.tool_form { + lines.push(text_field_line("Logical name", &form.name, form.focus == 0)); + lines.push(text_field_line("Command basename", &form.command, form.focus == 1)); + lines.push(action_line(&format!("Allow args: {}", if form.allow_args { "yes" } else { "no" }), form.focus == 2)); + lines.push(action_line("Save", form.focus == 3)); + lines.push(action_line("Cancel", form.focus == 4)); + } + setup_error_line(setup, &mut lines); + ("Remote Tool".to_string(), lines, 11) + } + SetupScreen::Environment | SetupScreen::Exclusions | SetupScreen::Artifacts => { + if let Some(state) = &setup.strings { + lines.push(Line::from("A Add Enter/E Edit D Delete Esc Done")); + for (index, entry) in state.entries.iter().enumerate() { + let marker = if index == state.selected { "> " } else { " " }; + lines.push(Line::from(format!("{marker}{entry}"))); + } + if let Some(editing) = &state.editing { + lines.push(text_field_line("Value", editing, true)); + lines.push(Line::from("Enter accept Esc cancel")); + } + if let Some(error) = &state.error { + lines.push(Line::from(Span::styled(error.clone(), Style::default().fg(palette::ERROR)))); + } + setup_error_line(setup, &mut lines); + let height = (lines.len() as u16 + 4).max(8); + (state.kind.title().to_string(), lines, height) + } else { + ("Remote List".to_string(), vec![Line::from("No list state available")], 7) + } + } + SetupScreen::Resources | SetupScreen::AdvancedResources => { + if let Some(resources) = &setup.resources { + lines.push(Line::from("Blank unset Tab/Up/Down navigate A or Page toggles core/advanced")); + for index in resources.visible_indices() { + let input = &resources.inputs[index]; + lines.push(text_field_line( + input.label, + &input.value, + resources.focus == resources.visible_indices().iter().position(|candidate| *candidate == index).unwrap_or(0), + )); + } + lines.push(action_line("Use defaults (clear overrides)", resources.focus == resources.action_index(0))); + lines.push(action_line("Save", resources.focus == resources.action_index(1))); + lines.push(action_line("Cancel", resources.focus == resources.action_index(2))); + setup_error_line(setup, &mut lines); + let height = (lines.len() as u16 + 4).max(10); + (if resources.advanced { "Advanced Resources".to_string() } else { "Core Resources".to_string() }, lines, height) + } else { + ("Resources".to_string(), vec![Line::from("No resource state available")], 7) + } + } + SetupScreen::ConfirmDelete => { + let label = setup.delete_label.as_deref().unwrap_or(""); + ( + "Delete Remote Target".to_string(), + vec![Line::from(format!("Delete target \"{label}\"?")), Line::from("Enter Delete"), Line::from("Esc Cancel")], + 8, + ) + } + } +} + +fn setup_error_line(setup: &RemoteSetupState, lines: &mut Vec>) { + if let Some(error) = &setup.error { + lines.push(Line::from(Span::styled(error.clone(), Style::default().fg(palette::ERROR)))); + } +} + +fn text_field_line(label: &str, field: &TextField, focused: bool) -> Line<'static> { + let marker = if focused { "> " } else { " " }; + let value = text_field_value(field, focused); + Line::from(vec![ + Span::styled(marker, Style::default().fg(palette::ACCENT)), + Span::styled(format!("{label}: "), Style::default().fg(palette::MUTED)), + Span::styled(value, Style::default().fg(palette::FG)), + ]) +} + +fn action_line(label: &str, focused: bool) -> Line<'static> { + Line::from(vec![ + Span::styled(if focused { "> " } else { " " }, Style::default().fg(palette::ACCENT)), + Span::styled(label.to_string(), Style::default().fg(if focused { palette::FG } else { palette::MUTED })), + ]) +} + +fn text_field_value(field: &TextField, focused: bool) -> String { + if !focused { + return field.value.clone(); + } + let mut value = field.value.clone(); + let index = field.byte_index(field.cursor); + value.insert(index, '|'); + value +} + /// Renders one frame: writes the guest vt100 screen into the reduced guest /// viewport, then draws host-owned status and popup controls. fn render_frame( From 524abc4f4d5a3315fb645f1a6eaabec7c3eef35b Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 15:26:55 +0200 Subject: [PATCH 46/52] Add unit tests for TUI setup --- src/remote_target_ut.rs | 111 ++++++++++++++++++++++++++++++++++++++++ src/tui_ut.rs | 104 ++++++++++++++++++++++++++++++++++++- 2 files changed, 213 insertions(+), 2 deletions(-) diff --git a/src/remote_target_ut.rs b/src/remote_target_ut.rs index 30ee1c5..ef6c86b 100644 --- a/src/remote_target_ut.rs +++ b/src/remote_target_ut.rs @@ -463,3 +463,114 @@ fn lifecycle_and_admission_limits_reject_zero_or_excessive_values() { assert!(RemoteTargetConfig::load_from(&fixture.config).is_err()); } } + +#[test] +fn remote_project_draft_missing_file_is_empty_and_omits_optional_sections() { + let fixture = Fixture::new(); + let draft = RemoteConfigDraft::load_optional(&fixture.project, &ProjectConfig::default()).unwrap(); + assert_eq!(draft, None); + + let mut base = ProjectConfig::default(); + base.project.remote.environment = vec!["CC".to_string()]; + let mut draft = RemoteConfigDraft::default(); + draft.targets.insert( + "netbsd".to_string(), + RemoteTargetDraft { + ssh: "builder@build.example.test:2222".to_string(), + workspace: "/var/tmp/bunkerbox".to_string(), + project: None, + resources: None, + }, + ); + let yaml = draft.to_yaml().unwrap(); + assert!(yaml.contains("targets:")); + assert!(!yaml.contains("project:")); + assert!(!yaml.contains("resources:")); + draft.validate(&base).unwrap(); + + draft.targets.get_mut("netbsd").unwrap().resources = Some(RemoteResourceOverridesDraft::default()); + assert!(!draft.to_yaml().unwrap().contains("resources:")); +} + +#[test] +fn remote_project_draft_round_trips_deterministically_and_writes_private_file() { + let fixture = Fixture::new(); + let mut draft = RemoteConfigDraft::default(); + draft.targets.insert( + "zeta".to_string(), + RemoteTargetDraft { ssh: "builder@zeta.example.test".to_string(), workspace: "/var/tmp/zeta".to_string(), project: None, resources: None }, + ); + draft.targets.insert( + "alpha".to_string(), + RemoteTargetDraft { + ssh: "builder@alpha.example.test:2200".to_string(), + workspace: "/var/tmp/alpha".to_string(), + project: Some(RemoteProjectOverlayDraft { + remote: Some(RemoteOverlayDraft { + tools: Some(vec![RemoteToolSpec { name: "make".to_string(), command: Some("gmake".to_string()), allow_args: true }]), + environment: Some(vec!["CC".to_string()]), + ..RemoteOverlayDraft::default() + }), + }), + resources: Some(RemoteResourceOverridesDraft { + build_timeout: Some(RemoteQuantity::Text("10m".to_string())), + max_active_builds: Some(2), + ..RemoteResourceOverridesDraft::default() + }), + }, + ); + + draft.write_atomic(&fixture.project, &ProjectConfig::default()).unwrap(); + let path = fixture.project.join(".bunkerbox").join(REMOTE_PROJECT_CONFIG_FILE_NAME); + let first = fs::read_to_string(&path).unwrap(); + let loaded = RemoteConfigDraft::load_optional(&fixture.project, &ProjectConfig::default()).unwrap().unwrap(); + assert_eq!(loaded, draft); + assert_eq!(first, draft.to_yaml().unwrap()); + assert!(first.find("alpha:").unwrap() < first.find("zeta:").unwrap()); + #[cfg(unix)] + assert_eq!(fs::metadata(path).unwrap().permissions().mode() & 0o777, 0o600); +} + +#[test] +fn remote_project_draft_rejects_malformed_or_unsafe_destination_without_replacement() { + let fixture = Fixture::new(); + let bunkerbox = fixture.project.join(".bunkerbox"); + fs::create_dir(&bunkerbox).unwrap(); + let path = bunkerbox.join(REMOTE_PROJECT_CONFIG_FILE_NAME); + fs::write(&path, "targets: [not-a-map]\n").unwrap(); + assert!(RemoteConfigDraft::load_optional(&fixture.project, &ProjectConfig::default()).is_err()); + assert_eq!(fs::read_to_string(&path).unwrap(), "targets: [not-a-map]\n"); + + let draft = RemoteConfigDraft::default(); + let replacement = fixture.temp.path().join("replacement"); + fs::write(&replacement, b"do not replace").unwrap(); + fs::remove_file(&path).unwrap(); + #[cfg(unix)] + std::os::unix::fs::symlink(&replacement, &path).unwrap(); + #[cfg(unix)] + { + assert!(draft.write_atomic(&fixture.project, &ProjectConfig::default()).is_err()); + assert!(fs::symlink_metadata(&path).unwrap().file_type().is_symlink()); + } +} + +#[test] +fn remote_project_draft_validates_compact_fields_and_resource_units() { + let mut draft = RemoteConfigDraft::default(); + draft.targets.insert( + "builder".to_string(), + RemoteTargetDraft { + ssh: "builder@build.example.test".to_string(), + workspace: "/var/tmp/work".to_string(), + project: None, + resources: Some(RemoteResourceOverridesDraft { + max_output: Some(RemoteQuantity::Text("64M".to_string())), + ..RemoteResourceOverridesDraft::default() + }), + }, + ); + draft.validate(&ProjectConfig::default()).unwrap(); + + draft.targets.get_mut("builder").unwrap().ssh = "builder@bad host".to_string(); + assert!(draft.validate(&ProjectConfig::default()).is_err()); +} diff --git a/src/tui_ut.rs b/src/tui_ut.rs index 82b0ef2..ac799a1 100644 --- a/src/tui_ut.rs +++ b/src/tui_ut.rs @@ -1,8 +1,9 @@ use super::{ - dispatch_ui_command, guest_rows, mouse_to_bytes, process_status_bytes, HostPopup, HostUiState, MouseEncoding, MouseTracking, OverlayState, Term, + dispatch_ui_command, guest_rows, mouse_to_bytes, process_status_bytes, HostPopup, HostUiState, MouseEncoding, MouseTracking, OverlayState, + RemoteSetupState, SetupScreen, Term, }; use crate::cfg::ProjectConfig; -use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; +use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog, RemoteConfigDraft, RemoteTargetDraft}; use crossterm::event::{KeyCode, KeyEvent, KeyModifiers, MouseButton, MouseEvent, MouseEventKind}; use std::path::PathBuf; use std::sync::{Arc, Mutex}; @@ -148,3 +149,102 @@ fn mouse_motion_requires_the_requested_tracking_level() { assert_eq!(mouse_to_bytes(event, MouseTracking::Button, MouseEncoding::Sgr), None); assert_eq!(mouse_to_bytes(event, MouseTracking::Any, MouseEncoding::Sgr), Some(b"\x1b[<35;6;7M".to_vec())); } + +#[test] +fn remote_setup_shortcut_is_consumed_and_opens_empty_first_run_form() { + let catalog = BuildTargetCatalog::localhost_only(PathBuf::from("/tmp/project"), ProjectConfig::default()).unwrap(); + let active = ActiveBuildTarget::new(); + let mut host = HostUiState::new(&catalog); + let key = KeyEvent::new(KeyCode::Char('s'), KeyModifiers::CONTROL | KeyModifiers::ALT); + + assert!(host.handle_key(key, &catalog, &active)); + assert_eq!(host.popup, HostPopup::Setup); + assert!(host.setup.as_ref().unwrap().draft.targets.is_empty()); + assert!(host.handle_key(KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE), &catalog, &active)); + assert_eq!(host.popup, HostPopup::None); +} + +#[test] +fn remote_setup_target_form_editing_is_draft_only_until_save() { + let catalog = BuildTargetCatalog::localhost_only(PathBuf::from("/tmp/project"), ProjectConfig::default()).unwrap(); + let active = ActiveBuildTarget::new(); + let mut host = HostUiState::new(&catalog); + host.open_setup(&catalog); + + host.handle_setup_key(KeyEvent::new(KeyCode::Char('a'), KeyModifiers::NONE), &catalog); + assert_eq!(host.setup.as_ref().unwrap().screen, SetupScreen::TargetForm); + host.handle_setup_key(KeyEvent::new(KeyCode::Char('n'), KeyModifiers::NONE), &catalog); + host.handle_setup_key(KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE), &catalog); + assert_eq!(host.setup.as_ref().unwrap().screen, SetupScreen::List); + assert!(host.setup.as_ref().unwrap().draft.targets.is_empty()); + + assert!(host.handle_key(KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE), &catalog, &active)); + assert_eq!(host.popup, HostPopup::None); +} + +#[test] +fn remote_setup_delete_requires_confirmation_and_only_changes_the_draft() { + let catalog = BuildTargetCatalog::localhost_only(PathBuf::from("/tmp/project"), ProjectConfig::default()).unwrap(); + let active = ActiveBuildTarget::new(); + let mut draft = RemoteConfigDraft::default(); + draft.targets.insert( + "builder".to_string(), + RemoteTargetDraft { ssh: "builder@build.example.test".to_string(), workspace: "/var/tmp/work".to_string(), project: None, resources: None }, + ); + let mut host = HostUiState::new(&catalog); + host.setup = Some(RemoteSetupState::new(draft)); + host.popup = HostPopup::Setup; + + host.handle_setup_key(KeyEvent::new(KeyCode::Char('d'), KeyModifiers::NONE), &catalog); + assert_eq!(host.setup.as_ref().unwrap().screen, SetupScreen::ConfirmDelete); + host.handle_setup_key(KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE), &catalog); + assert!(host.setup.as_ref().unwrap().draft.targets.contains_key("builder")); + host.handle_setup_key(KeyEvent::new(KeyCode::Char('d'), KeyModifiers::NONE), &catalog); + host.handle_setup_key(KeyEvent::new(KeyCode::Enter, KeyModifiers::NONE), &catalog); + assert!(!host.setup.as_ref().unwrap().draft.targets.contains_key("builder")); + assert_eq!(host.popup, HostPopup::Setup); + assert_eq!(active.current(), "localhost"); +} + +#[test] +fn malformed_remote_setup_is_reported_without_creating_an_editable_draft() { + let project = tempfile::tempdir().unwrap(); + std::fs::create_dir(project.path().join(".bunkerbox")).unwrap(); + let config = project.path().join(".bunkerbox").join(crate::remote_target::REMOTE_PROJECT_CONFIG_FILE_NAME); + std::fs::write(&config, "targets: [broken]\n").unwrap(); + let catalog = BuildTargetCatalog::localhost_only(project.path().to_path_buf(), ProjectConfig::default()).unwrap(); + let active = ActiveBuildTarget::new(); + let mut host = HostUiState::new(&catalog); + + assert!(host.handle_key(KeyEvent::new(KeyCode::Char('s'), KeyModifiers::CONTROL | KeyModifiers::ALT), &catalog, &active)); + assert_eq!(host.popup, HostPopup::ConfigError); + assert!(host.setup.is_none()); + host.handle_key(KeyEvent::new(KeyCode::Char('v'), KeyModifiers::NONE), &catalog, &active); + assert!(host.config_error.as_ref().unwrap().view); + host.handle_key(KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE), &catalog, &active); + assert_eq!(host.popup, HostPopup::None); + assert_eq!(std::fs::read_to_string(config).unwrap(), "targets: [broken]\n"); +} + +#[test] +fn remote_setup_save_persists_for_next_run_without_mutating_frozen_runtime_state() { + let project = tempfile::tempdir().unwrap(); + let catalog = BuildTargetCatalog::localhost_only(project.path().to_path_buf(), ProjectConfig::default()).unwrap(); + let active = ActiveBuildTarget::new(); + let mut draft = RemoteConfigDraft::default(); + draft.targets.insert( + "builder".to_string(), + RemoteTargetDraft { ssh: "builder@build.example.test".to_string(), workspace: "/var/tmp/work".to_string(), project: None, resources: None }, + ); + let mut host = HostUiState::new(&catalog); + host.setup = Some(RemoteSetupState::new(draft)); + host.popup = HostPopup::Setup; + + host.handle_setup_key(KeyEvent::new(KeyCode::Char('s'), KeyModifiers::NONE), &catalog); + assert_eq!(host.popup, HostPopup::None); + assert_eq!(active.current(), "localhost"); + assert_eq!(catalog.summaries().iter().map(|target| target.label()).collect::>(), vec!["localhost"]); + assert!(host.confirmation.as_ref().unwrap().0.contains("next run")); + let saved = crate::remote_target::RemoteConfigDraft::load_optional(project.path(), catalog.base_project()).unwrap().unwrap(); + assert!(saved.targets.contains_key("builder")); +} From 8352875981103d06bbd21bb6d53b95a6c4feceda Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Tue, 18 Aug 2026 18:22:15 +0200 Subject: [PATCH 47/52] Fix remote regression --- src/daemon.rs | 164 +++++++++++++++++++++++++++++++++++--- src/daemon_ut.rs | 199 +++++++++++++++++++++++++++++++++++++++++++++-- src/remote.rs | 8 ++ 3 files changed, 355 insertions(+), 16 deletions(-) diff --git a/src/daemon.rs b/src/daemon.rs index 9d76216..bd08799 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -4,8 +4,9 @@ use crate::logging; use crate::loopback::{LoopbackBackend, RunRemoteSession}; use crate::proxy::{FilterProxy, UnixProxyHandle}; use crate::remote::{ - RemoteAdmissionLimits, RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, - RemoteEnvironmentPolicy, RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteTargetId, RemoteToolPolicy, + RemoteAdmissionLimits, RemoteAuthorizationError, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, RemoteBuild, + RemoteEnvironmentPolicy, RemoteExecutionContext, RemoteRequest, RemoteResourcePolicy, RemoteSnapshotId, RemoteTargetId, RemoteTool, + RemoteToolPolicy, RequestId, WorkspaceRelativePath, WorkspaceSessionId, }; use crate::remote_target::SshTarget; use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; @@ -290,6 +291,14 @@ impl RemoteBroker { self.cleanup_timeout } + fn allows_tool(&self, tool: &str) -> bool { + self.policy.allowed_tools().any(|allowed| allowed == tool) + } + + fn selected_environment_for_tool(&self, tool: &str, entries: &[(String, String)]) -> Result, String> { + self.policy.selected_environment_for_tool(tool, entries).map_err(|error| format!("remote environment rejected: {error:?}")) + } + pub async fn dispatch(&self, request: RemoteRequest, events: tokio::sync::mpsc::Sender) -> Result<(), RemoteDispatchError> { let authorized = match self.policy.authorize(&self.context, request) { Ok(authorized) => authorized, @@ -469,6 +478,121 @@ impl RemoteRouter { Ok((label, broker)) } + fn is_remote_selected(&self) -> bool { + self.active_target.current() != "localhost" + } + + async fn dispatch_transparent_exec( + &self, request: &ExecRequest, session_id: WorkspaceSessionId, writer: &mut W, + ) -> Result<(), String> { + let (label, broker) = self.active_remote_broker().map_err(|error| format!("remote dispatch failed: {error:?}"))?; + if !broker.allows_tool(&request.command) { + return Err(format!("remote tool is not configured for target {label}: {}", request.command)); + } + let cwd = Path::new(&request.cwd) + .strip_prefix("/workspace") + .map_err(|_| "current directory must be under /workspace".to_string())? + .to_str() + .ok_or_else(|| "current directory is not valid UTF-8".to_string())?; + let cwd = WorkspaceRelativePath::new(cwd)?; + let environment = broker.selected_environment_for_tool(&request.command, &request.env)?; + let sync_request = RemoteRequest::sync(new_remote_request_id(), session_id); + let snapshot_id = self + .dispatch_transparent_operation(&label, broker.clone(), sync_request, writer, true) + .await? + .ok_or_else(|| "remote sync did not return a retained capability".to_string())?; + let build = RemoteBuild::new(cwd, RemoteTool::new(request.command.clone())?, request.args.clone(), environment, snapshot_id)?; + let build_request = RemoteRequest::build(new_remote_request_id(), session_id, build); + let result = self.dispatch_transparent_operation(&label, broker, build_request, writer, false).await; + self.snapshots.lock().map_err(|_| "remote target snapshot lock poisoned".to_string())?.remove(&snapshot_id); + result.map(|_| ()) + } + + async fn dispatch_transparent_operation( + &self, label: &str, broker: Arc, request: RemoteRequest, writer: &mut W, sync: bool, + ) -> Result, String> { + let (event_tx, mut event_rx) = mpsc::channel(64); + let mut dispatch = Box::pin(broker.dispatch(request, event_tx)); + let mut dispatch_result = None; + let mut snapshot_id = None; + let mut terminal = false; + + loop { + if let Some(result) = dispatch_result.take() { + while let Ok(event) = event_rx.try_recv() { + self.forward_transparent_event(label, event, sync, writer, &mut snapshot_id, &mut terminal).await?; + } + dispatch_result = Some(result); + break; + } + + tokio::select! { + result = &mut dispatch => dispatch_result = Some(result), + event = event_rx.recv() => { + let Some(event) = event else { return Err("remote event stream closed before completion".to_string()) }; + self.forward_transparent_event(label, event, sync, writer, &mut snapshot_id, &mut terminal).await?; + } + } + } + + let result = dispatch_result.expect("transparent dispatch result is present"); + result.map_err(|error| format!("remote dispatch failed: {error:?}"))?; + if sync { + snapshot_id.ok_or_else(|| "remote sync completed without a capability".to_string()).map(Some) + } else if terminal { + Ok(None) + } else { + Err("remote build completed without an exit status".to_string()) + } + } + + async fn forward_transparent_event( + &self, label: &str, event: RemoteBackendEvent, sync: bool, writer: &mut W, snapshot_id: &mut Option, terminal: &mut bool, + ) -> Result<(), String> { + match event { + RemoteBackendEvent::SyncProgress { .. } => { + if !sync { + return Err("remote build returned a sync progress event".to_string()); + } + } + RemoteBackendEvent::SyncCompleted { snapshot_id: completed } => { + if !sync { + return Err("remote build returned a sync completion".to_string()); + } + self.snapshots.lock().map_err(|_| "remote target snapshot lock poisoned".to_string())?.insert(completed, label.to_string()); + *snapshot_id = Some(completed); + *terminal = true; + } + RemoteBackendEvent::Stdout(data) => { + if sync { + return Err("remote sync returned build output".to_string()); + } + write_frame(writer, &Frame::new(FrameType::Stdout, data)).await?; + } + RemoteBackendEvent::Stderr(data) => { + if sync { + return Err("remote sync returned build diagnostics".to_string()); + } + write_frame(writer, &Frame::new(FrameType::Stderr, data)).await?; + } + RemoteBackendEvent::Error { message } => { + if !sync { + write_frame(writer, &Frame::new(FrameType::Stderr, message.as_bytes().to_vec())).await?; + } + return Err(format!("remote target {label} request failed: {message}")); + } + RemoteBackendEvent::Cancelled => return Err(format!("remote target {label} request was cancelled")), + RemoteBackendEvent::Completed { exit_code } => { + if sync { + return Err("remote sync returned a build completion".to_string()); + } + write_frame(writer, &Frame::new(FrameType::Exit, exit_code.to_le_bytes().to_vec())).await?; + *terminal = true; + } + } + Ok(()) + } + async fn dispatch(&self, request: RemoteRequest, writer: &mut W) -> Result<(), String> { let (broker, label) = match request.operation() { crate::remote::RemoteOperation::Sync(_) => { @@ -914,14 +1038,8 @@ async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSess return Ok(()); } - if !is_allowed(&session.passthrough, &req.command, &req.args) { - logging::diagnostic(&format!("bunkerbox-vscomm: command '{}' not whitelisted", req.command)); - write_frame(&mut writer, &Frame::new(FrameType::Exit, 1i32.to_le_bytes().to_vec())).await?; - return Ok(()); - } - let command = req.command.clone(); - if let Err(err) = execute_request(&mut writer, session, &req).await { + if let Err(err) = dispatch_exec_request_for_session(&req, session, &mut writer).await { logging::diagnostic(&format!("bunkerbox: toolchain command '{command}' failed: {err}")); let _ = write_frame(&mut writer, &Frame::new(FrameType::Exit, 1i32.to_le_bytes().to_vec())).await; return Ok(()); @@ -930,10 +1048,36 @@ async fn handle_connection(stream: tokio_vsock::VsockStream, session: &VsockSess Ok(()) } +async fn dispatch_exec_request_for_session( + req: &ExecRequest, session: &VsockSession, writer: &mut W, +) -> Result<(), String> { + if session.remote_router.is_remote_selected() { + return session.remote_router.dispatch_transparent_exec(req, session.local_session_id, writer).await; + } + + if !is_allowed(&session.passthrough, &req.command, &req.args) { + logging::diagnostic(&format!("bunkerbox-vscomm: command '{}' not whitelisted", req.command)); + write_frame(writer, &Frame::new(FrameType::Exit, 1i32.to_le_bytes().to_vec())).await?; + return Ok(()); + } + + execute_request(writer, session, req).await +} + +fn new_remote_request_id() -> RequestId { + let mut bytes = [0u8; 16]; + loop { + rand::thread_rng().fill_bytes(&mut bytes); + if bytes != [0; 16] { + return RequestId(bytes); + } + } +} + async fn dispatch_remote_frame_for_session(frame: Frame, session: &VsockSession, writer: &mut W) -> Result<(), String> { let request = crate::vscomm::RemoteRequest::from_frame(frame).map_err(|err| format!("decode remote request: {err}"))?.into_domain()?; let local_snapshot = match request.operation() { - crate::remote::RemoteOperation::Build(build) => { + crate::remote::RemoteOperation::Build(build) if session.remote_router.active_target.current() == "localhost" => { session.local_capabilities.lock().map_err(|_| "local target capability lock poisoned".to_string())?.contains_key(&build.snapshot_id()) } _ => false, diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index 127461c..7be82b7 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -1,17 +1,21 @@ -use super::{dispatch_remote_frame, is_allowed, RemoteBroker, RemoteDispatchError, RemoteRouter}; +use super::{ + dispatch_exec_request_for_session, dispatch_remote_frame, dispatch_remote_frame_for_session, is_allowed, RemoteBroker, RemoteDispatchError, + RemoteRouter, +}; use super::{monitor_bwrap_status, ChildEvent}; -use crate::cfg::EnvMode; +use crate::cfg::{EnvMode, ProjectConfig}; use crate::remote::{ AuthorizedRemoteRequest, RemoteAdmissionLimits, RemoteAuthorizationPolicy, RemoteBackend, RemoteBackendError, RemoteBackendEvent, - RemoteExecutionContext, RemoteExecutionControl, RemoteFuture, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, RemoteTargetId, - RequestId, WorkspaceRelativePath, WorkspaceSessionId, + RemoteEnvironmentPolicy, RemoteExecutionContext, RemoteExecutionControl, RemoteFuture, RemoteRequest, RemoteSnapshotAuthority, RemoteSnapshotId, + RemoteTargetId, RequestId, WorkspaceRelativePath, WorkspaceSessionId, }; -use crate::remote_target::ActiveBuildTarget; +use crate::remote_target::{ActiveBuildTarget, BuildTargetCatalog}; use crate::vscomm::{ - Frame, RemoteBuild as WireRemoteBuild, RemoteRequest as WireRemoteRequest, RemoteTool as WireRemoteTool, RequestId as WireRequestId, + Frame, FrameType, RemoteBuild as WireRemoteBuild, RemoteRequest as WireRemoteRequest, RemoteTool as WireRemoteTool, RequestId as WireRequestId, WorkspaceRelativePath as WireWorkspaceRelativePath, WorkspaceSessionId as WireWorkspaceSessionId, }; use std::collections::BTreeMap; +use std::fs; use std::io::Write; use std::path::Path; use std::sync::{Arc, Mutex}; @@ -50,6 +54,27 @@ struct RecordingBackend { result: Option, } +struct TargetRecordingBackend { + calls: Mutex>, +} + +impl RemoteBackend for TargetRecordingBackend { + fn execute<'a>( + &'a self, request: AuthorizedRemoteRequest, _control: RemoteExecutionControl, events: mpsc::Sender, + ) -> RemoteFuture<'a, Result<(), RemoteBackendError>> { + let sync = matches!(request.request().operation(), crate::remote::RemoteOperation::Sync(_)); + self.calls.lock().unwrap().push(request); + Box::pin(async move { + let event = if sync { + RemoteBackendEvent::SyncCompleted { snapshot_id: RemoteSnapshotId::from_bytes([9; 16]) } + } else { + RemoteBackendEvent::Completed { exit_code: 0 } + }; + events.send(event).await.map_err(|_| RemoteBackendError::Cancelled) + }) + } +} + struct TestSnapshotAuthority; impl RemoteSnapshotAuthority for TestSnapshotAuthority { @@ -183,6 +208,44 @@ fn remote_broker(backend: Arc) -> RemoteBroker { ) } +fn target_catalog_fixture() -> (tempfile::TempDir, BuildTargetCatalog) { + let root = tempfile::tempdir().unwrap(); + fs::create_dir(root.path().join(".bunkerbox")).unwrap(); + fs::write( + root.path().join(".bunkerbox").join(crate::remote_target::REMOTE_PROJECT_CONFIG_FILE_NAME), + "targets:\n bsdbox:\n ssh: builder@bsdbox\n workspace: /var/tmp/bunkerbox\n project:\n remote:\n tools:\n - name: cargo\n allow-args: true\n", + ) + .unwrap(); + let catalog = BuildTargetCatalog::load_optional(root.path(), &ProjectConfig::default()).unwrap().unwrap(); + (root, catalog) +} + +fn target_session( + workspace: std::path::PathBuf, active: ActiveBuildTarget, backend: Arc, target: RemoteTargetId, + session_id: WorkspaceSessionId, +) -> super::VsockSession { + let policy = RemoteAuthorizationPolicy::from_policies( + target, + session_id, + vec![("cargo".to_string(), crate::remote::RemoteToolPolicy::new(true).with_command("cargo"))], + RemoteEnvironmentPolicy::default(), + ) + .unwrap() + .with_snapshot_authority(Arc::new(TestSnapshotAuthority)); + let context = RemoteExecutionContext { target, workspace_session_id: session_id }; + let broker = Arc::new(RemoteBroker::new(policy, context, backend)); + super::VsockSession { + passthrough: Arc::new(vec!["cargo *".to_string()]), + env_mode: EnvMode::Relaxed, + workspace, + merged_profile: None, + proxy_config: None, + remote_router: Arc::new(RemoteRouter::new(active, BTreeMap::from([("bsdbox".to_string(), broker)]))), + local_session_id: session_id, + local_capabilities: Mutex::new(BTreeMap::new()), + } +} + #[tokio::test] async fn authorized_remote_request_reaches_typed_backend() { let backend = Arc::new(RecordingBackend { @@ -464,3 +527,127 @@ fn local_exec_request_still_builds_on_the_local_path() { assert!(super::build_command(&session, &request, &cwd).is_ok()); } + +#[tokio::test] +async fn transparent_exec_request_routes_fresh_managed_cargo_to_selected_remote_target() { + let (root, catalog) = target_catalog_fixture(); + let active = ActiveBuildTarget::new(); + active.select(&catalog, "bsdbox").unwrap(); + let backend = Arc::new(TargetRecordingBackend { calls: Mutex::new(Vec::new()) }); + let session_id = WorkspaceSessionId([2; 16]); + let target = RemoteTargetId([3; 16]); + let session = target_session(root.path().to_path_buf(), active, backend.clone(), target, session_id); + let request = + crate::vscomm::ExecRequest { cwd: "/workspace".to_string(), command: "cargo".to_string(), args: vec!["build".to_string()], env: Vec::new() }; + let (mut guest, mut host) = tokio::io::duplex(4096); + + dispatch_exec_request_for_session(&request, &session, &mut host).await.unwrap(); + let exit = Frame::read_async(&mut guest).await.unwrap(); + assert!(matches!(exit.frame_type, FrameType::Exit)); + assert_eq!(i32::from_le_bytes(exit.payload.try_into().unwrap()), 0); + + let calls = backend.calls.lock().unwrap(); + assert_eq!(calls.len(), 2, "fresh remote transaction must sync before building"); + assert!(matches!(calls[0].request().operation(), crate::remote::RemoteOperation::Sync(_))); + assert!(matches!(calls[1].request().operation(), crate::remote::RemoteOperation::Build(_))); + assert_eq!(calls[0].target(), target); + assert_eq!(calls[1].target(), target); +} + +#[tokio::test] +async fn transparent_exec_request_keeps_selected_localhost_on_secured_local_executor() { + let workspace = tempfile::tempdir().unwrap(); + let active = ActiveBuildTarget::new(); + let session = super::VsockSession { + passthrough: Arc::new(vec!["/bin/touch *".to_string()]), + env_mode: EnvMode::Relaxed, + workspace: workspace.path().to_path_buf(), + merged_profile: None, + proxy_config: None, + remote_router: Arc::new(RemoteRouter::new(active, BTreeMap::new())), + local_session_id: WorkspaceSessionId([2; 16]), + local_capabilities: Mutex::new(BTreeMap::new()), + }; + let marker = workspace.path().join("local-executor-called"); + let request = crate::vscomm::ExecRequest { + cwd: "/workspace".to_string(), + command: "/bin/touch".to_string(), + args: vec![marker.to_string_lossy().into_owned()], + env: Vec::new(), + }; + let (mut guest, mut host) = tokio::io::duplex(4096); + + dispatch_exec_request_for_session(&request, &session, &mut host).await.unwrap(); + let exit = Frame::read_async(&mut guest).await.unwrap(); + assert!(matches!(exit.frame_type, FrameType::Exit)); + assert_eq!(i32::from_le_bytes(exit.payload.try_into().unwrap()), 0); + assert!(marker.is_file()); +} + +#[tokio::test] +async fn retained_remote_capability_remains_bound_after_switching_to_localhost() { + let (root, catalog) = target_catalog_fixture(); + let active = ActiveBuildTarget::new(); + active.select(&catalog, "bsdbox").unwrap(); + let backend = Arc::new(TargetRecordingBackend { calls: Mutex::new(Vec::new()) }); + let session_id = WorkspaceSessionId([2; 16]); + let target = RemoteTargetId([3; 16]); + let session = target_session(root.path().to_path_buf(), active.clone(), backend.clone(), target, session_id); + let (mut guest, mut host) = tokio::io::duplex(4096); + let sync = WireRemoteRequest::sync(WireRequestId([6; 16]), WireWorkspaceSessionId([2; 16])); + + dispatch_remote_frame_for_session(sync.to_frame().unwrap(), &session, &mut host).await.unwrap(); + let sync_event = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + let crate::vscomm::RemoteEventKind::SyncCompleted { snapshot_id } = sync_event.kind else { panic!("expected sync completion") }; + + active.select(&catalog, "localhost").unwrap(); + let build = WireRemoteRequest::build( + WireRequestId([7; 16]), + WireWorkspaceSessionId([2; 16]), + WireRemoteBuild::new( + WireWorkspaceRelativePath::new("src").unwrap(), + WireRemoteTool::new("cargo").unwrap(), + vec!["build".to_string()], + Vec::new(), + snapshot_id, + ) + .unwrap(), + ); + dispatch_remote_frame_for_session(build.to_frame().unwrap(), &session, &mut host).await.unwrap(); + let build_event = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + assert_eq!(build_event.kind, crate::vscomm::RemoteEventKind::Completed { exit_code: 0 }); + let calls = backend.calls.lock().unwrap(); + assert_eq!(calls.len(), 2); + assert_eq!(calls[0].target(), target); + assert_eq!(calls[1].target(), target); +} + +#[tokio::test] +async fn local_capability_cannot_bypass_a_new_remote_target_selection() { + let (root, catalog) = target_catalog_fixture(); + let active = ActiveBuildTarget::new(); + let backend = Arc::new(TargetRecordingBackend { calls: Mutex::new(Vec::new()) }); + let session_id = WorkspaceSessionId([2; 16]); + let target = RemoteTargetId([3; 16]); + let session = target_session(root.path().to_path_buf(), active.clone(), backend.clone(), target, session_id); + let local_snapshot = RemoteSnapshotId::from_bytes([8; 16]); + session.local_capabilities.lock().unwrap().insert(local_snapshot, ()); + active.select(&catalog, "bsdbox").unwrap(); + let build = WireRemoteRequest::build( + WireRequestId([8; 16]), + WireWorkspaceSessionId([2; 16]), + WireRemoteBuild::new( + WireWorkspaceRelativePath::new("src").unwrap(), + WireRemoteTool::new("cargo").unwrap(), + vec!["build".to_string()], + Vec::new(), + crate::vscomm::RemoteSnapshotId(*local_snapshot.as_bytes()), + ) + .unwrap(), + ); + let (_, mut host) = tokio::io::duplex(4096); + + assert!(dispatch_remote_frame_for_session(build.to_frame().unwrap(), &session, &mut host).await.is_err()); + assert!(backend.calls.lock().unwrap().is_empty()); + assert!(session.local_capabilities.lock().unwrap().contains_key(&local_snapshot)); +} diff --git a/src/remote.rs b/src/remote.rs index 79891d5..1531d6d 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -572,6 +572,14 @@ impl RemoteAuthorizationPolicy { self.allowed_tools.keys().map(String::as_str) } + pub fn selected_environment_for_tool(&self, tool: &str, entries: &[(String, String)]) -> Result, RemoteAuthorizationError> { + if tool == "cargo" { + return Ok(Vec::new()); + } + let selected = entries.iter().filter(|(name, _)| self.environment.allows(name)).cloned().collect::>(); + self.environment.filter(&selected) + } + pub fn authorize(&self, context: &RemoteExecutionContext, request: RemoteRequest) -> Result { if request.request_id.is_zero() { return Err(RemoteAuthorizationError::InvalidRequestId); From ad985f1efd82cb263d5f8a513dbd249d7aa5709a Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Wed, 19 Aug 2026 15:01:44 +0200 Subject: [PATCH 48/52] Add mxrun --- Makefile | 71 ++++++++++++++++++++++++++++++++++--- docs/reference/makefile.md | 32 +++++++++++++++++ mxrun.conf | 1 + scripts/maybe-mxrun.sh | 35 ++++++++++++++++++ scripts/mxrun-set-local.sh | 5 +++ scripts/mxrun-set-remote.sh | 8 +++++ scripts/mxrun-status.sh | 28 +++++++++++++++ 7 files changed, 175 insertions(+), 5 deletions(-) create mode 100644 mxrun.conf create mode 100755 scripts/maybe-mxrun.sh create mode 100755 scripts/mxrun-set-local.sh create mode 100755 scripts/mxrun-set-remote.sh create mode 100755 scripts/mxrun-status.sh diff --git a/Makefile b/Makefile index 8793fb5..7fb5a1f 100644 --- a/Makefile +++ b/Makefile @@ -1,5 +1,6 @@ .DEFAULT_GOAL := help -.PHONY: help ensure-toolchain dev release check test integration-test setup image install-image prepare config docs docs-dev docs-clean musl-vscomm worker-netbsd clean +.PHONY: help ensure-toolchain mxrun mxrun-init mxrun-toggle set-local-builds set-remote-builds dev release check test integration-test setup image install-image prepare config docs docs-dev docs-clean musl-vscomm worker-netbsd clean +.PHONY: _dev _release _check _test _integration-test DOCS_VENV := .venv-docs DOCS_MKDOCS := $(DOCS_VENV)/bin/mkdocs @@ -7,12 +8,21 @@ VSCOMM_TARGET := x86_64-unknown-linux-musl WORKER_TARGET ?= x86_64-unknown-netbsd IMAGE ?= OCI ?= +MXRUN_BIN ?= mxrun +MXRUN_ARGS ?= +MX_ACTIVE := $(shell awk -F= '/^active=/ {print $$2}' .mxrun-env 2>/dev/null) +export MXRUN_ARGS +export MXRUN_BIN help: @printf " %-24s %s\n" "Development" "" @printf " %-24s %s\n" " dev" "Build all binaries (host + musl-static vscomm)" - @printf " %-24s %s\n" " release" "Build optimized release binaries" @printf " %-24s %s\n" " check" "Format and lint" + @printf " %-24s %s\n" "" "" + @printf " %-24s %s\n" "Release" "" + @printf " %-24s %s\n" " release" "Build optimized release binaries" + @printf " %-24s %s\n" "" "" + @printf " %-24s %s\n" "Testing" "" @printf " %-24s %s\n" " test" "Run tests (requires cargo-nextest)" @printf " %-24s %s\n" " integration-test" "Run sandbox integration tests" @printf " %-24s %s\n" "" "" @@ -37,13 +47,28 @@ help: @printf " %-24s %s\n" "" "" @printf " %-24s %s\n" "Cleanup" "" @printf " %-24s %s\n" " clean" "Remove build artifacts (cargo clean)" + @printf " %-24s %s\n" "" "" + @printf " %-24s %s\n" "mxrun" "" + @printf " %-24s %s\n" " mxrun" "Show mxrun status" + @printf " %-24s %s\n" " mxrun-init" "Initialise mxrun with a local target" + @printf " %-24s %s\n" " mxrun-toggle" "Toggle mxrun delegation" + @printf " %-24s %s\n" " set-local-builds" "Disable mxrun delegation" + @printf " %-24s %s\n" " set-remote-builds" "Enable mxrun delegation" + @if [ "$(MX_ACTIVE)" = "yes" ]; then \ + printf " mxrun enabled; builds use the configured target matrix.\n"; \ + else \ + printf " mxrun disabled; builds run through local targets.\n"; \ + fi ensure-toolchain: @command -v rustup >/dev/null 2>&1 || { echo "rustup is required: https://rustup.rs" >&2; exit 1; } rustup update stable rustup target add $(VSCOMM_TARGET) -dev: ensure-toolchain +dev: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _dev; else scripts/maybe-mxrun.sh dev || $(MAKE) _dev; fi + +_dev: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image cargo build -p bunkerbox-worker cargo build --bin bunkerbox-netrelay --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) @@ -58,7 +83,10 @@ dev: ensure-toolchain cp target/$(VSCOMM_TARGET)/debug/bunkerbox-status target/dist/ cp target/$(VSCOMM_TARGET)/debug/bunkerbox-netrelay target/debug/bunkerbox-netrelay -release: ensure-toolchain +release: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _release; else scripts/maybe-mxrun.sh release || $(MAKE) _release; fi + +_release: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image --release cargo build --bin bunkerbox-netrelay --bin bunkerbox-vscomm --bin bunkerbox-remote --target $(VSCOMM_TARGET) --release cargo build --bin bunkerbox-status --target $(VSCOMM_TARGET) --release @@ -73,15 +101,48 @@ release: ensure-toolchain cp target/$(VSCOMM_TARGET)/release/bunkerbox-netrelay target/release/bunkerbox-netrelay check: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _check; else scripts/maybe-mxrun.sh check || $(MAKE) _check; fi + +_check: cargo fmt --all cargo clippy --all-targets --all-features -- -D warnings || cargo clippy --fix --all-targets --all-features --allow-dirty --allow-staged -- -D warnings test: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _test; else scripts/maybe-mxrun.sh test || $(MAKE) _test; fi + +_test: cargo nextest run -integration-test: dev +integration-test: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _integration-test; else scripts/maybe-mxrun.sh integration-test || $(MAKE) _integration-test; fi + +_integration-test: _dev cargo nextest run --test test_base --test test_sandbox +mxrun-toggle: + @if [ -f .mxrun-env ] && grep -q '^active=yes' .mxrun-env 2>/dev/null; then \ + sh scripts/mxrun-set-local.sh; \ + else \ + sh scripts/mxrun-set-remote.sh; \ + fi + +set-local-builds: + sh scripts/mxrun-set-local.sh + +set-remote-builds: + sh scripts/mxrun-set-remote.sh + +mxrun-init: + @command -v $(MXRUN_BIN) >/dev/null 2>&1 || { echo "Missing $(MXRUN_BIN). Install it first." >&2; exit 1; } + @if [ ! -f mxrun.conf ]; then printf 'local\n' > mxrun.conf; fi + @printf 'active=yes\n' > .mxrun-env + @MXRUN_CONFIG=mxrun.conf MXRUN_LOCAL_MAKE='$(MAKE)' $(MXRUN_BIN) init || true + +mxrun: + @command -v $(MXRUN_BIN) >/dev/null 2>&1 || { echo "Missing $(MXRUN_BIN). Install it first." >&2; exit 1; } + @if [ ! -f mxrun.conf ] && [ ! -f .mxrun-env ]; then printf 'active=no\n' > .mxrun-env; fi + @sh scripts/mxrun-status.sh + setup: dev target/debug/bunkerbox setup diff --git a/docs/reference/makefile.md b/docs/reference/makefile.md index 0f3e933..4e14e52 100644 --- a/docs/reference/makefile.md +++ b/docs/reference/makefile.md @@ -22,6 +22,38 @@ make check It formats the code and runs lint checks. +## mxrun Builds + +Development, release, and testing targets use the local mxrun target when mxrun +is enabled: + +```sh +make mxrun-init +make dev +make release +make test +``` + +The integration test target is delegated in the same way: + +```sh +make integration-test +``` + +Use these targets to control delegation: + +```sh +make mxrun +make set-local-builds +make set-remote-builds +``` + +Extra mxrun command-line options can be passed with `MXRUN_ARGS`: + +```sh +MXRUN_ARGS="--mirror-results" make test +``` + ## Setup Use this to prepare the host runtime pieces needed by Bunkerbox: diff --git a/mxrun.conf b/mxrun.conf new file mode 100644 index 0000000..4083037 --- /dev/null +++ b/mxrun.conf @@ -0,0 +1 @@ +local diff --git a/scripts/maybe-mxrun.sh b/scripts/maybe-mxrun.sh new file mode 100755 index 0000000..9f31dc2 --- /dev/null +++ b/scripts/maybe-mxrun.sh @@ -0,0 +1,35 @@ +#!/usr/bin/env sh +set -eu + +ENTRY="${1:-}" +[ -n "$ENTRY" ] || { echo "maybe-mxrun.sh: missing entry argument" >&2; exit 1; } + +# mxrun clears MXRUN_LOCAL_MAKE when calling back into Make. The callback must +# run the private local target instead of recursively delegating. +[ "${MXRUN_LOCAL_MAKE+set}" = "set" ] && exit 1 + +ACTIVE=no +[ -f .mxrun-env ] && ACTIVE=$(awk -F= '/^active=/ {print $2}' .mxrun-env 2>/dev/null) + +if [ "$ACTIVE" = "yes" ] && [ -f mxrun.conf ]; then + command -v "${MXRUN_BIN:-mxrun}" >/dev/null 2>&1 || { echo "Missing ${MXRUN_BIN:-mxrun} binary. Install it first." >&2; exit 1; } + case "$ENTRY" in + dev|check) + LABEL="Development Build Mode" ;; + release) + LABEL="Release Build Mode" ;; + test) + LABEL="Testing Everything" ;; + integration-test) + LABEL="Integration Tests" ;; + *) + LABEL="Bunkerbox Build" ;; + esac + # shellcheck disable=SC2086 + MXRUN_CONFIG=mxrun.conf "${MXRUN_BIN:-mxrun}" run --label="$LABEL" ${MXRUN_ARGS:-} "$ENTRY" || true + # mxrun handled the request or was interrupted; do not fall through locally. + exit 0 +fi + +# mxrun is inactive; the caller falls back to its private local target. +exit 1 diff --git a/scripts/mxrun-set-local.sh b/scripts/mxrun-set-local.sh new file mode 100755 index 0000000..ab1b31f --- /dev/null +++ b/scripts/mxrun-set-local.sh @@ -0,0 +1,5 @@ +#!/usr/bin/env sh +set -eu + +printf 'active=no\n' > .mxrun-env +printf '\nLocal builds enabled. To re-enable mxrun:\n make set-remote-builds\n\n' diff --git a/scripts/mxrun-set-remote.sh b/scripts/mxrun-set-remote.sh new file mode 100755 index 0000000..9b7b36e --- /dev/null +++ b/scripts/mxrun-set-remote.sh @@ -0,0 +1,8 @@ +#!/usr/bin/env sh +set -eu + +command -v "${MXRUN_BIN:-mxrun}" >/dev/null 2>&1 || { echo "Missing ${MXRUN_BIN:-mxrun} binary. Install it first." >&2; exit 1; } +[ -f mxrun.conf ] || { echo "No mxrun.conf found in this project. Create one first." >&2; exit 1; } + +printf 'active=yes\n' > .mxrun-env +printf '\nmxrun builds enabled. To switch back to local-only:\n make set-local-builds\n\n' diff --git a/scripts/mxrun-status.sh b/scripts/mxrun-status.sh new file mode 100755 index 0000000..954580f --- /dev/null +++ b/scripts/mxrun-status.sh @@ -0,0 +1,28 @@ +#!/usr/bin/env sh +set -eu + +if [ -f .mxrun-env ]; then + ACTIVE=$(awk -F= '/^active=/ {print $2}' .mxrun-env 2>/dev/null) +else + ACTIVE=no +fi + +if [ -f mxrun.conf ]; then + CONFIG_AVAILABLE=yes +else + CONFIG_AVAILABLE=no +fi + +if [ "$ACTIVE" = "yes" ] && [ "$CONFIG_AVAILABLE" = "yes" ]; then + printf '\nmxrun mode active; builds use the configured target matrix.\n\n' + printf ' To switch to local-only builds:\n' + printf ' make set-local-builds\n\n' +elif [ "$CONFIG_AVAILABLE" = "yes" ]; then + printf '\nmxrun available but inactive.\n\n' + printf ' To enable:\n' + printf ' make set-remote-builds\n\n' +else + printf '\nmxrun not configured; no mxrun.conf found.\n\n' + printf ' Create mxrun.conf or run:\n' + printf ' make mxrun-init\n\n' +fi From b1bbc7ddde7d9d4ee59cdd3b527fb83d904aed92 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Wed, 19 Aug 2026 15:26:01 +0200 Subject: [PATCH 49/52] Split mxrun for the remote worker --- Makefile | 108 ++++++++++++++++++++++--------------- docs/reference/makefile.md | 12 +++++ mxrun.conf | 2 + remote-mxrun.conf | 1 + scripts/maybe-mxrun.sh | 11 ++-- 5 files changed, 88 insertions(+), 46 deletions(-) create mode 100644 remote-mxrun.conf diff --git a/Makefile b/Makefile index 7fb5a1f..1fcb80c 100644 --- a/Makefile +++ b/Makefile @@ -1,6 +1,6 @@ .DEFAULT_GOAL := help -.PHONY: help ensure-toolchain mxrun mxrun-init mxrun-toggle set-local-builds set-remote-builds dev release check test integration-test setup image install-image prepare config docs docs-dev docs-clean musl-vscomm worker-netbsd clean -.PHONY: _dev _release _check _test _integration-test +.PHONY: help ensure-toolchain mxrun mxrun-init mxrun-toggle set-local-builds set-remote-builds dev worker-dev release worker check test integration-test setup image install-image prepare config docs docs-dev docs-clean musl-vscomm worker-netbsd clean +.PHONY: _dev _worker-dev _release _worker _check _test _integration-test DOCS_VENV := .venv-docs DOCS_MKDOCS := $(DOCS_VENV)/bin/mkdocs @@ -13,51 +13,61 @@ MXRUN_ARGS ?= MX_ACTIVE := $(shell awk -F= '/^active=/ {print $$2}' .mxrun-env 2>/dev/null) export MXRUN_ARGS export MXRUN_BIN +C_TITLE := \033[1;38;2;215;0;175m +C_COMMAND := \033[38;2;175;255;215m +C_DESCRIPTION := \033[38;2;128;128;128m +C_OFF := \033[0m help: - @printf " %-24s %s\n" "Development" "" - @printf " %-24s %s\n" " dev" "Build all binaries (host + musl-static vscomm)" - @printf " %-24s %s\n" " check" "Format and lint" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Release" "" - @printf " %-24s %s\n" " release" "Build optimized release binaries" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Testing" "" - @printf " %-24s %s\n" " test" "Run tests (requires cargo-nextest)" - @printf " %-24s %s\n" " integration-test" "Run sandbox integration tests" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Toolchain" "" - @printf " %-24s %s\n" " ensure-toolchain" "Install/update Rust stable and musl target" - @printf " %-24s %s\n" " musl-vscomm" "Build static vscomm binary only" - @printf " %-24s %s\n" " worker-netbsd" "Build the portable worker for NetBSD" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Image" "" - @printf " %-24s %s\n" " image" "Build OCI agent image (requires IMAGE=)" - @printf " %-24s %s\n" " install-image" "Install OCI archive (requires OCI=)" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Setup" "" - @printf " %-24s %s\n" " setup" "Install containerd, CNI, Kata dependencies" - @printf " %-24s %s\n" " prepare" "Prepare workspace overlay layers" - @printf " %-24s %s\n" " config" "Configure project interactively" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Documentation" "" - @printf " %-24s %s\n" " docs" "Build documentation site" - @printf " %-24s %s\n" " docs-dev" "Serve docs locally with live reload" - @printf " %-24s %s\n" " docs-clean" "Remove docs build artifacts" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "Cleanup" "" - @printf " %-24s %s\n" " clean" "Remove build artifacts (cargo clean)" - @printf " %-24s %s\n" "" "" - @printf " %-24s %s\n" "mxrun" "" - @printf " %-24s %s\n" " mxrun" "Show mxrun status" - @printf " %-24s %s\n" " mxrun-init" "Initialise mxrun with a local target" - @printf " %-24s %s\n" " mxrun-toggle" "Toggle mxrun delegation" - @printf " %-24s %s\n" " set-local-builds" "Disable mxrun delegation" - @printf " %-24s %s\n" " set-remote-builds" "Enable mxrun delegation" + @printf '$(C_TITLE)%s$(C_OFF)\n' "Development" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "dev" "Build all binaries (host + musl-static vscomm)" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "worker-dev" "Build only the remote worker (development)" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "check" "Format and lint" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Release" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "release" "Build optimized release binaries" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "worker" "Build only the remote worker (release)" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Testing" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "test" "Run tests (requires cargo-nextest)" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "integration-test" "Run sandbox integration tests" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Toolchain" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "ensure-toolchain" "Install/update Rust stable and musl target" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "musl-vscomm" "Build static vscomm binary only" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "worker-netbsd" "Build the portable worker for NetBSD" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Image" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "image" "Build OCI agent image (requires IMAGE=)" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "install-image" "Install OCI archive (requires OCI=)" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Setup" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "setup" "Install containerd, CNI, Kata dependencies" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "prepare" "Prepare workspace overlay layers" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "config" "Configure project interactively" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Documentation" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "docs" "Build documentation site" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "docs-dev" "Serve docs locally with live reload" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "docs-clean" "Remove docs build artifacts" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Cleanup" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "clean" "Remove build artifacts (cargo clean)" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "Utils" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "help" "Show this help" + @printf '\n' + @printf '$(C_TITLE)%s$(C_OFF)\n' "mxrun" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "mxrun" "Show mxrun status" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "mxrun-init" "Initialise mxrun with a local target" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "mxrun-toggle" "Toggle mxrun delegation" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "set-local-builds" "Disable mxrun delegation" + @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "set-remote-builds" "Enable mxrun delegation" + @printf '\n' @if [ "$(MX_ACTIVE)" = "yes" ]; then \ - printf " mxrun enabled; builds use the configured target matrix.\n"; \ + printf "$(C_OFF)mxrun enabled; builds use the configured target matrix.$(C_OFF)\n"; \ else \ - printf " mxrun disabled; builds run through local targets.\n"; \ + printf "$(C_OFF)mxrun disabled; builds run through local targets.$(C_OFF)\n"; \ fi ensure-toolchain: @@ -83,6 +93,12 @@ _dev: ensure-toolchain cp target/$(VSCOMM_TARGET)/debug/bunkerbox-status target/dist/ cp target/$(VSCOMM_TARGET)/debug/bunkerbox-netrelay target/debug/bunkerbox-netrelay +worker-dev: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _worker-dev; else scripts/maybe-mxrun.sh worker-dev || $(MAKE) _worker-dev; fi + +_worker-dev: + cargo build -p bunkerbox-worker + release: @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _release; else scripts/maybe-mxrun.sh release || $(MAKE) _release; fi @@ -100,6 +116,12 @@ _release: ensure-toolchain cp target/$(VSCOMM_TARGET)/release/bunkerbox-status target/dist/ cp target/$(VSCOMM_TARGET)/release/bunkerbox-netrelay target/release/bunkerbox-netrelay +worker: + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _worker; else scripts/maybe-mxrun.sh worker || $(MAKE) _worker; fi + +_worker: + cargo build -p bunkerbox-worker --release + check: @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _check; else scripts/maybe-mxrun.sh check || $(MAKE) _check; fi diff --git a/docs/reference/makefile.md b/docs/reference/makefile.md index 4e14e52..37996aa 100644 --- a/docs/reference/makefile.md +++ b/docs/reference/makefile.md @@ -34,6 +34,18 @@ make release make test ``` +Remote worker-only builds use the same mxrun dispatch: + +```sh +make worker-dev +make worker +``` + +The remote worker target must be present in `remote-mxrun.conf`; these targets +do not build the host binaries. +The worker dispatch selects it with mxrun's `--config remote-mxrun.conf` +command-line option. + The integration test target is delegated in the same way: ```sh diff --git a/mxrun.conf b/mxrun.conf index 4083037..43d6de3 100644 --- a/mxrun.conf +++ b/mxrun.conf @@ -1 +1,3 @@ local +# Add the remote worker target here before using worker-dev or worker. +# Format: @:/absolute/path/to/bunkerbox diff --git a/remote-mxrun.conf b/remote-mxrun.conf new file mode 100644 index 0000000..4083037 --- /dev/null +++ b/remote-mxrun.conf @@ -0,0 +1 @@ +local diff --git a/scripts/maybe-mxrun.sh b/scripts/maybe-mxrun.sh index 9f31dc2..ed3c3d9 100755 --- a/scripts/maybe-mxrun.sh +++ b/scripts/maybe-mxrun.sh @@ -14,9 +14,9 @@ ACTIVE=no if [ "$ACTIVE" = "yes" ] && [ -f mxrun.conf ]; then command -v "${MXRUN_BIN:-mxrun}" >/dev/null 2>&1 || { echo "Missing ${MXRUN_BIN:-mxrun} binary. Install it first." >&2; exit 1; } case "$ENTRY" in - dev|check) + dev|worker-dev|check) LABEL="Development Build Mode" ;; - release) + release|worker) LABEL="Release Build Mode" ;; test) LABEL="Testing Everything" ;; @@ -26,7 +26,12 @@ if [ "$ACTIVE" = "yes" ] && [ -f mxrun.conf ]; then LABEL="Bunkerbox Build" ;; esac # shellcheck disable=SC2086 - MXRUN_CONFIG=mxrun.conf "${MXRUN_BIN:-mxrun}" run --label="$LABEL" ${MXRUN_ARGS:-} "$ENTRY" || true + case "$ENTRY" in + worker-dev|worker) + "${MXRUN_BIN:-mxrun}" --config remote-mxrun.conf run --label="$LABEL" ${MXRUN_ARGS:-} "$ENTRY" || true ;; + *) + MXRUN_CONFIG=mxrun.conf "${MXRUN_BIN:-mxrun}" run --label="$LABEL" ${MXRUN_ARGS:-} "$ENTRY" || true ;; + esac # mxrun handled the request or was interrupted; do not fall through locally. exit 0 fi From cd7bd84acb01773ff5cf8a347c7892676d299b75 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Wed, 19 Aug 2026 23:35:53 +0200 Subject: [PATCH 50/52] Add worker remote builds --- Makefile | 31 +++++++------- docs/reference/makefile.md | 6 +++ scripts/linux-amd64.sh | 84 ++++++++++++++++++++++++++++++++++++++ scripts/netbsd-amd64.sh | 84 ++++++++++++++++++++++++++++++++++++++ scripts/worker.sh | 28 +++++++++++++ 5 files changed, 216 insertions(+), 17 deletions(-) create mode 100755 scripts/linux-amd64.sh create mode 100755 scripts/netbsd-amd64.sh create mode 100755 scripts/worker.sh diff --git a/Makefile b/Makefile index 1fcb80c..8672d31 100644 --- a/Makefile +++ b/Makefile @@ -10,9 +10,6 @@ IMAGE ?= OCI ?= MXRUN_BIN ?= mxrun MXRUN_ARGS ?= -MX_ACTIVE := $(shell awk -F= '/^active=/ {print $$2}' .mxrun-env 2>/dev/null) -export MXRUN_ARGS -export MXRUN_BIN C_TITLE := \033[1;38;2;215;0;175m C_COMMAND := \033[38;2;175;255;215m C_DESCRIPTION := \033[38;2;128;128;128m @@ -64,7 +61,7 @@ help: @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "set-local-builds" "Disable mxrun delegation" @printf ' $(C_COMMAND)%-20s$(C_OFF) $(C_DESCRIPTION)%s$(C_OFF)\n' "set-remote-builds" "Enable mxrun delegation" @printf '\n' - @if [ "$(MX_ACTIVE)" = "yes" ]; then \ + @if [ "$$(awk -F= '/^active=/ {print $$2}' .mxrun-env 2>/dev/null)" = "yes" ]; then \ printf "$(C_OFF)mxrun enabled; builds use the configured target matrix.$(C_OFF)\n"; \ else \ printf "$(C_OFF)mxrun disabled; builds run through local targets.$(C_OFF)\n"; \ @@ -76,7 +73,7 @@ ensure-toolchain: rustup target add $(VSCOMM_TARGET) dev: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _dev; else scripts/maybe-mxrun.sh dev || $(MAKE) _dev; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _dev; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh dev || $(MAKE) _dev; fi _dev: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image @@ -94,13 +91,13 @@ _dev: ensure-toolchain cp target/$(VSCOMM_TARGET)/debug/bunkerbox-netrelay target/debug/bunkerbox-netrelay worker-dev: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _worker-dev; else scripts/maybe-mxrun.sh worker-dev || $(MAKE) _worker-dev; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _worker-dev; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh worker-dev || $(MAKE) _worker-dev; fi _worker-dev: - cargo build -p bunkerbox-worker + scripts/worker.sh worker-dev release: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _release; else scripts/maybe-mxrun.sh release || $(MAKE) _release; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _release; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh release || $(MAKE) _release; fi _release: ensure-toolchain cargo build --bin bunkerbox --bin bunkerbox-image --release @@ -117,42 +114,42 @@ _release: ensure-toolchain cp target/$(VSCOMM_TARGET)/release/bunkerbox-netrelay target/release/bunkerbox-netrelay worker: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _worker; else scripts/maybe-mxrun.sh worker || $(MAKE) _worker; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _worker; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh worker || $(MAKE) _worker; fi _worker: - cargo build -p bunkerbox-worker --release + scripts/worker.sh worker check: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _check; else scripts/maybe-mxrun.sh check || $(MAKE) _check; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _check; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh check || $(MAKE) _check; fi _check: cargo fmt --all cargo clippy --all-targets --all-features -- -D warnings || cargo clippy --fix --all-targets --all-features --allow-dirty --allow-staged -- -D warnings test: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _test; else scripts/maybe-mxrun.sh test || $(MAKE) _test; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _test; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh test || $(MAKE) _test; fi _test: cargo nextest run integration-test: - @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _integration-test; else scripts/maybe-mxrun.sh integration-test || $(MAKE) _integration-test; fi + @if [ -n "$$SSH_CONNECTION" ]; then $(MAKE) _integration-test; else MXRUN_BIN="$(MXRUN_BIN)" MXRUN_ARGS="$(MXRUN_ARGS)" scripts/maybe-mxrun.sh integration-test || $(MAKE) _integration-test; fi _integration-test: _dev cargo nextest run --test test_base --test test_sandbox mxrun-toggle: @if [ -f .mxrun-env ] && grep -q '^active=yes' .mxrun-env 2>/dev/null; then \ - sh scripts/mxrun-set-local.sh; \ + MXRUN_BIN="$(MXRUN_BIN)" sh scripts/mxrun-set-local.sh; \ else \ - sh scripts/mxrun-set-remote.sh; \ + MXRUN_BIN="$(MXRUN_BIN)" sh scripts/mxrun-set-remote.sh; \ fi set-local-builds: - sh scripts/mxrun-set-local.sh + MXRUN_BIN="$(MXRUN_BIN)" sh scripts/mxrun-set-local.sh set-remote-builds: - sh scripts/mxrun-set-remote.sh + MXRUN_BIN="$(MXRUN_BIN)" sh scripts/mxrun-set-remote.sh mxrun-init: @command -v $(MXRUN_BIN) >/dev/null 2>&1 || { echo "Missing $(MXRUN_BIN). Install it first." >&2; exit 1; } diff --git a/docs/reference/makefile.md b/docs/reference/makefile.md index 37996aa..d0cf7f3 100644 --- a/docs/reference/makefile.md +++ b/docs/reference/makefile.md @@ -45,6 +45,12 @@ The remote worker target must be present in `remote-mxrun.conf`; these targets do not build the host binaries. The worker dispatch selects it with mxrun's `--config remote-mxrun.conf` command-line option. +Once `make worker-dev` or `make worker` runs on a target, the Makefile detects +that target's `uname -s` and `uname -m` values and dispatches to the matching +`scripts/.sh` worker toolchain script. +The platform script provisions Rust under the invoking user's home directory +when needed. It does not use `sudo`; missing system prerequisites produce an +explicit root-required error. The integration test target is delegated in the same way: diff --git a/scripts/linux-amd64.sh b/scripts/linux-amd64.sh new file mode 100755 index 0000000..091e883 --- /dev/null +++ b/scripts/linux-amd64.sh @@ -0,0 +1,84 @@ +#!/usr/bin/env sh +set -eu + +ACTION="${1:-}" +CARGO="${CARGO:-cargo}" + +require_linker() { + command -v cc >/dev/null 2>&1 || command -v clang >/dev/null 2>&1 || command -v gcc >/dev/null 2>&1 || { + echo "Linux worker setup requires a C compiler (cc, clang, or gcc)." >&2 + echo "Install one as root, then rerun the worker build." >&2 + exit 1 + } +} + +setup() { + [ -n "${HOME:-}" ] || { + echo "Linux worker setup requires HOME for a user-local Rust toolchain" >&2 + exit 1 + } + + export PATH="$HOME/.cargo/bin:$PATH" + if command -v "$CARGO" >/dev/null 2>&1 && command -v rustc >/dev/null 2>&1; then + require_linker + return + fi + + if [ "$(id -u)" -eq 0 ]; then + echo "Linux worker setup needs Rust, but this process is root." >&2 + echo "Install cargo and rustc system-wide, or rerun as a non-root user." >&2 + exit 1 + fi + + if ! command -v rustup >/dev/null 2>&1; then + rustup_script=$(mktemp "${TMPDIR:-/tmp}/bunkerbox-rustup.XXXXXX") + cleanup() { rm -f "$rustup_script"; } + trap cleanup 0 1 2 15 + + if command -v curl >/dev/null 2>&1; then + curl --fail --silent --show-error --location https://sh.rustup.rs > "$rustup_script" + elif command -v fetch >/dev/null 2>&1; then + fetch -o "$rustup_script" https://sh.rustup.rs + elif command -v wget >/dev/null 2>&1; then + wget -qO "$rustup_script" https://sh.rustup.rs + else + echo "Linux worker setup requires curl, fetch, or wget." >&2 + echo "Install one as root, then rerun the worker build." >&2 + exit 1 + fi + + if ! sh "$rustup_script" -y --profile minimal --default-toolchain stable --no-modify-path; then + echo "Linux worker setup could not install Rust user-locally." >&2 + echo "Install cargo and rustc as root, or fix the user-local rustup setup." >&2 + exit 1 + fi + fi + + export PATH="$HOME/.cargo/bin:$PATH" + if ! rustup toolchain install stable --profile minimal; then + echo "Linux worker setup could not install the stable Rust toolchain." >&2 + exit 1 + fi + if ! rustup default stable; then + echo "Linux worker setup could not select the stable Rust toolchain." >&2 + exit 1 + fi + + command -v "$CARGO" >/dev/null 2>&1 && command -v rustc >/dev/null 2>&1 || { + echo "Linux worker setup could not provide cargo and rustc." >&2 + exit 1 + } + require_linker +} + +setup + +case "$ACTION" in + worker-dev) + exec "$CARGO" build -p bunkerbox-worker ;; + worker) + exec "$CARGO" build -p bunkerbox-worker --release ;; + *) + echo "Unsupported Linux worker action: $ACTION" >&2 + exit 1 ;; +esac diff --git a/scripts/netbsd-amd64.sh b/scripts/netbsd-amd64.sh new file mode 100755 index 0000000..99001f0 --- /dev/null +++ b/scripts/netbsd-amd64.sh @@ -0,0 +1,84 @@ +#!/usr/bin/env sh +set -eu + +ACTION="${1:-}" +CARGO="${CARGO:-cargo}" + +require_linker() { + command -v cc >/dev/null 2>&1 || command -v clang >/dev/null 2>&1 || command -v gcc >/dev/null 2>&1 || { + echo "NetBSD worker setup requires a C compiler (cc, clang, or gcc)." >&2 + echo "Install one as root, then rerun the worker build." >&2 + exit 1 + } +} + +setup() { + [ -n "${HOME:-}" ] || { + echo "NetBSD worker setup requires HOME for a user-local Rust toolchain" >&2 + exit 1 + } + + export PATH="$HOME/.cargo/bin:$PATH" + if command -v "$CARGO" >/dev/null 2>&1 && command -v rustc >/dev/null 2>&1; then + require_linker + return + fi + + if [ "$(id -u)" -eq 0 ]; then + echo "NetBSD worker setup needs Rust, but this process is root." >&2 + echo "Install cargo and rustc system-wide, or rerun as a non-root user." >&2 + exit 1 + fi + + if ! command -v rustup >/dev/null 2>&1; then + rustup_script=$(mktemp "${TMPDIR:-/tmp}/bunkerbox-rustup.XXXXXX") + cleanup() { rm -f "$rustup_script"; } + trap cleanup 0 1 2 15 + + if command -v curl >/dev/null 2>&1; then + curl --fail --silent --show-error --location https://sh.rustup.rs > "$rustup_script" + elif command -v fetch >/dev/null 2>&1; then + fetch -o "$rustup_script" https://sh.rustup.rs + elif command -v wget >/dev/null 2>&1; then + wget -qO "$rustup_script" https://sh.rustup.rs + else + echo "NetBSD worker setup requires curl, fetch, or wget." >&2 + echo "Install one as root, then rerun the worker build." >&2 + exit 1 + fi + + if ! sh "$rustup_script" -y --profile minimal --default-toolchain stable --no-modify-path; then + echo "NetBSD worker setup could not install Rust user-locally." >&2 + echo "Install cargo and rustc as root, or fix the user-local rustup setup." >&2 + exit 1 + fi + fi + + export PATH="$HOME/.cargo/bin:$PATH" + if ! rustup toolchain install stable --profile minimal; then + echo "NetBSD worker setup could not install the stable Rust toolchain." >&2 + exit 1 + fi + if ! rustup default stable; then + echo "NetBSD worker setup could not select the stable Rust toolchain." >&2 + exit 1 + fi + + command -v "$CARGO" >/dev/null 2>&1 && command -v rustc >/dev/null 2>&1 || { + echo "NetBSD worker setup could not provide cargo and rustc." >&2 + exit 1 + } + require_linker +} + +setup + +case "$ACTION" in + worker-dev) + exec "$CARGO" build -p bunkerbox-worker ;; + worker) + exec "$CARGO" build -p bunkerbox-worker --release ;; + *) + echo "Unsupported NetBSD worker action: $ACTION" >&2 + exit 1 ;; +esac diff --git a/scripts/worker.sh b/scripts/worker.sh new file mode 100755 index 0000000..ab8b9b0 --- /dev/null +++ b/scripts/worker.sh @@ -0,0 +1,28 @@ +#!/usr/bin/env sh +set -eu + +ACTION="${1:-}" +[ -n "$ACTION" ] || { echo "worker.sh: missing worker action" >&2; exit 1; } + +OS=$(uname -s) +ARCH=$(uname -m) + +case "$OS:$ARCH" in + Linux:x86_64|Linux:amd64) + PLATFORM=linux-amd64 ;; + NetBSD:x86_64|NetBSD:amd64) + PLATFORM=netbsd-amd64 ;; + *) + echo "Unsupported worker platform: OS=$OS ARCH=$ARCH" >&2 + exit 1 ;; +esac + +SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd) +PLATFORM_SCRIPT="$SCRIPT_DIR/$PLATFORM.sh" + +[ -x "$PLATFORM_SCRIPT" ] || { + echo "Missing worker platform script: $PLATFORM_SCRIPT" >&2 + exit 1 +} + +exec "$PLATFORM_SCRIPT" "$ACTION" From 7737fab2d0bdfda52d0fb0ec3d103c75f597b8a4 Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Thu, 20 Aug 2026 01:20:26 +0200 Subject: [PATCH 51/52] Worker snapshotting --- crates/bunkerbox-worker-protocol/src/lib.rs | 146 +++++++++++- .../bunkerbox-worker-protocol/src/lib_ut.rs | 28 +++ crates/bunkerbox-worker/src/platform.rs | 9 + crates/bunkerbox-worker/src/storage.rs | 32 ++- crates/bunkerbox-worker/src/storage_ut.rs | 45 ++++ crates/bunkerbox-worker/src/worker.rs | 7 +- src/daemon_ut.rs | 38 +++ src/snapshot.rs | 217 ++++++++++++++++-- src/snapshot_ut.rs | 101 ++++++-- src/ssh.rs | 17 +- src/ssh_ut.rs | 5 +- 11 files changed, 584 insertions(+), 61 deletions(-) diff --git a/crates/bunkerbox-worker-protocol/src/lib.rs b/crates/bunkerbox-worker-protocol/src/lib.rs index 757c7d8..ccbb007 100644 --- a/crates/bunkerbox-worker-protocol/src/lib.rs +++ b/crates/bunkerbox-worker-protocol/src/lib.rs @@ -10,7 +10,7 @@ //! allocating a payload buffer, and every length and count inside a payload is //! checked before it can drive an allocation. -use std::collections::BTreeSet; +use std::collections::{BTreeMap, BTreeSet}; use std::fmt; use std::io::{self, Read, Write}; @@ -24,6 +24,8 @@ pub const WORKER_VERSION: u16 = WORKER_PROTOCOL_VERSION; pub const WORKER_ARTIFACT_PROTOCOL_VERSION: u16 = 2; /// Adds trusted target-PATH command resolution without changing V1/V2 fields. pub const WORKER_COMMAND_PROTOCOL_VERSION: u16 = 3; +/// Adds validated workspace-relative symlink entries to upload manifests. +pub const WORKER_SYMLINK_PROTOCOL_VERSION: u16 = 4; pub const WORKER_FRAME_HEADER_LEN: usize = 4 + 2 + 1 + 4; pub const WORKER_FRAME_HEADER_SIZE: usize = WORKER_FRAME_HEADER_LEN; pub const WORKER_ID_LEN: usize = 16; @@ -435,6 +437,7 @@ impl WorkerArtifactEntry { pub enum WorkerEntryKind { Directory = 1, File = 2, + Symlink = 3, } impl WorkerEntryKind { @@ -442,6 +445,7 @@ impl WorkerEntryKind { match value { 1 => Ok(Self::Directory), 2 => Ok(Self::File), + 3 => Ok(Self::Symlink), _ => Err(invalid(format!("unknown worker entry kind: {value}"))), } } @@ -458,11 +462,18 @@ pub struct WorkerUploadEntry { pub mode: u32, pub size: u64, pub digest: Option, + pub symlink_target: Option, } impl WorkerUploadEntry { pub fn new(path: impl Into, kind: WorkerEntryKind, mode: u32, size: u64, digest: Option) -> WorkerResult { - let entry = Self { path: WorkerRelativePath::for_entry(path)?, kind, mode, size, digest }; + Self::new_with_target(path, kind, mode, size, digest, None) + } + + pub fn new_with_target( + path: impl Into, kind: WorkerEntryKind, mode: u32, size: u64, digest: Option, symlink_target: Option, + ) -> WorkerResult { + let entry = Self { path: WorkerRelativePath::for_entry(path)?, kind, mode, size, digest, symlink_target }; entry.validate()?; Ok(entry) } @@ -475,6 +486,10 @@ impl WorkerUploadEntry { Self::new(path, WorkerEntryKind::File, mode, size, Some(digest)) } + pub fn symlink(path: impl Into, mode: u32, target: impl Into) -> WorkerResult { + Self::new_with_target(path, WorkerEntryKind::Symlink, mode, 0, None, Some(target.into())) + } + pub fn path(&self) -> &WorkerRelativePath { &self.path } @@ -495,6 +510,10 @@ impl WorkerUploadEntry { self.digest.as_ref() } + pub fn symlink_target(&self) -> Option<&str> { + self.symlink_target.as_deref() + } + pub fn validate(&self) -> WorkerResult<()> { validate_relative_path("worker entry path", self.path.as_str(), false)?; if self.mode & !0o7777 != 0 { @@ -509,6 +528,9 @@ impl WorkerUploadEntry { if self.digest.is_some() { return Err(invalid("worker directory entry must not have a digest")); } + if self.symlink_target.is_some() { + return Err(invalid("worker directory entry must not have a symlink target")); + } } WorkerEntryKind::File => { if self.size > MAX_WORKER_FILE_BYTES { @@ -517,6 +539,19 @@ impl WorkerUploadEntry { if self.digest.is_none() { return Err(invalid("worker file entry is missing a digest")); } + if self.symlink_target.is_some() { + return Err(invalid("worker file entry must not have a symlink target")); + } + } + WorkerEntryKind::Symlink => { + if self.size != 0 { + return Err(invalid("worker symlink entry must have zero size")); + } + if self.digest.is_some() { + return Err(invalid("worker symlink entry must not have a digest")); + } + let target = self.symlink_target.as_deref().ok_or_else(|| invalid("worker symlink entry is missing a target"))?; + validate_symlink_target(self.path.as_str(), target)?; } } Ok(()) @@ -936,6 +971,14 @@ impl WorkerMessage { } } + fn requires_symlink_version(&self) -> bool { + match self { + Self::UploadBegin { entries, .. } => entries.iter().any(|entry| entry.kind() == WorkerEntryKind::Symlink), + Self::UploadEntry { entry, .. } => entry.kind() == WorkerEntryKind::Symlink, + _ => false, + } + } + pub fn encode(&self) -> WorkerResult> { self.encode_version(WORKER_PROTOCOL_VERSION) } @@ -948,6 +991,9 @@ impl WorkerMessage { if version < WORKER_COMMAND_PROTOCOL_VERSION && matches!(self, Self::Build { build, .. } if build.target_command.is_some()) { return Err(invalid("worker command identity requires command-capable protocol version")); } + if version < WORKER_SYMLINK_PROTOCOL_VERSION && self.requires_symlink_version() { + return Err(invalid("worker symlink entries require symlink-capable protocol version")); + } self.validate()?; let mut payload = WireWriter::new(); encode_payload(self, &mut payload, version)?; @@ -1143,6 +1189,7 @@ pub fn validate_upload_manifest(entries: &[WorkerUploadEntry]) -> WorkerResult MAX_WORKER_MANIFEST_BYTES { @@ -1153,6 +1200,7 @@ pub fn validate_upload_manifest(entries: &[WorkerUploadEntry]) -> WorkerResult { encode_correlation(writer, *request_id, *session_id)?; writer.id(upload_id.0)?; writer.u32(*entry_index)?; - encode_entry(writer, entry)?; + encode_entry(writer, entry, version)?; } WorkerMessage::UploadFileChunk { request_id, session_id, upload_id, path, offset, data } => { encode_correlation(writer, *request_id, *session_id)?; @@ -1268,7 +1316,7 @@ fn decode_payload(kind: WorkerFrameKind, payload: &[u8], version: u16) -> Worker let count = reader.count(MAX_WORKER_UPLOAD_ENTRIES, "worker upload entries")?; let mut entries = Vec::with_capacity(count); for _ in 0..count { - entries.push(decode_entry(&mut reader)?); + entries.push(decode_entry(&mut reader, version)?); } WorkerMessage::UploadBegin { request_id, session_id, upload_id, entries } } @@ -1276,7 +1324,7 @@ fn decode_payload(kind: WorkerFrameKind, payload: &[u8], version: u16) -> Worker let (request_id, session_id) = decode_correlation(&mut reader)?; let upload_id = WorkerUploadId(reader.array16()?); let entry_index = reader.u32()?; - let entry = decode_entry(&mut reader)?; + let entry = decode_entry(&mut reader, version)?; WorkerMessage::UploadEntry { request_id, session_id, upload_id, entry_index, entry } } WorkerFrameKind::UploadFileChunk => { @@ -1480,7 +1528,7 @@ fn decode_environment(reader: &mut WireReader<'_>, field: &str) -> WorkerResult< Ok(environment) } -fn encode_entry(writer: &mut WireWriter, entry: &WorkerUploadEntry) -> WorkerResult<()> { +fn encode_entry(writer: &mut WireWriter, entry: &WorkerUploadEntry, version: u16) -> WorkerResult<()> { entry.validate()?; writer.string(entry.path.as_str(), MAX_WORKER_PATH_BYTES, "worker entry path")?; writer.u8(entry.kind.as_u8())?; @@ -1493,10 +1541,19 @@ fn encode_entry(writer: &mut WireWriter, entry: &WorkerUploadEntry) -> WorkerRes } None => writer.boolean(false)?, } + if version >= WORKER_SYMLINK_PROTOCOL_VERSION { + match entry.symlink_target() { + Some(target) => { + writer.boolean(true)?; + writer.string(target, MAX_WORKER_PATH_BYTES, "worker symlink target")?; + } + None => writer.boolean(false)?, + } + } Ok(()) } -fn decode_entry(reader: &mut WireReader<'_>) -> WorkerResult { +fn decode_entry(reader: &mut WireReader<'_>, version: u16) -> WorkerResult { let path = reader.string(MAX_WORKER_PATH_BYTES, "worker entry path")?; let kind = WorkerEntryKind::from_u8(reader.u8()?)?; let mode = reader.u32()?; @@ -1505,7 +1562,15 @@ fn decode_entry(reader: &mut WireReader<'_>) -> WorkerResult true => Some(reader.array32()?), false => None, }; - WorkerUploadEntry::new(path, kind, mode, size, digest) + let symlink_target = if version >= WORKER_SYMLINK_PROTOCOL_VERSION { + match reader.boolean("worker symlink target flag")? { + true => Some(reader.string(MAX_WORKER_PATH_BYTES, "worker symlink target")?), + false => None, + } + } else { + None + }; + WorkerUploadEntry::new_with_target(path, kind, mode, size, digest, symlink_target) } fn encode_correlation(writer: &mut WireWriter, request_id: WorkerRequestId, session_id: WorkerSessionId) -> WorkerResult<()> { @@ -1680,6 +1745,67 @@ fn validate_relative_path(field: &str, value: &str, allow_empty: bool) -> Worker Ok(()) } +fn validate_symlink_target(path: &str, target: &str) -> WorkerResult<()> { + normalize_symlink_target(path, target).map(|_| ()) +} + +fn normalize_symlink_target(path: &str, target: &str) -> WorkerResult { + validate_text("worker symlink target", target, MAX_WORKER_PATH_BYTES)?; + if target.is_empty() || target.starts_with('/') || target.starts_with('\\') || target.contains('\\') { + return Err(invalid("worker symlink target must be non-empty and relative")); + } + let mut components = path.rsplit_once('/').map_or_else(Vec::new, |(parent, _)| parent.split('/').collect::>()); + let target_components = target.split('/').collect::>(); + if target_components.len() > MAX_WORKER_PATH_DEPTH { + return Err(invalid("worker symlink target exceeds maximum depth")); + } + for (index, component) in target_components.iter().enumerate() { + if component.is_empty() { + if index + 1 == target_components.len() { + continue; + } + return Err(invalid("worker symlink target contains an empty component")); + } + match *component { + "." => {} + ".." => { + components.pop().ok_or_else(|| invalid("worker symlink target escapes workspace"))?; + } + value => { + if value.len() > MAX_WORKER_PATH_COMPONENT_BYTES || value.chars().any(char::is_control) || value.contains(':') { + return Err(invalid("worker symlink target contains an invalid component")); + } + components.push(value); + } + } + } + Ok(components.join("/")) +} + +fn validate_upload_symlinks(entries: &[WorkerUploadEntry]) -> WorkerResult<()> { + let by_path = entries.iter().map(|entry| (entry.path().as_str(), entry)).collect::>(); + for entry in entries.iter().filter(|entry| entry.kind() == WorkerEntryKind::Symlink) { + let target = entry.symlink_target().ok_or_else(|| invalid("worker symlink entry is missing a target"))?; + let mut current = normalize_symlink_target(entry.path().as_str(), target)?; + let mut visited = BTreeSet::new(); + loop { + if current.is_empty() { + break; + } + let target_entry = by_path.get(current.as_str()).ok_or_else(|| invalid("worker symlink target is missing from upload manifest"))?; + if target_entry.kind() != WorkerEntryKind::Symlink { + break; + } + if !visited.insert(current.clone()) { + return Err(invalid("worker symlink loop detected")); + } + let nested_target = target_entry.symlink_target().ok_or_else(|| invalid("worker symlink entry is missing a target"))?; + current = normalize_symlink_target(target_entry.path().as_str(), nested_target)?; + } + } + Ok(()) +} + fn validate_text(field: &str, value: &str, maximum: usize) -> WorkerResult<()> { if value.len() > maximum { return Err(invalid(format!("{field} exceeds maximum length {maximum}"))); @@ -1702,7 +1828,7 @@ fn invalid(message: impl Into) -> WorkerProtocolError { } fn is_supported_version(version: u16) -> bool { - matches!(version, WORKER_PROTOCOL_VERSION | WORKER_ARTIFACT_PROTOCOL_VERSION | WORKER_COMMAND_PROTOCOL_VERSION) + matches!(version, WORKER_PROTOCOL_VERSION | WORKER_ARTIFACT_PROTOCOL_VERSION | WORKER_COMMAND_PROTOCOL_VERSION | WORKER_SYMLINK_PROTOCOL_VERSION) } fn validate_version(version: u16) -> WorkerResult<()> { diff --git a/crates/bunkerbox-worker-protocol/src/lib_ut.rs b/crates/bunkerbox-worker-protocol/src/lib_ut.rs index 6ea7f74..a63dae9 100644 --- a/crates/bunkerbox-worker-protocol/src/lib_ut.rs +++ b/crates/bunkerbox-worker-protocol/src/lib_ut.rs @@ -211,6 +211,34 @@ fn manifests_require_sorted_paths_and_valid_metadata() { assert!(WorkerUploadEntry::new("file", WorkerEntryKind::File, 0o644, 1, Some([1; 32])).is_ok()); } +#[test] +fn symlink_entries_round_trip_only_in_the_symlink_protocol() { + let (request_id, session_id, upload_id) = ids(); + let message = WorkerMessage::UploadBegin { + request_id, + session_id, + upload_id, + entries: vec![ + WorkerUploadEntry::directory("src", 0o755).unwrap(), + WorkerUploadEntry::symlink("src/link", 0o777, "../target").unwrap(), + WorkerUploadEntry::directory("target", 0o755).unwrap(), + ], + }; + + assert!(message.encode_version(WORKER_COMMAND_PROTOCOL_VERSION).is_err()); + let frame = message.encode_version(WORKER_SYMLINK_PROTOCOL_VERSION).unwrap(); + let (version, decoded) = WorkerMessage::decode_versioned(&frame).unwrap(); + assert_eq!(version, WORKER_SYMLINK_PROTOCOL_VERSION); + assert_eq!(decoded, message); +} + +#[test] +fn symlink_targets_cannot_be_absolute_or_escape() { + assert!(WorkerUploadEntry::symlink("link", 0o777, "/etc").is_err()); + assert!(WorkerUploadEntry::symlink("src/link", 0o777, "../../outside").is_err()); + assert!(WorkerUploadEntry::symlink("link", 0o777, "").is_err()); +} + #[test] fn duplicate_environment_keys_and_bad_argv_are_rejected() { assert!(WorkerBuild::new( diff --git a/crates/bunkerbox-worker/src/platform.rs b/crates/bunkerbox-worker/src/platform.rs index ea69974..2119908 100644 --- a/crates/bunkerbox-worker/src/platform.rs +++ b/crates/bunkerbox-worker/src/platform.rs @@ -77,6 +77,15 @@ pub fn create_dir_at(parent: &File, name: &str, mode: u32) -> io::Result<()> { Ok(()) } +pub fn create_symlink_at(parent: &File, name: &str, target: &str) -> io::Result<()> { + let name = c_string(name)?; + let target = c_string(target)?; + if unsafe { libc::symlinkat(target.as_ptr(), parent.as_raw_fd(), name.as_ptr()) } != 0 { + return Err(io::Error::last_os_error()); + } + Ok(()) +} + pub fn chmod_fd(file: &File, mode: u32) -> io::Result<()> { if unsafe { libc::fchmod(file.as_raw_fd(), mode as libc::mode_t) } != 0 { return Err(io::Error::last_os_error()); diff --git a/crates/bunkerbox-worker/src/storage.rs b/crates/bunkerbox-worker/src/storage.rs index c2d519a..239f956 100644 --- a/crates/bunkerbox-worker/src/storage.rs +++ b/crates/bunkerbox-worker/src/storage.rs @@ -21,7 +21,7 @@ const LOCK_FILE: &str = "lock"; const FILES_DIRECTORY: &str = "files"; const QUOTA_LOCK_FILE: &str = "quota.lock"; const MANIFEST_MAGIC: [u8; 4] = *b"BBWM"; -const MANIFEST_VERSION: u16 = 1; +const MANIFEST_VERSION: u16 = 2; const COPY_BUFFER_BYTES: usize = 64 * 1024; const MAX_STALE_SESSIONS: usize = 64; const MAX_STALE_UPLOADS_PER_SESSION: usize = 256; @@ -511,6 +511,7 @@ impl UploadTransaction { }, ); } + WorkerEntryKind::Symlink => {} } } platform::sync_fd(&self.files).map_err(|error| format!("flush worker upload files: {error}"))?; @@ -604,6 +605,10 @@ impl StoredUpload { let target = create_relative_file(destination, entry.path().as_str(), entry.mode())?; copy_and_verify_with_interrupt(&source, &target, entry, disconnected)?; } + WorkerEntryKind::Symlink => { + let target = entry.symlink_target().ok_or_else(|| "worker symlink entry has no target".to_string())?; + create_relative_symlink(destination, entry.path().as_str(), target)?; + } } } platform::sync_fd(destination).map_err(|error| format!("flush worker job workspace: {error}"))?; @@ -782,6 +787,13 @@ impl StoredManifest { } None => bytes.push(0), } + match entry.symlink_target() { + Some(target) => { + bytes.push(1); + put_string(&mut bytes, target)?; + } + None => bytes.push(0), + } if bytes.len() > MAX_WORKER_MANIFEST_BYTES { return Err("worker stored manifest exceeds maximum length".to_string()); } @@ -812,7 +824,12 @@ impl StoredManifest { 1 => Some(reader.array32()?), _ => return Err("worker stored manifest has invalid digest flag".to_string()), }; - entries.push(WorkerUploadEntry::new(path, kind, mode, size, digest).map_err(protocol_error)?); + let symlink_target = match reader.u8()? { + 0 => None, + 1 => Some(reader.string()?), + _ => return Err("worker stored manifest has invalid symlink target flag".to_string()), + }; + entries.push(WorkerUploadEntry::new_with_target(path, kind, mode, size, digest, symlink_target).map_err(protocol_error)?); } reader.finish()?; validate_upload_manifest(&entries).map_err(protocol_error)?; @@ -986,6 +1003,13 @@ fn create_relative_file(root: &File, relative: &str, mode: u32) -> Result Result<(), String> { + let components = components(relative)?; + let (file_name, parents) = components.split_last().ok_or_else(|| "worker symlink path is empty".to_string())?; + let parent = open_relative_directory_components(root, parents)?; + platform::create_symlink_at(&parent, file_name, target).map_err(|error| format!("create worker symlink {relative}: {error}")) +} + fn open_relative_file(root: &File, relative: &str) -> Result { let components = components(relative)?; let (file_name, parents) = components.split_last().ok_or_else(|| "worker file path is empty".to_string())?; @@ -1023,8 +1047,8 @@ fn validate_manifest_structure(entries: &[WorkerUploadEntry]) -> Result<(), Stri prefix.push('/'); } prefix.push_str(component); - if matches!(kinds.get(prefix.as_str()), Some(WorkerEntryKind::File)) { - return Err(format!("worker manifest path collides with a file: {}", entry.path().as_str())); + if matches!(kinds.get(prefix.as_str()), Some(kind) if *kind != WorkerEntryKind::Directory) { + return Err(format!("worker manifest path collides with a non-directory: {}", entry.path().as_str())); } } } diff --git a/crates/bunkerbox-worker/src/storage_ut.rs b/crates/bunkerbox-worker/src/storage_ut.rs index b01101a..5bc6549 100644 --- a/crates/bunkerbox-worker/src/storage_ut.rs +++ b/crates/bunkerbox-worker/src/storage_ut.rs @@ -38,6 +38,14 @@ fn entries(contents: &[u8]) -> Vec { ] } +fn symlink_entries() -> Vec { + vec![ + WorkerUploadEntry::directory("src", 0o755).unwrap(), + WorkerUploadEntry::symlink("src/link", 0o777, "../target").unwrap(), + WorkerUploadEntry::directory("target", 0o755).unwrap(), + ] +} + #[test] fn completed_upload_is_reopened_and_materialized_by_opaque_identity() { let (_temp, store) = store_fixture(); @@ -58,6 +66,43 @@ fn completed_upload_is_reopened_and_materialized_by_opaque_identity() { assert!(store.open_completed(WorkerSessionId([9; 16]), UPLOAD).is_err()); } +#[test] +fn completed_upload_recreates_relative_symlinks_without_file_contents() { + let (_temp, store) = store_fixture(); + let mut transaction = store.begin(SESSION, UPLOAD, symlink_entries()).unwrap(); + transaction.commit().unwrap(); + drop(transaction); + + let stored = store.open_completed(SESSION, UPLOAD).unwrap(); + let jobs = store.jobs_directory().unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + stored.materialize(job.root()).unwrap(); + + let link = open_relative_file(job.root(), "src/link"); + assert!(link.is_err()); + let src = platform::open_dir_at(job.root(), "src").unwrap(); + let metadata = platform::stat_at(&src, "link").unwrap(); + assert_eq!(metadata.st_mode & libc::S_IFMT, libc::S_IFLNK); +} + +#[test] +fn worker_materialization_rejects_replaced_symlink_parent() { + let (temp, store) = store_fixture(); + let mut transaction = store.begin(SESSION, UPLOAD, symlink_entries()).unwrap(); + transaction.commit().unwrap(); + drop(transaction); + + let stored = store.open_completed(SESSION, UPLOAD).unwrap(); + let jobs = store.jobs_directory().unwrap(); + let job = JobWorkspace::create(&jobs).unwrap(); + let outside = temp.path().join("outside"); + fs::create_dir(&outside).unwrap(); + platform::create_symlink_at(job.root(), "src", outside.to_str().unwrap()).unwrap(); + + assert!(stored.materialize(job.root()).is_err()); + assert!(!outside.join("link").exists()); +} + #[test] fn incomplete_and_digest_failed_uploads_are_removed_and_token_cannot_be_reused_while_active() { let (_temp, store) = store_fixture(); diff --git a/crates/bunkerbox-worker/src/worker.rs b/crates/bunkerbox-worker/src/worker.rs index c7be00a..472895f 100644 --- a/crates/bunkerbox-worker/src/worker.rs +++ b/crates/bunkerbox-worker/src/worker.rs @@ -3,6 +3,7 @@ use crate::storage::{ArtifactSpool, UploadStore, UploadTransaction}; use bunkerbox_worker_protocol::{ WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRequestId, WorkerSessionId, WorkerUploadId, MAX_WORKER_CHUNK_BYTES, MAX_WORKER_ERROR_BYTES, WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_COMMAND_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, + WORKER_SYMLINK_PROTOCOL_VERSION, }; use std::io::{self, Read, Write}; use std::sync::atomic::{AtomicBool, AtomicU16, AtomicUsize, Ordering}; @@ -86,7 +87,11 @@ impl WorkerService { if version != hello_version { return Err("worker Hello version does not match frame version".to_string()); } - if version != WORKER_PROTOCOL_VERSION && version != WORKER_ARTIFACT_PROTOCOL_VERSION && version != WORKER_COMMAND_PROTOCOL_VERSION { + if version != WORKER_PROTOCOL_VERSION + && version != WORKER_ARTIFACT_PROTOCOL_VERSION + && version != WORKER_COMMAND_PROTOCOL_VERSION + && version != WORKER_SYMLINK_PROTOCOL_VERSION + { return Err(format!("unsupported worker protocol version: {version}")); } if session_id.0 == [0; 16] { diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index 7be82b7..2046689 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -554,6 +554,44 @@ async fn transparent_exec_request_routes_fresh_managed_cargo_to_selected_remote_ assert_eq!(calls[1].target(), target); } +#[tokio::test] +async fn managed_remote_wrapper_requests_route_to_target_selected_after_daemon_setup() { + let (root, catalog) = target_catalog_fixture(); + let active = ActiveBuildTarget::new(); + let backend = Arc::new(TargetRecordingBackend { calls: Mutex::new(Vec::new()) }); + let session_id = WorkspaceSessionId([2; 16]); + let target = RemoteTargetId([3; 16]); + let session = target_session(root.path().to_path_buf(), active.clone(), backend.clone(), target, session_id); + active.select(&catalog, "bsdbox").unwrap(); + + let (mut guest, mut host) = tokio::io::duplex(4096); + let sync = WireRemoteRequest::sync(WireRequestId([6; 16]), WireWorkspaceSessionId([2; 16])); + dispatch_remote_frame_for_session(sync.to_frame().unwrap(), &session, &mut host).await.unwrap(); + let sync_event = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + let crate::vscomm::RemoteEventKind::SyncCompleted { snapshot_id } = sync_event.kind else { panic!("expected sync completion") }; + + let build = WireRemoteRequest::build( + WireRequestId([7; 16]), + WireWorkspaceSessionId([2; 16]), + WireRemoteBuild::new( + WireWorkspaceRelativePath::new("").unwrap(), + WireRemoteTool::new("cargo").unwrap(), + vec!["build".to_string()], + Vec::new(), + snapshot_id, + ) + .unwrap(), + ); + dispatch_remote_frame_for_session(build.to_frame().unwrap(), &session, &mut host).await.unwrap(); + let build_event = crate::vscomm::RemoteEvent::from_frame(Frame::read_async(&mut guest).await.unwrap()).unwrap(); + assert_eq!(build_event.kind, crate::vscomm::RemoteEventKind::Completed { exit_code: 0 }); + + let calls = backend.calls.lock().unwrap(); + assert_eq!(calls.len(), 2); + assert_eq!(calls[0].target(), target); + assert_eq!(calls[1].target(), target); +} + #[tokio::test] async fn transparent_exec_request_keeps_selected_localhost_on_secured_local_executor() { let workspace = tempfile::tempdir().unwrap(); diff --git a/src/snapshot.rs b/src/snapshot.rs index e462462..2eda12a 100644 --- a/src/snapshot.rs +++ b/src/snapshot.rs @@ -3,7 +3,7 @@ use crate::remote::{RemoteExecutionControl, WorkspaceSessionId}; use crate::workspace::WorkspaceHandle; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; -use std::collections::BTreeSet; +use std::collections::{BTreeMap, BTreeSet}; use std::ffi::{CStr, CString, OsStr, OsString}; use std::fs::{self, File, OpenOptions}; use std::io::{self, Read, Write}; @@ -53,6 +53,7 @@ impl SnapshotRelativePath { pub enum SnapshotEntryKind { Directory, RegularFile, + Symlink, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -62,6 +63,7 @@ pub struct SnapshotEntry { mode: u16, size: u64, content_digest: Option<[u8; 32]>, + symlink_target: Option, } impl SnapshotEntry { @@ -84,6 +86,10 @@ impl SnapshotEntry { pub fn content_digest(&self) -> Option<&[u8; 32]> { self.content_digest.as_ref() } + + pub fn symlink_target(&self) -> Option<&str> { + self.symlink_target.as_deref() + } } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -306,6 +312,10 @@ impl SnapshotStore { &control, )?; } + SnapshotEntryKind::Symlink => { + let target = entry.symlink_target.as_deref().ok_or_else(|| "symlink snapshot entry has no target".to_string())?; + create_relative_symlink(&destination_root, entry.path.as_str(), target)?; + } } } if control.is_cancelled() { @@ -515,8 +525,10 @@ impl SnapshotBuilder { next_buffer: vec![0; SNAPSHOT_COPY_BUFFER_BYTES], control, }; - walk_directory(root.as_raw_fd(), "", 0, &self.limits, &self.exclusions, &mut state)?; + walk_directory(&canonical_root, root.as_raw_fd(), "", 0, &self.limits, &self.exclusions, &mut state)?; state.entries.sort_by(|left, right| left.path.as_str().cmp(right.path.as_str())); + validate_snapshot_structure(&state.entries)?; + validate_snapshot_symlinks(&state.entries)?; let id = snapshot_id(&state.entries); let stored = StoredSnapshot::from_entries(session_id, &state.entries, state.total_file_bytes); let manifest = serde_json::to_vec(&stored).map_err(|error| format!("encode snapshot manifest: {error}"))?; @@ -576,7 +588,8 @@ struct WalkState { } fn walk_directory( - directory_fd: RawFd, parent: &str, depth: usize, limits: &SnapshotLimits, exclusions: &SnapshotExclusionPolicy, state: &mut WalkState, + canonical_root: &Path, directory_fd: RawFd, parent: &str, depth: usize, limits: &SnapshotLimits, exclusions: &SnapshotExclusionPolicy, + state: &mut WalkState, ) -> Result<(), String> { if state.control.is_cancelled() { return Err("snapshot creation cancelled".to_string()); @@ -613,8 +626,9 @@ fn walk_directory( mode, size: 0, content_digest: None, + symlink_target: None, }); - walk_directory(child.as_raw_fd(), &relative, depth + 1, limits, exclusions, state)?; + walk_directory(canonical_root, child.as_raw_fd(), &relative, depth + 1, limits, exclusions, state)?; } SnapshotEntryKind::RegularFile => { add_entry_limit(state.entries.len(), limits)?; @@ -637,6 +651,20 @@ fn walk_directory( mode: normalized_mode(child_stat.st_mode), size, content_digest: Some(digest), + symlink_target: None, + }); + } + SnapshotEntryKind::Symlink => { + add_entry_limit(state.entries.len(), limits)?; + let target = read_symlink_at(directory_fd, &name, &relative)?; + validate_symlink_target(canonical_root, &relative, &target, limits)?; + state.entries.push(SnapshotEntry { + path: SnapshotRelativePath::new(relative)?, + kind: SnapshotEntryKind::Symlink, + mode: normalized_mode(child_stat.st_mode), + size: 0, + content_digest: None, + symlink_target: Some(target), }); } } @@ -743,11 +771,119 @@ fn validate_component(value: &str, max: usize) -> Result<(), String> { Ok(()) } +fn read_symlink_at(parent: RawFd, name: &OsStr, path: &str) -> Result { + let name = CString::new(name.as_bytes()).map_err(|_| format!("snapshot symlink contains NUL: {path}"))?; + let mut buffer = vec![0u8; MAX_SNAPSHOT_PATH_BYTES + 1]; + let length = unsafe { libc::readlinkat(parent, name.as_ptr(), buffer.as_mut_ptr().cast(), buffer.len()) }; + if length < 0 { + return Err(format!("read snapshot symlink {path}: {}", io::Error::last_os_error())); + } + let length = usize::try_from(length).map_err(|_| format!("snapshot symlink target length is invalid: {path}"))?; + if length >= buffer.len() { + return Err(format!("snapshot symlink target is too long: {path}")); + } + std::str::from_utf8(&buffer[..length]).map(str::to_string).map_err(|_| format!("snapshot symlink target is not UTF-8: {path}")) +} + +fn validate_symlink_target_syntax(target: &str, max_path: usize, max_component: usize, max_depth: usize) -> Result<(), String> { + if target.is_empty() || target.starts_with('/') || target.contains('\\') || target.len() > max_path || target.as_bytes().contains(&0) { + return Err(format!("invalid snapshot symlink target: {target}")); + } + let components = target.split('/').collect::>(); + if components.len() > max_depth { + return Err(format!("snapshot symlink target exceeds maximum depth: {target}")); + } + for (index, component) in components.iter().enumerate() { + if component.is_empty() { + if index + 1 == components.len() { + continue; + } + return Err(format!("snapshot symlink target contains an empty component: {target}")); + } + if *component != "." && *component != ".." { + validate_component(component, max_component)?; + } + } + Ok(()) +} + +fn normalize_symlink_target(link_path: &str, target: &str) -> Result { + let mut components = link_path.rsplit_once('/').map_or_else(Vec::new, |(parent, _)| parent.split('/').collect::>()); + validate_symlink_target_syntax(target, MAX_SNAPSHOT_PATH_BYTES, MAX_SNAPSHOT_COMPONENT_BYTES, MAX_SNAPSHOT_DEPTH)?; + for component in target.split('/') { + match component { + "" | "." => {} + ".." => { + components.pop().ok_or_else(|| format!("snapshot symlink target escapes workspace: {link_path} -> {target}"))?; + } + value => components.push(value), + } + } + Ok(components.join("/")) +} + +fn validate_symlink_target(canonical_root: &Path, link_path: &str, target: &str, limits: &SnapshotLimits) -> Result<(), String> { + let _normalized = normalize_symlink_target(link_path, target)?; + let parent = link_path.rsplit_once('/').map_or(Path::new(""), |(parent, _)| Path::new(parent)); + let candidate = canonical_root.join(parent).join(target); + let resolved = fs::canonicalize(&candidate).map_err(|error| format!("resolve snapshot symlink {link_path}: {error}"))?; + if resolved.strip_prefix(canonical_root).is_err() { + return Err(format!("snapshot symlink target escapes workspace: {link_path} -> {target}")); + } + validate_symlink_target_syntax(target, limits.max_path_bytes, limits.max_component_bytes, limits.max_depth) +} + +fn validate_snapshot_structure(entries: &[SnapshotEntry]) -> Result<(), String> { + let kinds = entries.iter().map(|entry| (entry.path.as_str(), entry.kind)).collect::>(); + for entry in entries { + let mut prefix = String::new(); + for component in entry.path.as_str().split('/').take(entry.path.as_str().split('/').count().saturating_sub(1)) { + if !prefix.is_empty() { + prefix.push('/'); + } + prefix.push_str(component); + if kinds.get(prefix.as_str()) != Some(&SnapshotEntryKind::Directory) { + return Err(format!("snapshot entry is missing directory parent: {prefix}")); + } + } + } + Ok(()) +} + +fn validate_snapshot_symlinks(entries: &[SnapshotEntry]) -> Result<(), String> { + let by_path = entries.iter().map(|entry| (entry.path.as_str(), entry)).collect::>(); + for entry in entries.iter().filter(|entry| entry.kind == SnapshotEntryKind::Symlink) { + let target = entry.symlink_target.as_deref().ok_or_else(|| format!("symlink snapshot entry has no target: {}", entry.path.as_str()))?; + let mut current = normalize_symlink_target(entry.path.as_str(), target)?; + let mut visited = BTreeSet::new(); + loop { + if current.is_empty() { + break; + } + let target_entry = by_path + .get(current.as_str()) + .ok_or_else(|| format!("snapshot symlink target is missing from snapshot: {} -> {target}", entry.path.as_str()))?; + if target_entry.kind != SnapshotEntryKind::Symlink { + break; + } + if !visited.insert(current.clone()) { + return Err(format!("snapshot symlink loop detected at: {}", entry.path.as_str())); + } + let nested_target = target_entry + .symlink_target + .as_deref() + .ok_or_else(|| format!("symlink snapshot entry has no target: {}", target_entry.path.as_str()))?; + current = normalize_symlink_target(target_entry.path.as_str(), nested_target)?; + } + } + Ok(()) +} + fn child_kind(stat: &libc::stat, path: &str) -> Result { match stat.st_mode & libc::S_IFMT { libc::S_IFDIR => Ok(SnapshotEntryKind::Directory), libc::S_IFREG => Ok(SnapshotEntryKind::RegularFile), - libc::S_IFLNK => Err(format!("snapshot rejects symlink: {path}")), + libc::S_IFLNK => Ok(SnapshotEntryKind::Symlink), _ => Err(format!("snapshot rejects special file: {path}")), } } @@ -875,6 +1011,23 @@ fn create_relative_file(root: &File, relative: &str, mode: u16) -> Result Result<(), String> { + let _normalized_target = normalize_symlink_target(relative, target)?; + let mut components = relative.split('/').collect::>(); + let file_name = components.pop().ok_or_else(|| "empty materialization symlink path".to_string())?; + let mut parent = root.try_clone().map_err(|error| format!("clone materialization symlink root: {error}"))?; + for component in components { + parent = open_at(parent.as_raw_fd(), OsStr::new(component), libc::O_RDONLY | libc::O_DIRECTORY | libc::O_NOFOLLOW | libc::O_CLOEXEC) + .map_err(|error| format!("open materialized symlink parent {relative}: {error}"))?; + } + let file_name = CString::new(file_name.as_bytes()).map_err(|_| "NUL in materialization symlink path".to_string())?; + let target = CString::new(target.as_bytes()).map_err(|_| format!("NUL in materialization symlink target: {relative}"))?; + if unsafe { libc::symlinkat(target.as_ptr(), parent.as_raw_fd(), file_name.as_ptr()) } != 0 { + return Err(format!("create materialized symlink {relative}: {}", io::Error::last_os_error())); + } + Ok(()) +} + fn open_relative_file(root: &File, relative: &str, flags: i32) -> Result { let mut components = relative.split('/').collect::>(); let file_name = components.pop().ok_or_else(|| "empty snapshot content path".to_string())?; @@ -995,6 +1148,8 @@ struct StoredEntry { mode: u16, size: u64, content_digest: Option<[u8; 32]>, + #[serde(default)] + symlink_target: Option, } impl StoredSnapshot { @@ -1009,6 +1164,7 @@ impl StoredSnapshot { mode: entry.mode, size: entry.size, content_digest: entry.content_digest, + symlink_target: entry.symlink_target.clone(), }) .collect(), total_file_bytes, @@ -1030,24 +1186,44 @@ impl StoredSnapshot { return Err("snapshot manifest entries are not strictly ordered".to_string()); } previous_path = Some(path.as_str().to_string()); - if matches!(entry.kind, SnapshotEntryKind::Directory) && (entry.size != 0 || entry.content_digest.is_some()) { - return Err("invalid directory snapshot entry".to_string()); - } - if matches!(entry.kind, SnapshotEntryKind::RegularFile) && entry.content_digest.is_none() { - return Err("invalid regular-file snapshot entry".to_string()); - } - if matches!(entry.kind, SnapshotEntryKind::RegularFile) { - if entry.size > MAX_SNAPSHOT_FILE_BYTES { - return Err("stored snapshot file exceeds maximum size".to_string()); + match entry.kind { + SnapshotEntryKind::Directory => { + if entry.size != 0 || entry.content_digest.is_some() || entry.symlink_target.is_some() { + return Err("invalid directory snapshot entry".to_string()); + } } - total_file_bytes = total_file_bytes.checked_add(entry.size).ok_or_else(|| "snapshot manifest size overflow".to_string())?; - if total_file_bytes > MAX_SNAPSHOT_TOTAL_BYTES { - return Err("stored snapshot exceeds maximum total size".to_string()); + SnapshotEntryKind::RegularFile => { + if entry.content_digest.is_none() || entry.symlink_target.is_some() { + return Err("invalid regular-file snapshot entry".to_string()); + } + if entry.size > MAX_SNAPSHOT_FILE_BYTES { + return Err("stored snapshot file exceeds maximum size".to_string()); + } + total_file_bytes = total_file_bytes.checked_add(entry.size).ok_or_else(|| "snapshot manifest size overflow".to_string())?; + if total_file_bytes > MAX_SNAPSHOT_TOTAL_BYTES { + return Err("stored snapshot exceeds maximum total size".to_string()); + } + } + SnapshotEntryKind::Symlink => { + let target = entry.symlink_target.as_deref().ok_or_else(|| "invalid symlink snapshot entry".to_string())?; + if entry.size != 0 || entry.content_digest.is_some() { + return Err("invalid symlink snapshot entry".to_string()); + } + validate_symlink_target_syntax(target, MAX_SNAPSHOT_PATH_BYTES, MAX_SNAPSHOT_COMPONENT_BYTES, MAX_SNAPSHOT_DEPTH)?; } } - Ok(SnapshotEntry { path, kind: entry.kind, mode: entry.mode & 0o777, size: entry.size, content_digest: entry.content_digest }) + Ok(SnapshotEntry { + path, + kind: entry.kind, + mode: entry.mode & 0o777, + size: entry.size, + content_digest: entry.content_digest, + symlink_target: entry.symlink_target, + }) }) .collect::, String>>()?; + validate_snapshot_structure(&entries)?; + validate_snapshot_symlinks(&entries)?; if total_file_bytes != self.total_file_bytes { return Err("snapshot manifest total size mismatch".to_string()); } @@ -1061,6 +1237,7 @@ fn snapshot_id(entries: &[SnapshotEntry]) -> SnapshotId { canonical.push(match entry.kind { SnapshotEntryKind::Directory => 0, SnapshotEntryKind::RegularFile => 1, + SnapshotEntryKind::Symlink => 2, }); canonical.extend_from_slice(&(entry.path.as_str().len() as u32).to_le_bytes()); canonical.extend_from_slice(entry.path.as_str().as_bytes()); @@ -1069,6 +1246,10 @@ fn snapshot_id(entries: &[SnapshotEntry]) -> SnapshotId { if let Some(digest) = entry.content_digest { canonical.extend_from_slice(&digest); } + if let Some(target) = entry.symlink_target() { + canonical.extend_from_slice(&(target.len() as u32).to_le_bytes()); + canonical.extend_from_slice(target.as_bytes()); + } } SnapshotId(Sha256::digest(canonical).into()) } diff --git a/src/snapshot_ut.rs b/src/snapshot_ut.rs index b9f408e..44c463c 100644 --- a/src/snapshot_ut.rs +++ b/src/snapshot_ut.rs @@ -168,28 +168,74 @@ fn malformed_exclusions_are_rejected() { } #[test] -fn every_symlink_is_rejected_without_following_it() { - let cases = ["internal-file", "internal-dir", "external", "dangling", "loop-a"]; - for case in cases { - let source = TempDir::new().unwrap(); - let store = TempDir::new().unwrap(); - write_file(source.path(), "real/file", b"data"); - match case { - "internal-file" => symlink("real/file", source.path().join(case)).unwrap(), - "internal-dir" => symlink("real", source.path().join(case)).unwrap(), - "external" => { - let outside = TempDir::new().unwrap(); - symlink(outside.path(), source.path().join(case)).unwrap(); - } - "dangling" => symlink("missing", source.path().join(case)).unwrap(), - "loop-a" => { - symlink("loop-b", source.path().join("loop-a")).unwrap(); - symlink("loop-a", source.path().join("loop-b")).unwrap(); - } - _ => unreachable!(), - } - assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err(), "{case}"); +fn internal_relative_symlink_is_accepted_and_preserved() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "Lib/module.py", b"module"); + fs::create_dir_all(source.path().join("crates/pylib")).unwrap(); + symlink("../../Lib/", source.path().join("crates/pylib/Lib")).unwrap(); + + let snapshot = build_at(&source, &store, SnapshotLimits::default(), &[]).unwrap(); + let entry = snapshot.entries().iter().find(|entry| entry.path().as_str() == "crates/pylib/Lib").unwrap(); + assert_eq!(entry.kind(), SnapshotEntryKind::Symlink); + assert_eq!(entry.symlink_target(), Some("../../Lib/")); + assert_eq!(entry.size(), 0); + assert!(entry.content_digest().is_none()); + + let destination = store.path().join("materialized"); + SnapshotStore::new(store.path()).materialize(snapshot.handle(), &destination).unwrap(); + let materialized_link = destination.join("crates/pylib/Lib"); + assert!(fs::symlink_metadata(&materialized_link).unwrap().file_type().is_symlink()); + assert_eq!(fs::read_link(&materialized_link).unwrap(), Path::new("../../Lib/")); + assert_eq!(fs::read(destination.join("crates/pylib/Lib/module.py")).unwrap(), b"module"); +} + +#[test] +fn nested_internal_symlinks_are_accepted_without_dereferencing() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + write_file(source.path(), "Lib/module.py", b"module"); + symlink("Lib", source.path().join("alias")).unwrap(); + fs::create_dir_all(source.path().join("nested/a")).unwrap(); + symlink("../../alias", source.path().join("nested/a/link")).unwrap(); + + let snapshot = build_at(&source, &store, SnapshotLimits::default(), &[]).unwrap(); + for (path, target) in [("alias", "Lib"), ("nested/a/link", "../../alias")] { + let entry = snapshot.entries().iter().find(|entry| entry.path().as_str() == path).unwrap(); + assert_eq!(entry.kind(), SnapshotEntryKind::Symlink); + assert_eq!(entry.symlink_target(), Some(target)); } + + let destination = store.path().join("materialized"); + SnapshotStore::new(store.path()).materialize(snapshot.handle(), &destination).unwrap(); + assert!(fs::symlink_metadata(destination.join("alias")).unwrap().file_type().is_symlink()); + assert!(fs::symlink_metadata(destination.join("nested/a/link")).unwrap().file_type().is_symlink()); + assert_eq!(fs::read(destination.join("nested/a/link/module.py")).unwrap(), b"module"); +} + +#[test] +fn unsafe_symlinks_are_rejected_without_following_them() { + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + symlink("/etc", source.path().join("absolute")).unwrap(); + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err()); + + let parent = TempDir::new().unwrap(); + let escaping_root = parent.path().join("root"); + let outside = parent.path().join("outside"); + fs::create_dir_all(&escaping_root).unwrap(); + fs::create_dir_all(&outside).unwrap(); + symlink("../outside", escaping_root.join("escape")).unwrap(); + assert!(builder(&store, SnapshotLimits::default(), &[]).build_root(&escaping_root, session(2)).is_err()); + + let source = TempDir::new().unwrap(); + let store = TempDir::new().unwrap(); + symlink("missing", source.path().join("broken")).unwrap(); + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err()); + + symlink("loop-b", source.path().join("loop-a")).unwrap(); + symlink("loop-a", source.path().join("loop-b")).unwrap(); + assert!(build_at(&source, &store, SnapshotLimits::default(), &[]).is_err()); } #[test] @@ -358,6 +404,19 @@ fn materialization_rejects_destination_symlink() { assert!(outside.path().read_dir().unwrap().next().is_none()); } +#[test] +fn materialization_rejects_symlink_replacement_in_a_parent_directory() { + let destination_parent = TempDir::new().unwrap(); + let destination = destination_parent.path().join("materialized"); + let outside = TempDir::new().unwrap(); + fs::create_dir(&destination).unwrap(); + symlink(outside.path(), destination.join("nested")).unwrap(); + + let root = open_directory(&destination).unwrap(); + assert!(create_relative_symlink(&root, "nested/link", "../target").is_err()); + assert!(!outside.path().join("link").exists()); +} + #[test] fn export_reads_validated_manifest_files_with_bounded_no_follow_access() { let source = TempDir::new().unwrap(); diff --git a/src/ssh.rs b/src/ssh.rs index 3a12147..9e0bf1c 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -9,7 +9,7 @@ use crate::snapshot::SnapshotEntryKind; use crate::worker_protocol::{ self, WorkerArtifactEntry, WorkerArtifactPath, WorkerArtifactSetId, WorkerBuild, WorkerErrorKind, WorkerMessage, WorkerOperation, WorkerRelativePath, WorkerRequestId, WorkerSessionId, WorkerUploadEntry, WorkerUploadId, MAX_WORKER_CHUNK_BYTES, MAX_WORKER_FILE_BYTES, - WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_COMMAND_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, + WORKER_ARTIFACT_PROTOCOL_VERSION, WORKER_COMMAND_PROTOCOL_VERSION, WORKER_PROTOCOL_VERSION, WORKER_SYMLINK_PROTOCOL_VERSION, }; use rand::RngCore; use std::collections::BTreeMap; @@ -28,13 +28,14 @@ const MAX_SSH_DIAGNOSTIC_BYTES: usize = 16 * 1024; const PROCESS_REAP_TIMEOUT: Duration = Duration::from_secs(2); fn build_protocol_version(target: &SshTarget, artifacts: bool) -> u16 { - if target.compact_destination() { + let version = if target.compact_destination() { WORKER_COMMAND_PROTOCOL_VERSION } else if artifacts { WORKER_ARTIFACT_PROTOCOL_VERSION } else { WORKER_PROTOCOL_VERSION - } + }; + version.max(WORKER_SYMLINK_PROTOCOL_VERSION) } type WorkerReader = Box; @@ -371,7 +372,7 @@ impl SshExecution { let mut connection = match phase( async { let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; - connection.handshake(WorkerRequestId(request_id), session_id_for(&self.session), WORKER_PROTOCOL_VERSION).await?; + connection.handshake(WorkerRequestId(request_id), session_id_for(&self.session), build_protocol_version(&self.target, false)).await?; Ok(connection) }, remaining(sync_deadline).min(self.target.resources().connect_timeout()), @@ -1085,6 +1086,12 @@ fn worker_entry(entry: &crate::snapshot::SnapshotEntry) -> Result WorkerUploadEntry::symlink( + entry.path().as_str(), + entry.mode() as u32, + entry.symlink_target().ok_or_else(|| worker_protocol("symlink snapshot entry has no target"))?, + ) + .map_err(worker_io_error), } } @@ -1093,7 +1100,7 @@ fn preflight_upload( ) -> Result<(), RemoteBackendError> { worker_protocol::validate_upload_manifest(entries).map_err(|error| upload_preflight_error(error.to_string()))?; WorkerMessage::UploadBegin { request_id: WorkerRequestId(request_id), session_id, upload_id, entries: entries.to_vec() } - .encode() + .encode_version(WORKER_SYMLINK_PROTOCOL_VERSION) .map_err(|error| upload_preflight_error(error.to_string()))?; Ok(()) } diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index 1a77428..c1faac8 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -158,13 +158,14 @@ where let Ok((_, first)) = worker_protocol::read_message_versioned(&mut reader).await else { return 72 }; messages.lock().unwrap().push(first.clone()); if matches!(mode, ScriptMode::ProtocolError) { - worker_protocol::write_message( + worker_protocol::write_message_versioned( &mut writer, &WorkerMessage::error(request_id, session_id, WorkerOperation::Upload, WorkerErrorKind::WorkerProtocol, "scripted protocol failure"), + hello_version, ) .await .unwrap(); - while worker_protocol::read_message(&mut reader).await.is_ok() {} + while worker_protocol::read_message_versioned(&mut reader).await.is_ok() {} return 0; } match first { From f173ca6db9cc4c00849654fe790c8ba2902d532a Mon Sep 17 00:00:00 2001 From: Bo Maryniuk Date: Thu, 20 Aug 2026 13:19:14 +0200 Subject: [PATCH 52/52] Fix absolute paths of the workspace --- src/daemon.rs | 13 +++++++- src/daemon_ut.rs | 27 +++++++++++++++- src/ssh.rs | 77 ++++++++++++++++++++++++++++++++++----------- src/ssh_ut.rs | 82 ++++++++++++++++++++++++++++++++++++++++++++++-- 4 files changed, 176 insertions(+), 23 deletions(-) diff --git a/src/daemon.rs b/src/daemon.rs index bd08799..8246b6a 100644 --- a/src/daemon.rs +++ b/src/daemon.rs @@ -1147,7 +1147,10 @@ async fn dispatch_local_remote_request( .await; } let exec_request = ExecRequest { - cwd: build.cwd().as_str().to_string(), + cwd: guest_workspace_cwd(build.cwd()) + .into_os_string() + .into_string() + .map_err(|_| "local workspace cwd is not valid UTF-8".to_string())?, command: build.tool().as_str().to_string(), args: build.argv().to_vec(), env: build.env().to_vec(), @@ -1193,6 +1196,14 @@ async fn dispatch_local_remote_request( } } +fn guest_workspace_cwd(cwd: &WorkspaceRelativePath) -> PathBuf { + if cwd.as_str().is_empty() { + PathBuf::from("/workspace") + } else { + Path::new("/workspace").join(cwd.as_str()) + } +} + async fn write_local_remote_event( writer: &mut W, request_id: crate::remote::RequestId, event: RemoteBackendEvent, ) -> Result<(), String> { diff --git a/src/daemon_ut.rs b/src/daemon_ut.rs index 2046689..8726939 100644 --- a/src/daemon_ut.rs +++ b/src/daemon_ut.rs @@ -528,6 +528,30 @@ fn local_exec_request_still_builds_on_the_local_path() { assert!(super::build_command(&session, &request, &cwd).is_ok()); } +#[test] +fn logical_workspace_root_maps_to_the_absolute_local_workspace_root() { + let cwd = super::guest_workspace_cwd(&WorkspaceRelativePath::new("").unwrap()); + assert_eq!(cwd, Path::new("/workspace")); +} + +#[test] +fn logical_workspace_subdirectory_maps_to_an_absolute_local_workspace_path() { + let cwd = super::guest_workspace_cwd(&WorkspaceRelativePath::new("nested/project").unwrap()); + assert_eq!(cwd, Path::new("/workspace/nested/project")); +} + +#[test] +fn empty_logical_workspace_cwd_never_becomes_an_empty_local_executor_cwd() { + let cwd = super::guest_workspace_cwd(&WorkspaceRelativePath::new("").unwrap()); + assert!(!cwd.as_os_str().is_empty()); +} + +#[test] +fn escaping_logical_workspace_cwd_remains_rejected() { + assert!(WorkspaceRelativePath::new("../escape").is_err()); + assert!(WorkspaceRelativePath::new("/workspace/escape").is_err()); +} + #[tokio::test] async fn transparent_exec_request_routes_fresh_managed_cargo_to_selected_remote_target() { let (root, catalog) = target_catalog_fixture(); @@ -555,7 +579,7 @@ async fn transparent_exec_request_routes_fresh_managed_cargo_to_selected_remote_ } #[tokio::test] -async fn managed_remote_wrapper_requests_route_to_target_selected_after_daemon_setup() { +async fn managed_remote_wrapper_requests_with_bsdbox_selected_never_enter_local_cwd_conversion() { let (root, catalog) = target_catalog_fixture(); let active = ActiveBuildTarget::new(); let backend = Arc::new(TargetRecordingBackend { calls: Mutex::new(Vec::new()) }); @@ -590,6 +614,7 @@ async fn managed_remote_wrapper_requests_route_to_target_selected_after_daemon_s assert_eq!(calls.len(), 2); assert_eq!(calls[0].target(), target); assert_eq!(calls[1].target(), target); + assert!(session.local_capabilities.lock().unwrap().is_empty()); } #[tokio::test] diff --git a/src/ssh.rs b/src/ssh.rs index 9e0bf1c..96f19bc 100644 --- a/src/ssh.rs +++ b/src/ssh.rs @@ -370,11 +370,14 @@ impl SshExecution { } self.control.set_phase(crate::remote::RemoteLifecyclePhase::Connecting); let mut connection = match phase( - async { - let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; - connection.handshake(WorkerRequestId(request_id), session_id_for(&self.session), build_protocol_version(&self.target, false)).await?; - Ok(connection) - }, + WorkerConnection::connect( + &self.factory, + &self.target, + WorkerRequestId(request_id), + session_id_for(&self.session), + build_protocol_version(&self.target, false), + self.target.resources().cleanup_timeout(), + ), remaining(sync_deadline).min(self.target.resources().connect_timeout()), &self.control, crate::remote::RemoteTimeoutCause::Connection, @@ -408,7 +411,7 @@ impl SshExecution { drop(export); if let Err(error) = result { - connection.kill_and_reap(self.target.resources().cleanup_timeout()).await; + let error = connection.capture_early_failure(error, self.target.resources().cleanup_timeout()).await; self.cleanup_upload(request_id, upload_id).await; let _ = self.session.abort_snapshot_capability(snapshot_id); return Err(error); @@ -523,17 +526,14 @@ impl SshExecution { self.control.set_phase(crate::remote::RemoteLifecyclePhase::Connecting); let mut connection = match phase( - async { - let mut connection = WorkerConnection::spawn(&self.factory, &self.target)?; - connection - .handshake( - WorkerRequestId(request_id), - WorkerSessionId(self.session.session_id().0), - build_protocol_version(&self.target, self.artifact_policy.is_enabled()), - ) - .await?; - Ok(connection) - }, + WorkerConnection::connect( + &self.factory, + &self.target, + WorkerRequestId(request_id), + WorkerSessionId(self.session.session_id().0), + build_protocol_version(&self.target, self.artifact_policy.is_enabled()), + self.target.resources().cleanup_timeout(), + ), self.target.resources().connect_timeout(), &self.control, crate::remote::RemoteTimeoutCause::Connection, @@ -579,7 +579,7 @@ impl SshExecution { } else { false }; - connection.kill_and_reap(self.target.resources().cleanup_timeout()).await; + let error = connection.capture_early_failure(error, self.target.resources().cleanup_timeout()).await; if !cleanup_succeeded { self.cleanup_upload(request_id, upload_id).await; } @@ -651,6 +651,17 @@ impl WorkerConnection { Ok(Self { process, writer: Some(writer), reader, stderr_task: Some(stderr_task), version: WORKER_PROTOCOL_VERSION }) } + async fn connect( + factory: &Arc, target: &SshTarget, request_id: WorkerRequestId, session_id: WorkerSessionId, version: u16, + cleanup_timeout: Duration, + ) -> Result { + let mut connection = Self::spawn(factory, target)?; + match connection.handshake(request_id, session_id, version).await { + Ok(()) => Ok(connection), + Err(error) => Err(connection.capture_early_failure(error, cleanup_timeout).await), + } + } + async fn handshake(&mut self, request_id: WorkerRequestId, session_id: WorkerSessionId, version: u16) -> Result<(), RemoteBackendError> { self.version = version; self.write(&WorkerMessage::hello_for_version(request_id, session_id, false, version)).await?; @@ -750,6 +761,27 @@ impl WorkerConnection { task.abort(); } } + + async fn capture_early_failure(&mut self, error: RemoteBackendError, cleanup_timeout: Duration) -> RemoteBackendError { + self.writer.take(); + self.process.kill_group(); + let deadline = Instant::now() + cleanup_timeout.min(PROCESS_REAP_TIMEOUT); + let _ = timeout(remaining(deadline), self.process.wait()).await; + let diagnostic = self.collect_diagnostic(remaining(deadline)).await; + append_worker_diagnostic(error, &diagnostic) + } + + async fn collect_diagnostic(&mut self, duration: Duration) -> Vec { + let Some(mut task) = self.stderr_task.take() else { return Vec::new() }; + match timeout(duration, &mut task).await { + Ok(Ok(diagnostic)) => diagnostic, + Ok(Err(_)) => Vec::new(), + Err(_) => { + task.abort(); + Vec::new() + } + } + } } impl Drop for WorkerConnection { @@ -1171,6 +1203,15 @@ fn worker_io_error(error: worker_protocol::WorkerProtocolError) -> RemoteBackend } } +fn append_worker_diagnostic(error: RemoteBackendError, diagnostic: &[u8]) -> RemoteBackendError { + match error { + RemoteBackendError::Transport { class, message } if !diagnostic.is_empty() => { + RemoteBackendError::Transport { class, message: format!("{message}\nworker stderr: {}", String::from_utf8_lossy(diagnostic)) } + } + error => error, + } +} + async fn phase( future: F, duration: Duration, control: &RemoteExecutionControl, cause: crate::remote::RemoteTimeoutCause, ) -> Result diff --git a/src/ssh_ut.rs b/src/ssh_ut.rs index c1faac8..ea3144a 100644 --- a/src/ssh_ut.rs +++ b/src/ssh_ut.rs @@ -51,6 +51,8 @@ enum ScriptMode { Success, WrongVersion, ProtocolError, + ExitBeforeHello { diagnostic_bytes: usize }, + ExitDuringUpload { diagnostic_bytes: usize }, ArtifactSuccess, ArtifactManifestError, ArtifactTransferError, @@ -77,11 +79,11 @@ impl SshProcessFactory for ScriptedFactory { let (host_writer, worker_reader) = tokio::io::duplex(8192); let (worker_writer, host_reader) = tokio::io::duplex(8192); - let (host_stderr, _worker_stderr) = tokio::io::duplex(256); + let (host_stderr, worker_stderr) = tokio::io::duplex(MAX_SSH_DIAGNOSTIC_BYTES + 1024); let (status_tx, status_rx) = oneshot::channel(); let messages = self.messages.clone(); tokio::spawn(async move { - let status = scripted_worker(worker_reader, worker_writer, mode, messages).await; + let status = scripted_worker(worker_reader, worker_writer, worker_stderr, mode, messages).await; let _ = status_tx.send(status); }); @@ -133,11 +135,17 @@ impl SshProcess for ScriptedProcess { } } -async fn scripted_worker(mut reader: R, mut writer: W, mode: ScriptMode, messages: Arc>>) -> i32 +async fn scripted_worker(mut reader: R, mut writer: W, mut stderr: E, mode: ScriptMode, messages: Arc>>) -> i32 where R: AsyncRead + Unpin, W: AsyncWrite + Unpin, + E: AsyncWrite + Unpin, { + if let ScriptMode::ExitBeforeHello { diagnostic_bytes } = mode { + let _ = worker_protocol::read_message_versioned(&mut reader).await; + write_scripted_diagnostic(&mut stderr, diagnostic_bytes).await; + return 79; + } let Ok((hello_version, WorkerMessage::Hello { request_id, session_id, .. })) = worker_protocol::read_message_versioned(&mut reader).await else { return 71; }; @@ -157,6 +165,10 @@ where let Ok((_, first)) = worker_protocol::read_message_versioned(&mut reader).await else { return 72 }; messages.lock().unwrap().push(first.clone()); + if let ScriptMode::ExitDuringUpload { diagnostic_bytes } = mode { + write_scripted_diagnostic(&mut stderr, diagnostic_bytes).await; + return 80; + } if matches!(mode, ScriptMode::ProtocolError) { worker_protocol::write_message_versioned( &mut writer, @@ -302,6 +314,16 @@ where } } +async fn write_scripted_diagnostic(stderr: &mut W, bytes: usize) { + let mut diagnostic = Vec::new(); + while diagnostic.len() < bytes { + diagnostic.extend_from_slice(b"scripted worker diagnostic: missing worker dependency\n"); + } + diagnostic.truncate(bytes); + let _ = stderr.write_all(&diagnostic).await; + let _ = stderr.flush().await; +} + struct Fixture { _temp: TempDir, session: Arc, @@ -479,6 +501,60 @@ async fn worker_protocol_failure_is_terminal_without_local_fallback() { assert_eq!(fixture.session.snapshot_capability_count(), 0); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn early_handshake_failure_includes_worker_stderr() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::ExitBeforeHello { diagnostic_bytes: 96 }]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory); + let (tx, _rx) = tokio::sync::mpsc::channel(16); + + let error = backend.execute(authorize_sync(&fixture), tx).await.unwrap_err(); + let RemoteBackendError::Transport { class, message } = error else { panic!("expected transport failure") }; + assert_eq!(class, RemoteFailureClass::WorkerProtocol); + assert!(message.contains("invalid worker protocol: truncated worker frame header")); + assert!(message.contains("\nworker stderr: scripted worker diagnostic: missing worker dependency")); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn early_upload_failure_includes_worker_stderr() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::ExitDuringUpload { diagnostic_bytes: 96 }]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory); + let (tx, _rx) = tokio::sync::mpsc::channel(16); + + let error = backend.execute(authorize_sync(&fixture), tx).await.unwrap_err(); + let RemoteBackendError::Transport { message, .. } = error else { panic!("expected transport failure") }; + assert!(message.contains("worker I/O error") || message.contains("invalid worker protocol")); + assert!(message.contains("\nworker stderr: scripted worker diagnostic: missing worker dependency")); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn early_failure_without_worker_stderr_preserves_the_original_error() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::ExitBeforeHello { diagnostic_bytes: 0 }]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory); + let (tx, _rx) = tokio::sync::mpsc::channel(16); + + let error = backend.execute(authorize_sync(&fixture), tx).await.unwrap_err(); + let RemoteBackendError::Transport { class, message } = error else { panic!("expected transport failure") }; + assert_eq!(class, RemoteFailureClass::WorkerProtocol); + assert_eq!(message, "invalid worker protocol: truncated worker frame header"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn early_worker_stderr_diagnostic_is_bounded() { + let fixture = fixture(); + let factory = ScriptedFactory::new(vec![ScriptMode::ExitBeforeHello { diagnostic_bytes: MAX_SSH_DIAGNOSTIC_BYTES + 1024 }]); + let backend = SshBackend::new(fixture.session.clone(), fixture.ssh_target.clone()).unwrap().with_process_factory(factory); + let (tx, _rx) = tokio::sync::mpsc::channel(16); + + let error = backend.execute(authorize_sync(&fixture), tx).await.unwrap_err(); + let RemoteBackendError::Transport { message, .. } = error else { panic!("expected transport failure") }; + let diagnostic = message.split_once("\nworker stderr: ").map(|(_, diagnostic)| diagnostic).expect("worker diagnostic is present"); + assert_eq!(diagnostic.len(), MAX_SSH_DIAGNOSTIC_BYTES); + assert!(diagnostic.starts_with("scripted worker diagnostic: missing worker dependency")); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn artifact_capable_worker_manifest_is_fetched_verified_published_and_cleaned() { let fixture = fixture();