From bc144bce15a525563f50229c62b915d2108c9e1b Mon Sep 17 00:00:00 2001 From: leynos Date: Thu, 9 Jul 2026 18:45:26 +0200 Subject: [PATCH 1/5] Resolve Whitaker Dylint findings across the crate and its tests Prepare the codebase for adopting the Whitaker Dylint suite by fixing every finding it reports, without suppressions: - bumpy_road_function: extract build_agent_overrides and serialize_agent_override from build_overrides in the config loader, and append_file_entry/append_symlink_entry from append_non_directory_entry in the credential archive writer. - module_max_lines: split the protocol exec module into stdin_io and output_io sibling submodules; move the error module's tests to a sibling error_tests.rs; split the github tests into client_tests and property_tests submodules; split the ACP frame and protocol ACP test modules; move the configuration layer-precedence BDD steps into tests/bdd_config_steps/layer.rs. - module_must_have_inner_docs: add //! inner documentation to every unit-test module that lacked it. - no_expect_outside_tests: helper and fixture functions no longer panic; they return Result (io::Result, serde_json::Error, or PodbotError) and propagate failures, with the calling test bodies unwrapping as the test verdict. Mutex-guard recorders in test doubles now recover from poisoning via PoisonError::into_inner instead of panicking. - no_std_fs_operations: the Makefile audit-target test and the BDD config-loader helpers now perform fixture filesystem work through cap-std directory handles instead of ambient std::fs calls. Also resolve the Clippy too_many_arguments and needless_pass_by_value diagnostics introduced by the refactoring. --- src/api/configure_git_identity.rs | 13 +- src/api/repository_clone.rs | 2 + src/config/loader.rs | 55 ++- src/engine/connection/error_classification.rs | 2 + .../connection/exec/acp_frame_split_tests.rs | 118 ++++++ src/engine/connection/exec/acp_frame_tests.rs | 143 +------ .../connection/exec/acp_policy_tests.rs | 21 +- .../connection/exec/acp_runtime_bdd_tests.rs | 41 +- .../connection/exec/acp_runtime_tests.rs | 45 +- src/engine/connection/exec/helpers_tests.rs | 2 +- src/engine/connection/exec/host_io.rs | 2 + src/engine/connection/exec/protocol.rs | 279 +----------- .../connection/exec/protocol_acp_bdd_tests.rs | 17 +- .../exec/protocol_acp_forwarding_tests.rs | 17 +- .../exec/protocol_acp_masking_tests.rs | 142 +++++++ .../protocol_acp_policy_integration_tests.rs | 75 ++-- .../exec/protocol_acp_routing_tests.rs | 127 ++++++ .../connection/exec/protocol_acp_tests.rs | 400 ++++-------------- src/engine/connection/exec/protocol_output.rs | 132 ++++++ src/engine/connection/exec/protocol_stdin.rs | 179 ++++++++ src/engine/connection/exec/runtime_helpers.rs | 12 +- src/engine/connection/exec/session.rs | 2 + src/engine/connection/exec/terminal.rs | 2 + .../git_identity/container_configurator.rs | 18 +- src/engine/connection/repository_clone/mod.rs | 48 ++- src/engine/connection/tests.rs | 30 +- .../connection/upload_credentials/archive.rs | 60 ++- src/error.rs | 248 +---------- src/error_tests.rs | 242 +++++++++++ src/github/classify.rs | 2 + src/github/client_tests.rs | 187 ++++++++ src/github/property_tests.rs | 29 ++ src/github/tests.rs | 273 ++---------- tests/bdd_config_helpers.rs | 171 +------- tests/bdd_config_loader_helpers.rs | 6 +- tests/bdd_config_steps/layer.rs | 171 ++++++++ tests/bdd_hosting_config_loader_helpers.rs | 5 +- tests/make_audit_target.rs | 83 ++-- tests/test_utils.rs | 2 + 39 files changed, 1829 insertions(+), 1574 deletions(-) create mode 100644 src/engine/connection/exec/acp_frame_split_tests.rs create mode 100644 src/engine/connection/exec/protocol_acp_masking_tests.rs create mode 100644 src/engine/connection/exec/protocol_acp_routing_tests.rs create mode 100644 src/engine/connection/exec/protocol_output.rs create mode 100644 src/engine/connection/exec/protocol_stdin.rs create mode 100644 src/error_tests.rs create mode 100644 src/github/client_tests.rs create mode 100644 src/github/property_tests.rs create mode 100644 tests/bdd_config_steps/layer.rs diff --git a/src/api/configure_git_identity.rs b/src/api/configure_git_identity.rs index 5fb9a4af..cd19cf08 100644 --- a/src/api/configure_git_identity.rs +++ b/src/api/configure_git_identity.rs @@ -47,6 +47,8 @@ pub fn configure_container_git_identity PodbotResult { - let runtime = tokio::runtime::Runtime::new().expect("test requires a Tokio runtime"); + ) -> io::Result> { + let runtime = tokio::runtime::Runtime::new()?; let handle = runtime.handle().clone(); let params = GitIdentityParams { client: exec_client, @@ -152,7 +157,7 @@ mod tests { container_id, runtime_handle: &handle, }; - configure_container_git_identity(¶ms) + Ok(configure_container_git_identity(¶ms)) } #[test] @@ -163,6 +168,7 @@ mod tests { let exec_client = make_exec_client(0); let result = invoke(&host_runner, &exec_client, "sandbox-unit") + .expect("test requires a Tokio runtime") .expect("should succeed with Configured"); assert!( @@ -179,6 +185,7 @@ mod tests { let exec_client = MockExecClient::new(); // no exec calls expected let result = invoke(&host_runner, &exec_client, "sandbox-unit-2") + .expect("test requires a Tokio runtime") .expect("should succeed with NoneConfigured"); assert!( diff --git a/src/api/repository_clone.rs b/src/api/repository_clone.rs index 5656eb06..a3052fe7 100644 --- a/src/api/repository_clone.rs +++ b/src/api/repository_clone.rs @@ -188,6 +188,8 @@ impl AskpassPath { #[cfg(test)] mod tests { + //! Unit and property tests for repository-clone request value types. + use super::{AskpassPath, BranchName, RepositoryRef, WorkspacePath}; use crate::error::{ConfigError, PodbotError}; use proptest::prelude::*; diff --git a/src/config/loader.rs b/src/config/loader.rs index 03e256da..f32a6406 100644 --- a/src/config/loader.rs +++ b/src/config/loader.rs @@ -216,29 +216,42 @@ fn build_overrides(overrides: &ConfigOverrides) -> Result { json_overrides.insert("image".to_owned(), serde_json::Value::String(image.clone())); } - if overrides.agent_kind.is_some() || overrides.agent_mode.is_some() { - let mut agent = serde_json::Map::new(); - - if let Some(kind) = overrides.agent_kind { - agent.insert( - "kind".to_owned(), - serde_json::to_value(kind).map_err(|error| ConfigError::ParseError { - message: format!("failed to serialize agent kind override: {error}"), - })?, - ); - } - - if let Some(mode) = overrides.agent_mode { - agent.insert( - "mode".to_owned(), - serde_json::to_value(mode).map_err(|error| ConfigError::ParseError { - message: format!("failed to serialize agent mode override: {error}"), - })?, - ); - } - + if let Some(agent) = build_agent_overrides(overrides)? { json_overrides.insert("agent".to_owned(), serde_json::Value::Object(agent)); } Ok(serde_json::Value::Object(json_overrides)) } + +/// Build the nested agent override object, if any agent override is set. +fn build_agent_overrides( + overrides: &ConfigOverrides, +) -> Result>> { + if overrides.agent_kind.is_none() && overrides.agent_mode.is_none() { + return Ok(None); + } + + let mut agent = serde_json::Map::new(); + + if let Some(kind) = overrides.agent_kind { + agent.insert("kind".to_owned(), serialize_agent_override("kind", kind)?); + } + + if let Some(mode) = overrides.agent_mode { + agent.insert("mode".to_owned(), serialize_agent_override("mode", mode)?); + } + + Ok(Some(agent)) +} + +/// Serialize a single agent override field, labelling failures by field name. +fn serialize_agent_override( + field: &str, + value: T, +) -> Result { + Ok( + serde_json::to_value(value).map_err(|error| ConfigError::ParseError { + message: format!("failed to serialize agent {field} override: {error}"), + })?, + ) +} diff --git a/src/engine/connection/error_classification.rs b/src/engine/connection/error_classification.rs index 3411671f..f17ec2fc 100644 --- a/src/engine/connection/error_classification.rs +++ b/src/engine/connection/error_classification.rs @@ -113,6 +113,8 @@ fn io_error_kind_in_chain(error: &dyn std::error::Error) -> Option Result, serde_json::Error> { + let frames = [ + permitted_frame("session/new", b"\n")?, + permitted_frame("session/update", b"\n")?, + permitted_frame("session/cancel", b"\n")?, + permitted_frame("session/new", b"\r\n")?, + permitted_frame("session/update", b"\n")?, + ]; + Ok(frames.into_iter().fold(Vec::new(), |mut acc, mut frame| { + acc.append(&mut frame); + acc + })) +} + +fn assemble_with_two_chunks(stream: &[u8], split_at: usize) -> Vec { + let mut framer = assembler(); + let first = stream.get(..split_at).unwrap_or_default(); + let second = stream.get(split_at..).unwrap_or_default(); + let mut outputs = Vec::new(); + let (chunk_one, fallback_one) = framer.ingest_chunk(first); + assert!(fallback_one.is_none()); + outputs.extend(chunk_one); + let (chunk_two, fallback_two) = framer.ingest_chunk(second); + assert!(fallback_two.is_none()); + outputs.extend(chunk_two); + assert!(framer.finish().is_none()); + collect_forward_bytes(&outputs) +} + +fn assemble_with_three_chunks(stream: &[u8], first_split: usize, second_split: usize) -> Vec { + assert!(first_split <= second_split); + let mut framer = assembler(); + let first = stream.get(..first_split).unwrap_or_default(); + let second = stream.get(first_split..second_split).unwrap_or_default(); + let third = stream.get(second_split..).unwrap_or_default(); + let mut outputs = Vec::new(); + for chunk in [first, second, third] { + let (chunk_outputs, fallback) = framer.ingest_chunk(chunk); + assert!(fallback.is_none()); + outputs.extend(chunk_outputs); + } + assert!(framer.finish().is_none()); + collect_forward_bytes(&outputs) +} + +#[test] +fn every_two_way_split_reassembles_to_original_byte_stream() { + let stream = permitted_stream().expect("stream should serialize"); + for split_at in 1..stream.len() { + let reassembled = assemble_with_two_chunks(&stream, split_at); + assert_eq!( + reassembled, stream, + "split at byte {split_at} should reassemble byte-identically", + ); + } +} + +#[rstest] +#[case(1, 4)] +#[case(8, 32)] +#[case(16, 64)] +#[case(32, 96)] +#[case(40, 80)] +#[case(50, 120)] +#[case(60, 100)] +#[case(70, 140)] +#[case(80, 160)] +#[case(90, 150)] +#[case(95, 145)] +#[case(100, 200)] +#[case(110, 220)] +#[case(120, 180)] +#[case(125, 230)] +#[case(130, 240)] +#[case(135, 235)] +#[case(140, 250)] +#[case(150, 260)] +#[case(155, 265)] +#[case(160, 270)] +#[case(170, 280)] +#[case(180, 290)] +#[case(190, 300)] +#[case(200, 310)] +#[case(210, 320)] +#[case(220, 325)] +#[case(225, 330)] +#[case(230, 335)] +#[case(235, 340)] +#[case(240, 345)] +#[case(250, 350)] +fn three_way_splits_reassemble_to_original_byte_stream( + #[case] first_split: usize, + #[case] second_split: usize, +) { + let stream = permitted_stream().expect("stream should serialize"); + let first_clamped = first_split.min(stream.len()); + let second_clamped = second_split.min(stream.len()); + let reassembled = assemble_with_three_chunks(&stream, first_clamped, second_clamped); + assert_eq!( + reassembled, stream, + "three-way split at ({first_clamped}, {second_clamped}) should reassemble identically", + ); +} diff --git a/src/engine/connection/exec/acp_frame_tests.rs b/src/engine/connection/exec/acp_frame_tests.rs index b0e58155..b4b8dbd0 100644 --- a/src/engine/connection/exec/acp_frame_tests.rs +++ b/src/engine/connection/exec/acp_frame_tests.rs @@ -14,28 +14,26 @@ use super::{ }; use crate::engine::connection::exec::acp_policy::MethodDenylist; -fn permitted_frame(method: &str, line_ending: &[u8]) -> Vec { +fn permitted_frame(method: &str, line_ending: &[u8]) -> Result, serde_json::Error> { let mut bytes = serde_json::to_vec(&serde_json::json!({ "jsonrpc": "2.0", "id": 1, "method": method, "params": {}, - })) - .expect("frame serializes"); + }))?; bytes.extend_from_slice(line_ending); - bytes + Ok(bytes) } -fn blocked_request_frame(id: &Value, method: &str) -> Vec { +fn blocked_request_frame(id: &Value, method: &str) -> Result, serde_json::Error> { let mut bytes = serde_json::to_vec(&serde_json::json!({ "jsonrpc": "2.0", "id": id, "method": method, "params": {}, - })) - .expect("frame serializes"); + }))?; bytes.push(b'\n'); - bytes + Ok(bytes) } fn assembler() -> OutboundFrameAssembler { @@ -58,7 +56,7 @@ fn collect_forward_bytes(outputs: &[FrameOutput]) -> Vec { #[test] fn single_permitted_frame_is_forwarded_verbatim() { let mut framer = assembler(); - let frame = permitted_frame("session/new", b"\n"); + let frame = permitted_frame("session/new", b"\n").expect("frame should serialize"); let (outputs, fallback) = framer.ingest_chunk(&frame); @@ -70,8 +68,9 @@ fn single_permitted_frame_is_forwarded_verbatim() { #[test] fn multiple_frames_in_one_chunk_split_on_each_newline() { let mut framer = assembler(); - let mut chunk = permitted_frame("session/new", b"\n"); - chunk.extend_from_slice(&permitted_frame("session/update", b"\n")); + let mut chunk = permitted_frame("session/new", b"\n").expect("frame should serialize"); + let second_frame = permitted_frame("session/update", b"\n").expect("frame should serialize"); + chunk.extend_from_slice(&second_frame); let (outputs, fallback) = framer.ingest_chunk(&chunk); @@ -83,7 +82,7 @@ fn multiple_frames_in_one_chunk_split_on_each_newline() { #[test] fn frame_split_across_two_chunks_reassembles_correctly() { let mut framer = assembler(); - let frame = permitted_frame("session/new", b"\n"); + let frame = permitted_frame("session/new", b"\n").expect("frame should serialize"); let split_at = frame.len().div_euclid(2); let first = frame.get(..split_at).expect("split prefix"); let second = frame.get(split_at..).expect("split suffix"); @@ -102,7 +101,7 @@ fn frame_split_across_two_chunks_reassembles_correctly() { #[test] fn frame_split_across_three_chunks_reassembles_correctly() { let mut framer = assembler(); - let frame = permitted_frame("session/update", b"\n"); + let frame = permitted_frame("session/update", b"\n").expect("frame should serialize"); let third = frame.len().div_euclid(3); let two_thirds = frame.len().saturating_mul(2).div_euclid(3); let parts = [ @@ -124,7 +123,8 @@ fn frame_split_across_three_chunks_reassembles_correctly() { #[test] fn blocked_request_emits_decision_with_line_ending() { let mut framer = assembler(); - let frame = blocked_request_frame(&serde_json::json!(7), "terminal/create"); + let frame = blocked_request_frame(&serde_json::json!(7), "terminal/create") + .expect("frame should serialize"); let (outputs, fallback) = framer.ingest_chunk(&frame); @@ -210,8 +210,9 @@ fn frame_with_escaped_newline_in_string_treated_as_single_frame() { #[test] fn permitted_frame_after_blocked_frame_still_forwards() { let mut framer = assembler(); - let mut chunk = blocked_request_frame(&serde_json::json!(1), "terminal/create"); - let permitted = permitted_frame("session/new", b"\n"); + let mut chunk = blocked_request_frame(&serde_json::json!(1), "terminal/create") + .expect("frame should serialize"); + let permitted = permitted_frame("session/new", b"\n").expect("frame should serialize"); chunk.extend_from_slice(&permitted); let (outputs, _) = framer.ingest_chunk(&chunk); @@ -276,7 +277,7 @@ fn finish_drops_residual_partial_frame_and_reports_byte_count() { #[test] fn finish_returns_none_when_buffer_empty() { let mut framer = assembler(); - let frame = permitted_frame("session/new", b"\n"); + let frame = permitted_frame("session/new", b"\n").expect("frame should serialize"); let _ = framer.ingest_chunk(&frame); assert!(framer.finish().is_none()); @@ -299,109 +300,5 @@ fn empty_chunk_produces_no_output() { assert!(fallback.is_none()); } -/// Build a deterministic byte sequence containing several permitted ACP -/// frames. The sequence is reused across the exhaustive-split parameterized -/// tests below. -fn permitted_stream() -> Vec { - let frames = [ - permitted_frame("session/new", b"\n"), - permitted_frame("session/update", b"\n"), - permitted_frame("session/cancel", b"\n"), - permitted_frame("session/new", b"\r\n"), - permitted_frame("session/update", b"\n"), - ]; - frames.into_iter().fold(Vec::new(), |mut acc, mut frame| { - acc.append(&mut frame); - acc - }) -} - -fn assemble_with_two_chunks(stream: &[u8], split_at: usize) -> Vec { - let mut framer = assembler(); - let first = stream.get(..split_at).unwrap_or_default(); - let second = stream.get(split_at..).unwrap_or_default(); - let mut outputs = Vec::new(); - let (chunk_one, fallback_one) = framer.ingest_chunk(first); - assert!(fallback_one.is_none()); - outputs.extend(chunk_one); - let (chunk_two, fallback_two) = framer.ingest_chunk(second); - assert!(fallback_two.is_none()); - outputs.extend(chunk_two); - assert!(framer.finish().is_none()); - collect_forward_bytes(&outputs) -} - -fn assemble_with_three_chunks(stream: &[u8], first_split: usize, second_split: usize) -> Vec { - assert!(first_split <= second_split); - let mut framer = assembler(); - let first = stream.get(..first_split).unwrap_or_default(); - let second = stream.get(first_split..second_split).unwrap_or_default(); - let third = stream.get(second_split..).unwrap_or_default(); - let mut outputs = Vec::new(); - for chunk in [first, second, third] { - let (chunk_outputs, fallback) = framer.ingest_chunk(chunk); - assert!(fallback.is_none()); - outputs.extend(chunk_outputs); - } - assert!(framer.finish().is_none()); - collect_forward_bytes(&outputs) -} - -#[test] -fn every_two_way_split_reassembles_to_original_byte_stream() { - let stream = permitted_stream(); - for split_at in 1..stream.len() { - let reassembled = assemble_with_two_chunks(&stream, split_at); - assert_eq!( - reassembled, stream, - "split at byte {split_at} should reassemble byte-identically", - ); - } -} - -#[rstest] -#[case(1, 4)] -#[case(8, 32)] -#[case(16, 64)] -#[case(32, 96)] -#[case(40, 80)] -#[case(50, 120)] -#[case(60, 100)] -#[case(70, 140)] -#[case(80, 160)] -#[case(90, 150)] -#[case(95, 145)] -#[case(100, 200)] -#[case(110, 220)] -#[case(120, 180)] -#[case(125, 230)] -#[case(130, 240)] -#[case(135, 235)] -#[case(140, 250)] -#[case(150, 260)] -#[case(155, 265)] -#[case(160, 270)] -#[case(170, 280)] -#[case(180, 290)] -#[case(190, 300)] -#[case(200, 310)] -#[case(210, 320)] -#[case(220, 325)] -#[case(225, 330)] -#[case(230, 335)] -#[case(235, 340)] -#[case(240, 345)] -#[case(250, 350)] -fn three_way_splits_reassemble_to_original_byte_stream( - #[case] first_split: usize, - #[case] second_split: usize, -) { - let stream = permitted_stream(); - let first_clamped = first_split.min(stream.len()); - let second_clamped = second_split.min(stream.len()); - let reassembled = assemble_with_three_chunks(&stream, first_clamped, second_clamped); - assert_eq!( - reassembled, stream, - "three-way split at ({first_clamped}, {second_clamped}) should reassemble identically", - ); -} +#[path = "acp_frame_split_tests.rs"] +mod split_tests; diff --git a/src/engine/connection/exec/acp_policy_tests.rs b/src/engine/connection/exec/acp_policy_tests.rs index 64a5ab2b..1b041602 100644 --- a/src/engine/connection/exec/acp_policy_tests.rs +++ b/src/engine/connection/exec/acp_policy_tests.rs @@ -42,13 +42,13 @@ fn default_denylist_blocks_terminal_and_fs_only(#[case] method: &str, #[case] ex /// Serializes `value` to compact JSON bytes and appends a newline terminator, /// producing a well-formed ACP frame suitable for test input. -fn serialize_frame(value: &serde_json::Value) -> Vec { - let mut bytes = serde_json::to_vec(value).expect("frame serializes"); +fn serialize_frame(value: &serde_json::Value) -> Result, serde_json::Error> { + let mut bytes = serde_json::to_vec(value)?; bytes.push(b'\n'); - bytes + Ok(bytes) } -fn jsonrpc_request(id: &Value, method: &str) -> Vec { +fn jsonrpc_request(id: &Value, method: &str) -> Result, serde_json::Error> { serialize_frame(&serde_json::json!({ "jsonrpc": "2.0", "id": id, @@ -57,7 +57,7 @@ fn jsonrpc_request(id: &Value, method: &str) -> Vec { })) } -fn jsonrpc_notification(method: &str) -> Vec { +fn jsonrpc_notification(method: &str) -> Result, serde_json::Error> { serialize_frame(&serde_json::json!({ "jsonrpc": "2.0", "method": method, @@ -70,7 +70,7 @@ fn jsonrpc_notification(method: &str) -> Vec { #[case::string_id(serde_json::json!("call-7"))] #[case::null_id(serde_json::json!(null))] fn evaluate_returns_block_request_with_preserved_id(#[case] id: Value) { - let frame = jsonrpc_request(&id, "terminal/create"); + let frame = jsonrpc_request(&id, "terminal/create").expect("frame should serialize"); let denylist = MethodDenylist::default_families(); let decision = evaluate_agent_outbound_frame(&frame, &denylist); @@ -89,7 +89,7 @@ fn evaluate_returns_block_request_with_preserved_id(#[case] id: Value) { #[test] fn evaluate_returns_block_notification_for_blocked_method_without_id() { - let frame = jsonrpc_notification("fs/changed"); + let frame = jsonrpc_notification("fs/changed").expect("frame should serialize"); let denylist = MethodDenylist::default_families(); let decision = evaluate_agent_outbound_frame(&frame, &denylist); @@ -104,7 +104,8 @@ fn evaluate_returns_block_notification_for_blocked_method_without_id() { #[test] fn evaluate_forwards_permitted_request() { - let frame = jsonrpc_request(&serde_json::json!(1), "session/new"); + let frame = + jsonrpc_request(&serde_json::json!(1), "session/new").expect("frame should serialize"); let denylist = MethodDenylist::default_families(); assert_eq!( @@ -137,7 +138,7 @@ fn evaluate_forwards_malformed_json() { "method": 7, }))] fn evaluate_forwards_frames_without_dispatchable_method(#[case] payload: serde_json::Value) { - let frame = serialize_frame(&payload); + let frame = serialize_frame(&payload).expect("frame should serialize"); let denylist = MethodDenylist::default_families(); assert_eq!( @@ -162,7 +163,7 @@ fn evaluate_forwards_frames_without_dispatchable_method(#[case] payload: serde_j "method": "terminal/create", }))] fn evaluate_forwards_non_jsonrpc_objects_with_blocked_method(#[case] payload: serde_json::Value) { - let frame = serialize_frame(&payload); + let frame = serialize_frame(&payload).expect("frame should serialize"); let denylist = MethodDenylist::default_families(); assert_eq!( diff --git a/src/engine/connection/exec/acp_runtime_bdd_tests.rs b/src/engine/connection/exec/acp_runtime_bdd_tests.rs index a65760d7..15fa4f2f 100644 --- a/src/engine/connection/exec/acp_runtime_bdd_tests.rs +++ b/src/engine/connection/exec/acp_runtime_bdd_tests.rs @@ -66,16 +66,17 @@ fn denylist_state() -> DenylistState { DenylistState::default() } -/// Serialises `value` to compact JSON bytes and appends a newline terminator, +/// Serializes `value` to compact JSON bytes and appends a newline terminator, /// producing a well-formed ACP frame suitable for test input. -fn serialize_frame(value: serde_json::Value) -> Vec { - let mut bytes = serde_json::to_vec(&value).expect("frame serializes"); +fn serialize_frame(value: serde_json::Value) -> StepResult> { + let mut bytes = + serde_json::to_vec(&value).map_err(|err| format!("frame serialization failed: {err}"))?; drop(value); bytes.push(b'\n'); - bytes + Ok(bytes) } -fn make_request_frame(method: &str, id: i64) -> Vec { +fn make_request_frame(method: &str, id: i64) -> StepResult> { serialize_frame(serde_json::json!({ "jsonrpc": "2.0", "id": id, @@ -84,7 +85,7 @@ fn make_request_frame(method: &str, id: i64) -> Vec { })) } -fn make_notification_frame(method: &str) -> Vec { +fn make_notification_frame(method: &str) -> StepResult> { serialize_frame(serde_json::json!({ "jsonrpc": "2.0", "method": method, @@ -147,16 +148,14 @@ where adapter .handle_chunk(&chunk, &mut writer) .await - .expect("chunk handles cleanly"); + .map_err(|err| format!("chunk handling failed: {err}"))?; } adapter.finish(); drop(adapter); let commands = drain_sink(receiver).await; - let bytes = host_stdout_handle - .snapshot() - .expect("host stdout snapshot should succeed"); - (commands, bytes) - }); + let bytes = host_stdout_handle.snapshot()?; + Ok::<_, String>((commands, bytes)) + })?; state.host_stdout_bytes.set(host_bytes); state.sink_commands.set(commands); Ok(()) @@ -169,25 +168,25 @@ fn adapter_uses_default_denylist(denylist_state: &DenylistState) { #[when(r#"the agent emits a "terminal/create" request with id 7"#)] fn emit_blocked_terminal_create(denylist_state: &DenylistState) -> StepResult<()> { - let frame = make_request_frame("terminal/create", 7); + let frame = make_request_frame("terminal/create", 7)?; run_runtime(denylist_state, vec![frame], || {}) } #[when(r#"the agent emits a "session/new" request with id 1"#)] fn emit_permitted_session_new(denylist_state: &DenylistState) -> StepResult<()> { - let frame = make_request_frame("session/new", 1); + let frame = make_request_frame("session/new", 1)?; run_runtime(denylist_state, vec![frame], || {}) } #[when(r#"the agent emits an "fs/changed" notification"#)] fn emit_blocked_notification(denylist_state: &DenylistState) -> StepResult<()> { - let frame = make_notification_frame("fs/changed"); + let frame = make_notification_frame("fs/changed")?; run_runtime(denylist_state, vec![frame], || {}) } #[when("the agent emits a blocked frame split across two output chunks")] fn emit_blocked_split(denylist_state: &DenylistState) -> StepResult<()> { - let frame = make_request_frame("terminal/create", 2); + let frame = make_request_frame("terminal/create", 2)?; let split_at = frame.len().div_euclid(2); let first = frame .get(..split_at) @@ -202,8 +201,8 @@ fn emit_blocked_split(denylist_state: &DenylistState) -> StepResult<()> { #[when("the agent emits a blocked request followed by a permitted request")] fn emit_blocked_then_permitted(denylist_state: &DenylistState) -> StepResult<()> { - let mut chunk = make_request_frame("terminal/create", 5); - let permitted = make_request_frame("session/update", 6); + let mut chunk = make_request_frame("terminal/create", 5)?; + let permitted = make_request_frame("session/update", 6)?; chunk.extend_from_slice(&permitted); run_runtime(denylist_state, vec![chunk], || {}) } @@ -279,14 +278,16 @@ fn expect_synthesized_id(denylist_state: &DenylistState, expected_id: &Value) -> #[then("host stdout receives the permitted frame verbatim")] fn assert_host_stdout_matches_permitted(denylist_state: &DenylistState) -> StepResult<()> { - assert_host_stdout_matches(denylist_state, &make_request_frame("session/new", 1)) + let expected = make_request_frame("session/new", 1)?; + assert_host_stdout_matches(denylist_state, &expected) } #[then("host stdout receives only the permitted frame verbatim")] fn assert_host_stdout_matches_permitted_after_blocked( denylist_state: &DenylistState, ) -> StepResult<()> { - assert_host_stdout_matches(denylist_state, &make_request_frame("session/update", 6)) + let expected = make_request_frame("session/update", 6)?; + assert_host_stdout_matches(denylist_state, &expected) } #[then("container stdin receives no synthesized response")] diff --git a/src/engine/connection/exec/acp_runtime_tests.rs b/src/engine/connection/exec/acp_runtime_tests.rs index 07945a94..18eff121 100644 --- a/src/engine/connection/exec/acp_runtime_tests.rs +++ b/src/engine/connection/exec/acp_runtime_tests.rs @@ -3,7 +3,7 @@ use std::io; use std::pin::Pin; -use std::sync::{Arc, Mutex}; +use std::sync::{Arc, Mutex, PoisonError}; use std::task::{Context, Poll}; use ortho_config::serde_json::{self, Value}; @@ -26,11 +26,17 @@ struct RecordingWriter { impl RecordingWriter { fn snapshot(&self) -> Vec { - self.bytes.lock().expect("writer mutex").clone() + self.bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() } fn shutdown_observed(&self) -> bool { - *self.shutdown_called.lock().expect("shutdown mutex") + *self + .shutdown_called + .lock() + .unwrap_or_else(PoisonError::into_inner) } } @@ -42,7 +48,7 @@ impl AsyncWrite for RecordingWriter { ) -> Poll> { self.bytes .lock() - .expect("writer mutex") + .unwrap_or_else(PoisonError::into_inner) .extend_from_slice(buf); Poll::Ready(Ok(buf.len())) } @@ -52,7 +58,10 @@ impl AsyncWrite for RecordingWriter { } fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - *self.shutdown_called.lock().expect("shutdown mutex") = true; + *self + .shutdown_called + .lock() + .unwrap_or_else(PoisonError::into_inner) = true; Poll::Ready(Ok(())) } } @@ -85,7 +94,7 @@ impl AsyncWrite for BrokenPipeWriter { /// Builds a newline-terminated JSON-RPC 2.0 frame. /// /// Pass `id = Some(…)` for requests; `id = None` for notifications. -fn make_jsonrpc_frame(method: &str, id: Option<&Value>) -> Vec { +fn make_jsonrpc_frame(method: &str, id: Option<&Value>) -> Result, serde_json::Error> { let value = id.map_or_else( || { serde_json::json!({ @@ -103,20 +112,20 @@ fn make_jsonrpc_frame(method: &str, id: Option<&Value>) -> Vec { }) }, ); - let mut bytes = serde_json::to_vec(&value).expect("frame serializes"); + let mut bytes = serde_json::to_vec(&value)?; bytes.push(b'\n'); - bytes + Ok(bytes) } -fn permitted_frame() -> Vec { +fn permitted_frame() -> Result, serde_json::Error> { make_jsonrpc_frame("session/new", Some(&serde_json::json!(1))) } -fn blocked_request_frame(id: &Value) -> Vec { +fn blocked_request_frame(id: &Value) -> Result, serde_json::Error> { make_jsonrpc_frame("terminal/create", Some(id)) } -fn blocked_notification_frame() -> Vec { +fn blocked_notification_frame() -> Result, serde_json::Error> { make_jsonrpc_frame("fs/changed", None) } @@ -142,7 +151,7 @@ async fn permitted_frame_writes_to_host_stdout_only() { let host_stdout = RecordingWriter::default(); let recorder = host_stdout.clone(); let mut writer: Pin> = Box::pin(host_stdout); - let frame = permitted_frame(); + let frame = permitted_frame().expect("frame should serialize"); adapter .handle_chunk(&frame, &mut writer) @@ -162,7 +171,7 @@ async fn blocked_request_skips_host_stdout_and_queues_synthesized_response() { let host_stdout = RecordingWriter::default(); let recorder = host_stdout.clone(); let mut writer: Pin> = Box::pin(host_stdout); - let frame = blocked_request_frame(&serde_json::json!(7)); + let frame = blocked_request_frame(&serde_json::json!(7)).expect("frame should serialize"); adapter .handle_chunk(&frame, &mut writer) @@ -199,7 +208,7 @@ async fn blocked_notification_drops_silently_without_sink_command() { let host_stdout = RecordingWriter::default(); let recorder = host_stdout.clone(); let mut writer: Pin> = Box::pin(host_stdout); - let frame = blocked_notification_frame(); + let frame = blocked_notification_frame().expect("frame should serialize"); adapter .handle_chunk(&frame, &mut writer) @@ -221,8 +230,8 @@ async fn permitted_frame_after_blocked_frame_still_reaches_host_stdout() { let host_stdout = RecordingWriter::default(); let recorder = host_stdout.clone(); let mut writer: Pin> = Box::pin(host_stdout); - let mut chunk = blocked_request_frame(&serde_json::json!(1)); - let permitted = permitted_frame(); + let mut chunk = blocked_request_frame(&serde_json::json!(1)).expect("frame should serialize"); + let permitted = permitted_frame().expect("frame should serialize"); chunk.extend_from_slice(&permitted); adapter @@ -242,7 +251,7 @@ async fn frame_split_across_chunks_is_classified_after_assembly() { let host_stdout = RecordingWriter::default(); let recorder = host_stdout.clone(); let mut writer: Pin> = Box::pin(host_stdout); - let frame = blocked_request_frame(&serde_json::json!(2)); + let frame = blocked_request_frame(&serde_json::json!(2)).expect("frame should serialize"); let split_at = frame.len().div_euclid(2); let first = frame.get(..split_at).expect("split prefix"); let second = frame.get(split_at..).expect("split suffix"); @@ -355,7 +364,7 @@ async fn blocked_request_synthesized_before_channel_close_is_flushed() { let mut adapter = OutboundPolicyAdapter::new(assembler, tx.clone(), "container-test"); let host_stdout = RecordingWriter::default(); let mut host_writer: Pin> = Box::pin(host_stdout); - let frame = blocked_request_frame(&serde_json::json!(11)); + let frame = blocked_request_frame(&serde_json::json!(11)).expect("frame should serialize"); adapter .handle_chunk(&frame, &mut host_writer) diff --git a/src/engine/connection/exec/helpers_tests.rs b/src/engine/connection/exec/helpers_tests.rs index b783fb36..b89b47c0 100644 --- a/src/engine/connection/exec/helpers_tests.rs +++ b/src/engine/connection/exec/helpers_tests.rs @@ -191,7 +191,7 @@ impl AsyncWrite for SharedWriter { ) -> std::task::Poll> { self.bytes .lock() - .expect("writer buffer mutex should not be poisoned") + .unwrap_or_else(std::sync::PoisonError::into_inner) .extend_from_slice(buf); std::task::Poll::Ready(Ok(buf.len())) } diff --git a/src/engine/connection/exec/host_io.rs b/src/engine/connection/exec/host_io.rs index c166ace9..fe2ccd59 100644 --- a/src/engine/connection/exec/host_io.rs +++ b/src/engine/connection/exec/host_io.rs @@ -50,6 +50,8 @@ pub(super) fn default_host_stdin() -> Pin> { #[cfg(test)] mod tests { + //! Unit tests for host IO stdin-forwarding overrides. + use tokio::io::AsyncReadExt; use super::*; diff --git a/src/engine/connection/exec/protocol.rs b/src/engine/connection/exec/protocol.rs index da722f67..3f7520df 100644 --- a/src/engine/connection/exec/protocol.rs +++ b/src/engine/connection/exec/protocol.rs @@ -60,18 +60,14 @@ use std::io; use std::pin::Pin; use std::task::{Context, Poll}; -use std::time::Duration; use bollard::container::LogOutput; use bollard::errors::Error as BollardError; -use futures_util::{Stream, StreamExt}; -use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; -use tokio::task::JoinHandle; -use tokio::time::timeout; +use futures_util::Stream; +use tokio::io::{AsyncRead, AsyncWrite}; use super::ExecRequest; use super::acp_frame::OutboundFrameAssembler; -use super::acp_helpers; use super::acp_policy::MethodDenylist; use super::acp_runtime::{ OutboundPolicyAdapter, SINK_CHANNEL_CAPACITY, WriteCmd, run_container_stdin_sink, @@ -81,6 +77,16 @@ use super::host_io::stdin_forwarding_disabled_for_tests; use super::runtime_helpers::exec_failed; use super::session::CapabilityPolicy; use crate::error::PodbotError; + +#[path = "protocol_output.rs"] +mod output_io; +#[path = "protocol_stdin.rs"] +mod stdin_io; + +use self::output_io::{AdapterOutputIo, run_output_loop_async, run_output_loop_with_adapter}; +use self::stdin_io::{ + forward_host_stdin_to_channel, forward_host_stdin_to_exec_async, settle_stdin_forwarding_task, +}; /// Host-side stdio handles used by the protocol byte proxy. pub(super) struct ProtocolProxyIo { /// Host stdin reader supplied to the forwarding task. @@ -93,14 +99,6 @@ pub(super) struct ProtocolProxyIo { options: ProtocolSessionOptions, } -/// Allow a short grace period for EOF- and flush-driven completion paths to -/// finish before treating stdin forwarding as stalled. `50ms` matches typical -/// local pipe and socket flush timings in this code path without adding a -/// noticeable shutdown delay; if future benchmarks or transport changes show -/// different behaviour, adjust `STDIN_SETTLE_TIMEOUT` and extend the proxy -/// tests that exercise EOF and non-EOF stdin shutdown cases. -const STDIN_SETTLE_TIMEOUT: Duration = Duration::from_millis(50); - /// Maximum bytes buffered between host stdin reads and container input writes. /// This bounds memory consumption per read cycle and provides backpressure by /// limiting how many bytes can be in flight. A 64 KiB buffer aligns with common @@ -339,259 +337,6 @@ where }) } -/// Reads the first newline-delimited ACP frame from `buffered_stdin`, -/// applies capability masking, and forwards the resulting bytes to -/// `sender` as a [`WriteCmd::Forward`]. -/// -/// Returns `Ok(true)` when the caller should continue pumping host -/// stdin (either because the masker produced no bytes to forward, or -/// because the send to the sink succeeded). Returns `Ok(false)` when -/// the sink channel has closed and the caller should return cleanly. -async fn send_masked_initialize_frame( - buffered_stdin: &mut tokio::io::BufReader, - sender: &tokio::sync::mpsc::Sender, -) -> io::Result -where - R: AsyncRead + Unpin, -{ - let bytes = acp_helpers::read_and_mask_initial_acp_frame(buffered_stdin).await?; - if bytes.is_empty() { - return Ok(true); - } - Ok(sender.send(WriteCmd::Forward(bytes)).await.is_ok()) -} - -/// Pumps the remainder of host stdin into the container-stdin sink as -/// a sequence of [`WriteCmd::Forward`] chunks bounded by -/// [`STDIN_BUFFER_CAPACITY`]. Stops on EOF or when the sink channel -/// closes. -async fn pump_raw_frames( - buffered_stdin: &mut tokio::io::BufReader, - sender: &tokio::sync::mpsc::Sender, -) -> io::Result<()> -where - R: AsyncRead + Unpin, -{ - use tokio::io::AsyncReadExt; - - let mut buf = vec![0u8; STDIN_BUFFER_CAPACITY]; - loop { - let bytes_read = buffered_stdin.read(&mut buf).await?; - if bytes_read == 0 { - break; - } - let chunk = buf - .get(..bytes_read) - .map(<[u8]>::to_vec) - .unwrap_or_default(); - if sender.send(WriteCmd::Forward(chunk)).await.is_err() { - break; - } - } - Ok(()) -} -async fn forward_host_stdin_to_channel( - host_stdin: HostStdin, - sender: tokio::sync::mpsc::Sender, - rewrite_acp_initialize: bool, -) -> io::Result<()> -where - HostStdin: AsyncRead + Unpin, -{ - let mut buffered_stdin = tokio::io::BufReader::with_capacity(STDIN_BUFFER_CAPACITY, host_stdin); - - if rewrite_acp_initialize && !send_masked_initialize_frame(&mut buffered_stdin, &sender).await? - { - return Ok(()); - } - - pump_raw_frames(&mut buffered_stdin, &sender).await -} - -struct AdapterOutputIo<'a, HostStdout, HostStderr> { - adapter: &'a mut OutboundPolicyAdapter, - host_stdout: &'a mut HostStdout, - host_stderr: &'a mut HostStderr, -} -async fn run_output_loop_with_adapter( - container_id: &str, - output: &mut Pin> + Send>>, - io: &mut AdapterOutputIo<'_, HostStdout, HostStderr>, -) -> Result<(), PodbotError> -where - HostStdout: AsyncWrite + Unpin, - HostStderr: AsyncWrite + Unpin, -{ - while let Some(chunk_result) = output.next().await { - let chunk = chunk_result - .map_err(|error| exec_failed(container_id, format!("exec stream failed: {error}")))?; - match chunk { - LogOutput::StdOut { message } | LogOutput::Console { message } => { - io.adapter - .handle_chunk(message.as_ref(), io.host_stdout) - .await - .map_err(|error| { - exec_failed( - container_id, - format!("failed writing stdout output: {error}"), - ) - })?; - } - LogOutput::StdErr { message } => { - write_output_chunk(container_id, io.host_stderr, message.as_ref(), "stderr") - .await?; - } - LogOutput::StdIn { .. } => {} - } - } - Ok(()) -} -/// Wait for the stdin forwarding task to complete within a short grace -/// period. -async fn settle_stdin_forwarding_task( - container_id: &str, - mut stdin_task: JoinHandle>, - options: ProtocolSessionOptions, -) -> Result<(), PodbotError> { - let Ok(join_result) = timeout(STDIN_SETTLE_TIMEOUT, &mut stdin_task).await else { - // The container output path has already completed, so stdin can be - // cancelled instead of waiting indefinitely on a live host reader. A - // timeout here still indicates that stdin forwarding did not complete - // cleanly before shutdown, so protocol mode must surface that failure - // instead of reporting success with potentially truncated input. - abort_stdin_forwarding_task(stdin_task); - if options.disable_stdin_forwarding || stdin_forwarding_disabled_for_tests() { - return Ok(()); - } - return Err(exec_failed( - container_id, - "stdin forwarding did not complete before protocol session shutdown", - )); - }; - - classify_stdin_forwarding_task_result(container_id, join_result) -} - -/// Map a join result from the stdin forwarding task to a `PodbotError`. -fn classify_stdin_forwarding_task_result( - container_id: &str, - join_result: Result, tokio::task::JoinError>, -) -> Result<(), PodbotError> { - match join_result { - Ok(Ok(())) => Ok(()), - Ok(Err(error)) => Err(exec_failed( - container_id, - format!("failed forwarding stdin to exec input: {error}"), - )), - Err(error) if error.is_cancelled() => Ok(()), - Err(error) => Err(exec_failed( - container_id, - format!("stdin forwarding task failed: {error}"), - )), - } -} - -/// Abort and drop the stdin forwarding task without awaiting it. -fn abort_stdin_forwarding_task(stdin_task: JoinHandle>) { - if !stdin_task.is_finished() { - stdin_task.abort(); - // Avoid awaiting the aborted task here because host stdin may be - // blocked in a non-cancellable read. Dropping the handle mirrors the - // attached-session shutdown path and keeps teardown bounded. - drop(stdin_task); - } -} - -/// Copy host stdin to the container exec input, optionally rewriting the -/// first ACP `initialize` frame before the raw copy begins. -async fn forward_host_stdin_to_exec_async( - host_stdin: HostStdin, - mut input: Pin>, - rewrite_acp_initialize: bool, -) -> io::Result<()> -where - HostStdin: AsyncRead + Unpin, -{ - let mut buffered_stdin = tokio::io::BufReader::with_capacity(STDIN_BUFFER_CAPACITY, host_stdin); - - if rewrite_acp_initialize { - acp_helpers::forward_initial_acp_frame_async(&mut buffered_stdin, &mut input).await?; - } - tokio::io::copy(&mut buffered_stdin, &mut input).await?; - - input.flush().await?; - input.shutdown().await -} - #[cfg(test)] #[path = "protocol_acp_tests.rs"] mod acp_tests; - -/// Drain the container output stream, routing each chunk to host stdout or -/// stderr. -async fn run_output_loop_async( - container_id: &str, - output: &mut Pin> + Send>>, - host_stdout: &mut HostStdout, - host_stderr: &mut HostStderr, -) -> Result<(), PodbotError> -where - HostStdout: AsyncWrite + Unpin, - HostStderr: AsyncWrite + Unpin, -{ - while let Some(chunk_result) = output.next().await { - let chunk = chunk_result - .map_err(|error| exec_failed(container_id, format!("exec stream failed: {error}")))?; - handle_log_output_chunk(container_id, chunk, host_stdout, host_stderr).await?; - } - - Ok(()) -} - -/// Route a single container log-output chunk to the appropriate host stream. -async fn handle_log_output_chunk( - container_id: &str, - chunk: LogOutput, - host_stdout: &mut HostStdout, - host_stderr: &mut HostStderr, -) -> Result<(), PodbotError> -where - HostStdout: AsyncWrite + Unpin, - HostStderr: AsyncWrite + Unpin, -{ - match chunk { - LogOutput::StdOut { message } | LogOutput::Console { message } => { - write_output_chunk(container_id, host_stdout, message.as_ref(), "stdout").await - } - LogOutput::StdErr { message } => { - write_output_chunk(container_id, host_stderr, message.as_ref(), "stderr").await - } - LogOutput::StdIn { .. } => Ok(()), - } -} - -/// Write and flush a byte slice to `writer`, mapping I/O failures to -/// `PodbotError`. -async fn write_output_chunk( - container_id: &str, - writer: &mut Writer, - bytes: &[u8], - stream_name: &str, -) -> Result<(), PodbotError> -where - Writer: AsyncWrite + Unpin, -{ - writer.write_all(bytes).await.map_err(|error| { - exec_failed( - container_id, - format!("failed writing {stream_name} output: {error}"), - ) - })?; - writer.flush().await.map_err(|error| { - exec_failed( - container_id, - format!("failed flushing {stream_name} output: {error}"), - ) - })?; - Ok(()) -} diff --git a/src/engine/connection/exec/protocol_acp_bdd_tests.rs b/src/engine/connection/exec/protocol_acp_bdd_tests.rs index deafefbf..4114a526 100644 --- a/src/engine/connection/exec/protocol_acp_bdd_tests.rs +++ b/src/engine/connection/exec/protocol_acp_bdd_tests.rs @@ -26,10 +26,14 @@ fn acp_masking_state() -> AcpMaskingState { #[given( "ACP stdin contains an initialize request with blocked capabilities and a follow-up request" )] -fn acp_stdin_contains_blocked_initialize_and_follow_up(acp_masking_state: &AcpMaskingState) { - let (host_stdin_bytes, expected) = masked_initialize_with_follow_up(); +fn acp_stdin_contains_blocked_initialize_and_follow_up( + acp_masking_state: &AcpMaskingState, +) -> StepResult<()> { + let (host_stdin_bytes, expected) = masked_initialize_with_follow_up() + .map_err(|err| format!("initialize frames should serialize: {err}"))?; acp_masking_state.host_stdin.set(host_stdin_bytes); acp_masking_state.expected_forwarded.set(expected); + Ok(()) } #[given("ACP stdin contains malformed initialize bytes")] @@ -40,10 +44,12 @@ fn acp_stdin_contains_malformed_initialize(acp_masking_state: &AcpMaskingState) } #[given("ACP stdin contains initialize without blocked capabilities")] -fn acp_stdin_contains_safe_initialize(acp_masking_state: &AcpMaskingState) { - let initialize = initialize_without_blocked_capabilities(); +fn acp_stdin_contains_safe_initialize(acp_masking_state: &AcpMaskingState) -> StepResult<()> { + let initialize = initialize_without_blocked_capabilities() + .map_err(|err| format!("initialize frame should serialize: {err}"))?; acp_masking_state.host_stdin.set(initialize.clone()); acp_masking_state.expected_forwarded.set(initialize); + Ok(()) } #[when("ACP stdin forwarding runs")] @@ -52,7 +58,8 @@ fn acp_stdin_forwarding_runs(acp_masking_state: &AcpMaskingState) -> StepResult< .host_stdin .get() .ok_or_else(|| String::from("host stdin should be configured"))?; - let (forwarded, _) = run_forwarding(&host_stdin_bytes); + let (forwarded, _) = run_forwarding(&host_stdin_bytes) + .map_err(|err| format!("stdin forwarding failed: {err}"))?; acp_masking_state.actual_forwarded.set(forwarded); acp_masking_state.succeeded.set(true); Ok(()) diff --git a/src/engine/connection/exec/protocol_acp_forwarding_tests.rs b/src/engine/connection/exec/protocol_acp_forwarding_tests.rs index 4c184939..e8bf20e8 100644 --- a/src/engine/connection/exec/protocol_acp_forwarding_tests.rs +++ b/src/engine/connection/exec/protocol_acp_forwarding_tests.rs @@ -4,9 +4,10 @@ use super::*; #[test] fn forwarding_leaves_initialize_unchanged_when_acp_rewrite_is_disabled() { - let host_stdin_bytes = initialize_frame("\n"); + let host_stdin_bytes = initialize_frame("\n").expect("initialize frame should serialize"); - let (forwarded, shutdown_called) = run_forwarding_with_rewrite(&host_stdin_bytes, false); + let (forwarded, shutdown_called) = run_forwarding_with_rewrite(&host_stdin_bytes, false) + .expect("stdin forwarding should succeed"); assert_eq!( forwarded, host_stdin_bytes, @@ -17,11 +18,12 @@ fn forwarding_leaves_initialize_unchanged_when_acp_rewrite_is_disabled() { #[test] fn forwarding_masks_initialize_and_preserves_trailing_bytes() { - let mut host_stdin_bytes = initialize_frame("\n"); - let trailing = initialize_frame("\n"); + let mut host_stdin_bytes = initialize_frame("\n").expect("initialize frame should serialize"); + let trailing = initialize_frame("\n").expect("initialize frame should serialize"); host_stdin_bytes.extend_from_slice(&trailing); - let (forwarded, shutdown_called) = run_forwarding(&host_stdin_bytes); + let (forwarded, shutdown_called) = + run_forwarding(&host_stdin_bytes).expect("stdin forwarding should succeed"); let newline_index = forwarded .iter() .position(|byte| *byte == b'\n') @@ -32,9 +34,10 @@ fn forwarding_masks_initialize_and_preserves_trailing_bytes() { let trailing_forwarded = forwarded .get(newline_index + 1..) .expect("trailing bytes should remain addressable"); - let payload = parse_frame_payload(initialize_frame); + let payload = parse_frame_payload(initialize_frame).expect("frame should contain JSON payload"); - assert_masked_client_capabilities(&payload); + check_masked_client_capabilities(&payload) + .expect("blocked client capabilities should be masked"); assert_eq!( trailing_forwarded, trailing.as_slice(), diff --git a/src/engine/connection/exec/protocol_acp_masking_tests.rs b/src/engine/connection/exec/protocol_acp_masking_tests.rs new file mode 100644 index 00000000..09e88f56 --- /dev/null +++ b/src/engine/connection/exec/protocol_acp_masking_tests.rs @@ -0,0 +1,142 @@ +//! Unit tests for `mask_acp_initialize_frame`, covering blocked-capability +//! removal, preservation of unrelated capabilities, and pass-through of +//! non-initialize or malformed frames. + +use rstest::rstest; + +use super::{ + check_masked_client_capabilities, client_capabilities, initialize_frame, + initialize_frame_with_capabilities, initialize_with_only_blocked_capabilities, + malformed_initialize_bytes, mask_acp_initialize_frame, params, parse_frame_payload, + session_new_bytes, split_frame_line_ending, +}; + +#[rstest] +#[case("\n")] +#[case("\r\n")] +fn mask_acp_initialize_frame_removes_blocked_capabilities(#[case] line_ending: &str) { + let frame = initialize_frame(line_ending).expect("initialize frame should serialize"); + let masked = mask_acp_initialize_frame(&frame); + let payload = parse_frame_payload(&masked).expect("frame should contain JSON payload"); + + assert_eq!( + split_frame_line_ending(&masked).1, + line_ending.as_bytes(), + "line ending should be preserved" + ); + check_masked_client_capabilities(&payload) + .expect("blocked client capabilities should be masked"); +} + +#[test] +fn mask_acp_initialize_frame_removes_empty_client_capabilities() { + let frame = + initialize_with_only_blocked_capabilities("\n").expect("initialize frame should serialize"); + let masked = mask_acp_initialize_frame(&frame); + let payload = parse_frame_payload(&masked).expect("frame should contain JSON payload"); + let masked_params = params(&payload).expect("initialize params should remain present"); + + assert!( + !masked_params.contains_key("clientCapabilities"), + "clientCapabilities should be removed when all entries are masked" + ); + assert_eq!( + masked_params.get("protocolVersion"), + Some(&serde_json::json!(1)), + "protocolVersion should remain unchanged" + ); + assert_eq!( + masked_params.get("clientInfo"), + Some(&serde_json::json!({ + "name": "podbot-tests", + "version": "1.0.0" + })), + "clientInfo should remain unchanged" + ); +} + +#[rstest] +#[case( + serde_json::json!({ + "fs": { "readTextFile": true }, + "auth": { "token": true } + }), + &["fs"], + &["auth"] +)] +#[case( + serde_json::json!({ + "terminal": true, + "logging": { "level": "info" } + }), + &["terminal"], + &["logging"] +)] +#[case( + serde_json::json!({ + "fs": { "readTextFile": true }, + "terminal": true, + "auth": { "token": true }, + "logging": { "level": "debug" } + }), + &["fs", "terminal"], + &["auth", "logging"] +)] +fn mask_acp_initialize_frame_preserves_unrelated_capabilities( + #[case] capabilities: serde_json::Value, + #[case] removed_capabilities: &[&str], + #[case] preserved_capabilities: &[&str], +) { + let frame = initialize_frame_with_capabilities(&capabilities, "\n") + .expect("initialize frame should serialize"); + let masked = mask_acp_initialize_frame(&frame); + let result = parse_frame_payload(&masked).expect("frame should contain JSON payload"); + let caps = client_capabilities(&result).expect("clientCapabilities should remain"); + + for capability in removed_capabilities { + assert!( + !caps.contains_key(*capability), + "{capability} should be removed" + ); + } + for capability in preserved_capabilities { + assert!( + caps.contains_key(*capability), + "{capability} should be preserved" + ); + } +} + +#[test] +fn mask_acp_initialize_frame_passes_through_frame_without_line_ending() { + let frame = + initialize_with_only_blocked_capabilities("").expect("initialize frame should serialize"); + let masked = mask_acp_initialize_frame(&frame); + + assert!( + !masked.ends_with(b"\n"), + "masked frame should not gain a trailing newline" + ); + let result: serde_json::Value = + serde_json::from_slice(&masked).expect("result should be valid JSON"); + let masked_params = params(&result).expect("params should remain"); + assert!( + masked_params.get("clientCapabilities").is_none(), + "capabilities should still be masked even without a line ending" + ); +} + +#[test] +fn mask_acp_initialize_frame_leaves_non_initialize_messages_unchanged() { + let mut frame = session_new_bytes(); + frame.push(b'\n'); + + assert_eq!(mask_acp_initialize_frame(&frame), frame); +} + +#[test] +fn mask_acp_initialize_frame_leaves_malformed_input_unchanged() { + let frame = malformed_initialize_bytes(); + + assert_eq!(mask_acp_initialize_frame(&frame), frame); +} diff --git a/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs b/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs index adf1a143..6ef99780 100644 --- a/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs +++ b/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs @@ -2,7 +2,7 @@ use std::io; use std::pin::Pin; -use std::sync::{Arc, Mutex}; +use std::sync::{Arc, Mutex, PoisonError}; use std::task::{Context, Poll}; use bollard::container::LogOutput; @@ -21,7 +21,10 @@ struct RecordingOutput { impl RecordingOutput { fn snapshot(&self) -> Vec { - self.bytes.lock().expect("recording mutex").clone() + self.bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() } } @@ -33,7 +36,7 @@ impl AsyncWrite for RecordingOutput { ) -> Poll> { self.bytes .lock() - .expect("recording mutex") + .unwrap_or_else(PoisonError::into_inner) .extend_from_slice(buf); Poll::Ready(Ok(buf.len())) } @@ -55,55 +58,57 @@ fn protocol_request() -> Result { ) } -fn blocked_request_frame(id: i64) -> Vec { +fn blocked_request_frame(id: i64) -> Result, serde_json::Error> { let mut bytes = serde_json::to_vec(&serde_json::json!({ "jsonrpc": "2.0", "id": id, "method": "terminal/create", "params": {}, - })) - .expect("blocked request should serialize"); + }))?; bytes.push(b'\n'); - bytes + Ok(bytes) } -fn drive_policy_session(policy: CapabilityPolicy) -> (Vec, Vec, Vec) { - let runtime = tokio::runtime::Runtime::new().expect("runtime should build"); - let initialize = super::initialize_frame("\n"); - let host_stdin = runtime - .block_on(super::build_host_stdin(&initialize)) - .expect("host stdin should build"); +fn drive_policy_session(policy: CapabilityPolicy) -> io::Result<(Vec, Vec, Vec)> { + let runtime = tokio::runtime::Runtime::new()?; + let initialize = super::initialize_frame("\n").map_err(io::Error::other)?; + let host_stdin = runtime.block_on(super::build_host_stdin(&initialize))?; let host_stdout = RecordingOutput::default(); let host_stdout_handle = host_stdout.clone(); let host_stderr = RecordingOutput::default(); let host_stderr_handle = host_stderr.clone(); let container_input = super::RecordingInputWriter::new(); let container_stdin = container_input.bytes.clone(); + let blocked = blocked_request_frame(7).map_err(io::Error::other)?; let output = stream::iter([Ok(LogOutput::StdOut { - message: blocked_request_frame(7).into(), + message: blocked.into(), })]); + let request = protocol_request().map_err(io::Error::other)?; let stdio = ProtocolProxyIo::new(host_stdin, host_stdout, host_stderr) .with_options(ProtocolSessionOptions::new().with_capability_policy(policy)); runtime .block_on(run_protocol_session_with_io_async( - &protocol_request().expect("protocol request should build"), + &request, Box::pin(output), Box::pin(container_input), stdio, )) - .expect("policy session should complete"); - - ( - container_stdin.lock().expect("stdin mutex").clone(), + .map_err(io::Error::other)?; + + let recorded_container_stdin = container_stdin + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + Ok(( + recorded_container_stdin, host_stdout_handle.snapshot(), host_stderr_handle.snapshot(), - ) + )) } -fn parse_json_line(bytes: &[u8]) -> serde_json::Value { +fn parse_json_line(bytes: &[u8]) -> Result { serde_json::from_slice(bytes.strip_suffix(b"\n").unwrap_or(bytes)) - .expect("line should contain JSON") } fn split_lines(bytes: &[u8]) -> Vec<&[u8]> { @@ -117,8 +122,12 @@ fn line_matches(bytes: &[u8], expected: &[u8]) -> bool { bytes == expected } +/// Pure query: returns `true` when `bytes` parses as the synthesized denial +/// response for the blocked request with id 7 and the given `method`. fn synthesized_response_for_method(bytes: &[u8], method: &str) -> bool { - let response = parse_json_line(bytes); + let Ok(response) = parse_json_line(bytes) else { + return false; + }; response.get("id") == Some(&serde_json::json!(7)) && response .get("error") @@ -130,9 +139,10 @@ fn synthesized_response_for_method(bytes: &[u8], method: &str) -> bool { #[test] fn mask_and_deny_masks_initialize_and_synthesizes_blocked_response() { let (container_stdin, host_stdout, _host_stderr) = - drive_policy_session(CapabilityPolicy::MaskAndDeny); + drive_policy_session(CapabilityPolicy::MaskAndDeny).expect("policy session should run"); let lines = split_lines(&container_stdin); - let expected_initialize = super::mask_acp_initialize_frame(&super::initialize_frame("\n")); + let initialize = super::initialize_frame("\n").expect("initialize frame should serialize"); + let expected_initialize = super::mask_acp_initialize_frame(&initialize); assert!( host_stdout.is_empty(), @@ -155,12 +165,12 @@ fn mask_and_deny_masks_initialize_and_synthesizes_blocked_response() { #[test] fn disabled_policy_preserves_plain_streaming_proxy_behaviour() { let (container_stdin, host_stdout, _host_stderr) = - drive_policy_session(CapabilityPolicy::Disabled); - let blocked = blocked_request_frame(7); + drive_policy_session(CapabilityPolicy::Disabled).expect("policy session should run"); + let initialize = super::initialize_frame("\n").expect("initialize frame should serialize"); + let blocked = blocked_request_frame(7).expect("blocked request should serialize"); assert_eq!( - container_stdin, - super::initialize_frame("\n"), + container_stdin, initialize, "Disabled should forward host stdin without ACP masking", ); assert_eq!( @@ -172,12 +182,13 @@ fn disabled_policy_preserves_plain_streaming_proxy_behaviour() { #[test] fn mask_only_masks_initialize_without_runtime_denylist_enforcement() { let (container_stdin, host_stdout, _host_stderr) = - drive_policy_session(CapabilityPolicy::MaskOnly); - let blocked = blocked_request_frame(7); + drive_policy_session(CapabilityPolicy::MaskOnly).expect("policy session should run"); + let initialize = super::initialize_frame("\n").expect("initialize frame should serialize"); + let blocked = blocked_request_frame(7).expect("blocked request should serialize"); assert_eq!( container_stdin, - super::mask_acp_initialize_frame(&super::initialize_frame("\n")), + super::mask_acp_initialize_frame(&initialize), "MaskOnly should apply the initialize masking path", ); assert_eq!( diff --git a/src/engine/connection/exec/protocol_acp_routing_tests.rs b/src/engine/connection/exec/protocol_acp_routing_tests.rs new file mode 100644 index 00000000..15b4b67e --- /dev/null +++ b/src/engine/connection/exec/protocol_acp_routing_tests.rs @@ -0,0 +1,127 @@ +//! Verifies that `CapabilityPolicy` selects raw forwarding or runtime +//! enforcement for outbound ACP frames. + +use std::io; +use std::sync::PoisonError; + +use bollard::container::LogOutput; +use futures_util::stream; +use rstest::rstest; + +use super::{RecordingInputWriter, build_host_stdin}; +use crate::engine::connection::exec::protocol::{ + ProtocolProxyIo, ProtocolSessionOptions, run_protocol_session_with_io_async, +}; +use crate::engine::connection::exec::session::CapabilityPolicy; +use crate::engine::connection::exec::{ExecMode, ExecRequest}; +use crate::error::PodbotError; + +fn protocol_request() -> Result { + ExecRequest::new( + "capability-policy-routing", + vec![String::from("codex"), String::from("app-server")], + ExecMode::Protocol, + ) +} + +fn blocked_terminal_create_frame() -> Result, serde_json::Error> { + let mut bytes = serde_json::to_vec(&serde_json::json!({ + "jsonrpc": "2.0", + "id": 7, + "method": "terminal/create", + "params": {}, + }))?; + bytes.push(b'\n'); + Ok(bytes) +} + +fn run_policy_output_frame( + policy: CapabilityPolicy, + frame: &[u8], +) -> io::Result<(Vec, Vec)> { + let runtime = tokio::runtime::Runtime::new()?; + let request = protocol_request().map_err(io::Error::other)?; + let host_stdin = runtime.block_on(build_host_stdin(&[]))?; + let host_stdout = RecordingInputWriter::new(); + let host_stdout_bytes = host_stdout.bytes.clone(); + let host_stderr = RecordingInputWriter::new(); + let container_input = RecordingInputWriter::new(); + let container_stdin_bytes = container_input.bytes.clone(); + let output = stream::iter([Ok(LogOutput::StdOut { + message: frame.to_vec().into(), + })]); + let stdio = ProtocolProxyIo::new(host_stdin, host_stdout, host_stderr) + .with_options(ProtocolSessionOptions::new().with_capability_policy(policy)); + + runtime + .block_on(run_protocol_session_with_io_async( + &request, + Box::pin(output), + Box::pin(container_input), + stdio, + )) + .map_err(io::Error::other)?; + + let recorded_stdout = host_stdout_bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + let recorded_container_stdin = container_stdin_bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + Ok((recorded_stdout, recorded_container_stdin)) +} + +/// Pure query: returns `true` when `bytes` parses as the synthesized denial +/// response for the blocked `terminal/create` request with id 7. +fn synthesized_response_for_terminal_create(bytes: &[u8]) -> bool { + let Ok(response) = + serde_json::from_slice::(bytes.strip_suffix(b"\n").unwrap_or(bytes)) + else { + return false; + }; + response.get("id") == Some(&serde_json::json!(7)) + && response + .get("error") + .and_then(|error| error.get("data")) + .and_then(|data| data.get("method")) + == Some(&serde_json::json!("terminal/create")) +} + +#[rstest] +#[case::mask_and_deny_routes_through_enforcement_path(CapabilityPolicy::MaskAndDeny, false, true)] +#[case::disabled_policy_forwards_all_frames_raw(CapabilityPolicy::Disabled, true, false)] +#[case::mask_only_policy_forwards_blocked_frames_raw(CapabilityPolicy::MaskOnly, true, false)] +fn routes_output_frame_for_capability_policy( + #[case] policy: CapabilityPolicy, + #[case] expect_forward_raw: bool, + #[case] expect_synthesized_response: bool, +) { + let frame = blocked_terminal_create_frame().expect("blocked request should serialize"); + let (host_stdout, container_stdin) = + run_policy_output_frame(policy, &frame).expect("policy session should complete"); + + if expect_forward_raw { + assert_eq!( + host_stdout, frame, + "{policy:?} should preserve the byte-transparent output path", + ); + } else { + assert_ne!( + host_stdout, frame, + "{policy:?} must not forward blocked frames verbatim", + ); + } + if expect_synthesized_response { + assert!( + synthesized_response_for_terminal_create(&container_stdin), + "{policy:?} should write a synthesized denial response to container stdin", + ); + } else { + assert!( + container_stdin.is_empty(), + "{policy:?} should not write a synthesized denial response", + ); + } +} diff --git a/src/engine/connection/exec/protocol_acp_tests.rs b/src/engine/connection/exec/protocol_acp_tests.rs index 313a93ed..3989fe84 100644 --- a/src/engine/connection/exec/protocol_acp_tests.rs +++ b/src/engine/connection/exec/protocol_acp_tests.rs @@ -1,11 +1,14 @@ //! ACP capability masking tests for the protocol stdin proxy. +//! +//! This module hosts the shared test harness (recording writers, frame +//! builders, and the synchronous forwarding runner) and delegates the +//! individual test groups to sibling submodules. use std::io; use std::pin::Pin; -use std::sync::{Arc, Mutex}; +use std::sync::{Arc, Mutex, PoisonError}; use std::task::{Context, Poll}; -use rstest::rstest; use tokio::io::{AsyncWrite, AsyncWriteExt, DuplexStream}; use super::*; @@ -36,7 +39,7 @@ impl AsyncWrite for RecordingInputWriter { ) -> Poll> { self.bytes .lock() - .expect("writer mutex should not poison") + .unwrap_or_else(PoisonError::into_inner) .extend_from_slice(buf); Poll::Ready(Ok(buf.len())) } @@ -49,41 +52,36 @@ impl AsyncWrite for RecordingInputWriter { *self .shutdown_called .lock() - .expect("shutdown mutex should not poison") = true; + .unwrap_or_else(PoisonError::into_inner) = true; Poll::Ready(Ok(())) } } fn initialize_frame_with_capabilities( - capabilities: serde_json::Value, + capabilities: &serde_json::Value, line_ending: &str, -) -> Vec { - let mut payload = serde_json::json!({ +) -> Result, serde_json::Error> { + let payload = serde_json::json!({ "jsonrpc": "2.0", "id": 0, "method": "initialize", "params": { "protocolVersion": 1, - "clientCapabilities": null, + "clientCapabilities": capabilities, "clientInfo": { "name": "podbot-tests", "version": "1.0.0" } } }); - payload - .get_mut("params") - .and_then(serde_json::Value::as_object_mut) - .expect("initialize params should be present") - .insert("clientCapabilities".to_owned(), capabilities); - let mut frame = serde_json::to_vec(&payload).expect("initialize payload should serialise"); + let mut frame = serde_json::to_vec(&payload)?; frame.extend_from_slice(line_ending.as_bytes()); - frame + Ok(frame) } -fn initialize_frame(line_ending: &str) -> Vec { +fn initialize_frame(line_ending: &str) -> Result, serde_json::Error> { initialize_frame_with_capabilities( - serde_json::json!({ + &serde_json::json!({ "fs": { "readTextFile": true, "writeTextFile": true }, "terminal": true, "_meta": { "custom": true } @@ -94,7 +92,7 @@ fn initialize_frame(line_ending: &str) -> Vec { /// Builds a serialised ACP `initialize` frame whose `clientCapabilities` /// contains only `_meta` (no blocked entries), terminated with `\n`. -pub(super) fn initialize_without_blocked_capabilities() -> Vec { +pub(super) fn initialize_without_blocked_capabilities() -> Result, serde_json::Error> { let payload = serde_json::json!({ "jsonrpc": "2.0", "id": 0, @@ -109,14 +107,16 @@ pub(super) fn initialize_without_blocked_capabilities() -> Vec { } }); - let mut frame = serde_json::to_vec(&payload).expect("initialize payload should serialize"); + let mut frame = serde_json::to_vec(&payload)?; frame.push(b'\n'); - frame + Ok(frame) } -fn initialize_with_only_blocked_capabilities(line_ending: &str) -> Vec { +fn initialize_with_only_blocked_capabilities( + line_ending: &str, +) -> Result, serde_json::Error> { initialize_frame_with_capabilities( - serde_json::json!({ + &serde_json::json!({ "fs": { "readTextFile": true, "writeTextFile": true }, "terminal": true }), @@ -135,46 +135,41 @@ pub(super) fn malformed_initialize_bytes() -> Vec { .to_vec() } -fn parse_frame_payload(frame: &[u8]) -> serde_json::Value { +fn parse_frame_payload(frame: &[u8]) -> Result { let (payload, _) = split_frame_line_ending(frame); - serde_json::from_slice(payload).expect("frame should contain JSON payload") + serde_json::from_slice(payload) } -fn assert_masked_client_capabilities(message: &serde_json::Value) { - let client_capabilities = message - .get("params") - .and_then(serde_json::Value::as_object) - .and_then(|params| params.get("clientCapabilities")) - .and_then(serde_json::Value::as_object) - .expect("clientCapabilities should remain present"); - assert!( - !client_capabilities.contains_key(ACP_FILE_SYSTEM_CAPABILITY), - "fs capability should be removed" - ); - assert!( - !client_capabilities.contains_key(ACP_TERMINAL_CAPABILITY), - "terminal capability should be removed" - ); - assert!( - client_capabilities.contains_key("_meta"), - "unrelated capabilities should remain" - ); +/// Verifies that the blocked `fs` and `terminal` capabilities have been +/// removed while unrelated entries survive, returning a descriptive error on +/// the first violated expectation. +fn check_masked_client_capabilities(message: &serde_json::Value) -> Result<(), String> { + let caps = client_capabilities(message) + .ok_or_else(|| String::from("clientCapabilities should remain present"))?; + if caps.contains_key(ACP_FILE_SYSTEM_CAPABILITY) { + return Err(String::from("fs capability should be removed")); + } + if caps.contains_key(ACP_TERMINAL_CAPABILITY) { + return Err(String::from("terminal capability should be removed")); + } + if !caps.contains_key("_meta") { + return Err(String::from("unrelated capabilities should remain")); + } + Ok(()) } -fn client_capabilities(message: &serde_json::Value) -> &serde_json::Map { +fn client_capabilities( + message: &serde_json::Value, +) -> Option<&serde_json::Map> { message .get("params") .and_then(serde_json::Value::as_object) .and_then(|params| params.get("clientCapabilities")) .and_then(serde_json::Value::as_object) - .expect("clientCapabilities should remain") } -fn params(message: &serde_json::Value) -> &serde_json::Map { - message - .get("params") - .and_then(serde_json::Value::as_object) - .expect("params should remain") +fn params(message: &serde_json::Value) -> Option<&serde_json::Map> { + message.get("params").and_then(serde_json::Value::as_object) } async fn build_host_stdin(bytes: &[u8]) -> io::Result { @@ -188,175 +183,42 @@ async fn build_host_stdin(bytes: &[u8]) -> io::Result { /// Runs ACP stdin forwarding synchronously with `rewrite_acp_initialize = /// true`, returning the bytes written to the container input and whether /// `poll_shutdown` was called. -pub(super) fn run_forwarding(host_stdin_bytes: &[u8]) -> (Vec, bool) { +pub(super) fn run_forwarding(host_stdin_bytes: &[u8]) -> io::Result<(Vec, bool)> { run_forwarding_with_rewrite(host_stdin_bytes, true) } fn run_forwarding_with_rewrite( host_stdin_bytes: &[u8], rewrite_acp_initialize: bool, -) -> (Vec, bool) { - let runtime = tokio::runtime::Runtime::new().expect("runtime should build"); - let host_stdin = runtime - .block_on(build_host_stdin(host_stdin_bytes)) - .expect("host stdin should build"); +) -> io::Result<(Vec, bool)> { + let runtime = tokio::runtime::Runtime::new()?; + let host_stdin = runtime.block_on(build_host_stdin(host_stdin_bytes))?; let container_input = RecordingInputWriter::new(); let forwarded_bytes = container_input.bytes.clone(); let shutdown_called = container_input.shutdown_called.clone(); - runtime - .block_on(forward_host_stdin_to_exec_async( - host_stdin, - Box::pin(container_input), - rewrite_acp_initialize, - )) - .expect("stdin forwarding should succeed"); - - ( - forwarded_bytes - .lock() - .expect("writer mutex should not poison") - .clone(), - *shutdown_called - .lock() - .expect("shutdown mutex should not poison"), - ) -} - -#[rstest] -#[case("\n")] -#[case("\r\n")] -fn mask_acp_initialize_frame_removes_blocked_capabilities(#[case] line_ending: &str) { - let frame = initialize_frame(line_ending); - let masked = mask_acp_initialize_frame(&frame); - let payload = parse_frame_payload(&masked); - - assert_eq!( - split_frame_line_ending(&masked).1, - line_ending.as_bytes(), - "line ending should be preserved" - ); - assert_masked_client_capabilities(&payload); -} - -#[test] -fn mask_acp_initialize_frame_removes_empty_client_capabilities() { - let frame = initialize_with_only_blocked_capabilities("\n"); - let masked = mask_acp_initialize_frame(&frame); - let payload = parse_frame_payload(&masked); - let params = payload - .get("params") - .and_then(serde_json::Value::as_object) - .expect("initialize params should remain present"); - - assert!( - !params.contains_key("clientCapabilities"), - "clientCapabilities should be removed when all entries are masked" - ); - assert_eq!( - params.get("protocolVersion"), - Some(&serde_json::json!(1)), - "protocolVersion should remain unchanged" - ); - assert_eq!( - params.get("clientInfo"), - Some(&serde_json::json!({ - "name": "podbot-tests", - "version": "1.0.0" - })), - "clientInfo should remain unchanged" - ); -} - -#[rstest] -#[case( - serde_json::json!({ - "fs": { "readTextFile": true }, - "auth": { "token": true } - }), - &["fs"], - &["auth"] -)] -#[case( - serde_json::json!({ - "terminal": true, - "logging": { "level": "info" } - }), - &["terminal"], - &["logging"] -)] -#[case( - serde_json::json!({ - "fs": { "readTextFile": true }, - "terminal": true, - "auth": { "token": true }, - "logging": { "level": "debug" } - }), - &["fs", "terminal"], - &["auth", "logging"] -)] -fn mask_acp_initialize_frame_preserves_unrelated_capabilities( - #[case] capabilities: serde_json::Value, - #[case] removed_capabilities: &[&str], - #[case] preserved_capabilities: &[&str], -) { - let frame = initialize_frame_with_capabilities(capabilities, "\n"); - let masked = mask_acp_initialize_frame(&frame); - let result = parse_frame_payload(&masked); - let caps = client_capabilities(&result); - - for capability in removed_capabilities { - assert!( - !caps.contains_key(*capability), - "{capability} should be removed" - ); - } - for capability in preserved_capabilities { - assert!( - caps.contains_key(*capability), - "{capability} should be preserved" - ); - } -} - -#[test] -fn mask_acp_initialize_frame_passes_through_frame_without_line_ending() { - let frame = initialize_with_only_blocked_capabilities(""); - let masked = mask_acp_initialize_frame(&frame); - - assert!( - !masked.ends_with(b"\n"), - "masked frame should not gain a trailing newline" - ); - let result: serde_json::Value = - serde_json::from_slice(&masked).expect("result should be valid JSON"); - assert!( - params(&result).get("clientCapabilities").is_none(), - "capabilities should still be masked even without a line ending" - ); -} - -#[test] -fn mask_acp_initialize_frame_leaves_non_initialize_messages_unchanged() { - let mut frame = session_new_bytes(); - frame.push(b'\n'); - - assert_eq!(mask_acp_initialize_frame(&frame), frame); -} - -#[test] -fn mask_acp_initialize_frame_leaves_malformed_input_unchanged() { - let frame = malformed_initialize_bytes(); - - assert_eq!(mask_acp_initialize_frame(&frame), frame); + runtime.block_on(forward_host_stdin_to_exec_async( + host_stdin, + Box::pin(container_input), + rewrite_acp_initialize, + ))?; + + let forwarded = forwarded_bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone(); + let shutdown = *shutdown_called + .lock() + .unwrap_or_else(PoisonError::into_inner); + Ok((forwarded, shutdown)) } /// Constructs a host-stdin byte sequence containing a masked `initialize` /// frame followed by a follow-up frame, and returns both the raw input bytes /// and the expected post-masking output bytes for BDD assertion. -pub(super) fn masked_initialize_with_follow_up() -> (Vec, Vec) { - let mut host_stdin_bytes = initialize_frame("\n"); - let follow_up = initialize_frame("\n"); +pub(super) fn masked_initialize_with_follow_up() -> Result<(Vec, Vec), serde_json::Error> { + let mut host_stdin_bytes = initialize_frame("\n")?; + let follow_up = initialize_frame("\n")?; host_stdin_bytes.extend_from_slice(&follow_up); let expected_initialize = serde_json::json!({ @@ -377,136 +239,18 @@ pub(super) fn masked_initialize_with_follow_up() -> (Vec, Vec) { } }); - let mut expected = serde_json::to_vec(&expected_initialize) - .expect("expected initialize payload should serialize"); + let mut expected = serde_json::to_vec(&expected_initialize)?; expected.push(b'\n'); expected.extend_from_slice(&follow_up); - (host_stdin_bytes, expected) + Ok((host_stdin_bytes, expected)) } -#[cfg(test)] -mod capability_policy_routing { - //! Verifies that `CapabilityPolicy` selects raw forwarding or runtime - //! enforcement for outbound ACP frames. - - use bollard::container::LogOutput; - use futures_util::stream; - - use super::*; - use crate::engine::connection::exec::session::CapabilityPolicy; - use crate::engine::connection::exec::{ExecMode, ExecRequest}; - - fn protocol_request() -> ExecRequest { - ExecRequest::new( - "capability-policy-routing", - vec![String::from("codex"), String::from("app-server")], - ExecMode::Protocol, - ) - .expect("protocol request should build") - } - - fn blocked_terminal_create_frame() -> Vec { - let mut bytes = serde_json::to_vec(&serde_json::json!({ - "jsonrpc": "2.0", - "id": 7, - "method": "terminal/create", - "params": {}, - })) - .expect("blocked request should serialize"); - bytes.push(b'\n'); - bytes - } +#[path = "protocol_acp_masking_tests.rs"] +mod masking_tests; - fn run_policy_output_frame(policy: CapabilityPolicy, frame: &[u8]) -> (Vec, Vec) { - let runtime = tokio::runtime::Runtime::new().expect("runtime should build"); - let host_stdin = runtime - .block_on(build_host_stdin(&[])) - .expect("host stdin should build"); - let host_stdout = RecordingInputWriter::new(); - let host_stdout_bytes = host_stdout.bytes.clone(); - let host_stderr = RecordingInputWriter::new(); - let container_input = RecordingInputWriter::new(); - let container_stdin_bytes = container_input.bytes.clone(); - let output = stream::iter([Ok(LogOutput::StdOut { - message: frame.to_vec().into(), - })]); - let stdio = ProtocolProxyIo::new(host_stdin, host_stdout, host_stderr) - .with_options(ProtocolSessionOptions::new().with_capability_policy(policy)); - - runtime - .block_on(run_protocol_session_with_io_async( - &protocol_request(), - Box::pin(output), - Box::pin(container_input), - stdio, - )) - .expect("protocol session should complete"); - - ( - host_stdout_bytes - .lock() - .expect("stdout mutex should not poison") - .clone(), - container_stdin_bytes - .lock() - .expect("stdin mutex should not poison") - .clone(), - ) - } - - fn synthesized_response_for_terminal_create(bytes: &[u8]) -> bool { - let response: serde_json::Value = - serde_json::from_slice(bytes.strip_suffix(b"\n").unwrap_or(bytes)) - .expect("synthesized response should contain JSON"); - response.get("id") == Some(&serde_json::json!(7)) - && response - .get("error") - .and_then(|error| error.get("data")) - .and_then(|data| data.get("method")) - == Some(&serde_json::json!("terminal/create")) - } - - #[rstest] - #[case::mask_and_deny_routes_through_enforcement_path( - CapabilityPolicy::MaskAndDeny, - false, - true - )] - #[case::disabled_policy_forwards_all_frames_raw(CapabilityPolicy::Disabled, true, false)] - #[case::mask_only_policy_forwards_blocked_frames_raw(CapabilityPolicy::MaskOnly, true, false)] - fn routes_output_frame_for_capability_policy( - #[case] policy: CapabilityPolicy, - #[case] expect_forward_raw: bool, - #[case] expect_synthesized_response: bool, - ) { - let frame = blocked_terminal_create_frame(); - let (host_stdout, container_stdin) = run_policy_output_frame(policy, &frame); - - if expect_forward_raw { - assert_eq!( - host_stdout, frame, - "{policy:?} should preserve the byte-transparent output path", - ); - } else { - assert_ne!( - host_stdout, frame, - "{policy:?} must not forward blocked frames verbatim", - ); - } - if expect_synthesized_response { - assert!( - synthesized_response_for_terminal_create(&container_stdin), - "{policy:?} should write a synthesized denial response to container stdin", - ); - } else { - assert!( - container_stdin.is_empty(), - "{policy:?} should not write a synthesized denial response", - ); - } - } -} +#[path = "protocol_acp_routing_tests.rs"] +mod capability_policy_routing; #[path = "protocol_acp_forwarding_tests.rs"] mod forwarding_tests; diff --git a/src/engine/connection/exec/protocol_output.rs b/src/engine/connection/exec/protocol_output.rs new file mode 100644 index 00000000..4242cd47 --- /dev/null +++ b/src/engine/connection/exec/protocol_output.rs @@ -0,0 +1,132 @@ +//! Container-output routing loops for protocol exec sessions. +//! +//! These helpers drain the container output stream and route each chunk to +//! host stdout or stderr, either directly (plain byte proxy) or through the +//! outbound ACP policy adapter when runtime enforcement is active. The stdout +//! purity contract documented in the parent module is upheld here: only +//! container stdout and console bytes ever reach host stdout. + +use std::pin::Pin; + +use bollard::container::LogOutput; +use bollard::errors::Error as BollardError; +use futures_util::{Stream, StreamExt}; +use tokio::io::{AsyncWrite, AsyncWriteExt}; + +use super::super::acp_runtime::OutboundPolicyAdapter; +use super::super::runtime_helpers::exec_failed; +use crate::error::PodbotError; + +/// Borrowed IO handles threaded through the adapter-driven output loop. +pub(super) struct AdapterOutputIo<'a, HostStdout, HostStderr> { + /// Outbound ACP policy adapter that inspects container stdout frames. + pub(super) adapter: &'a mut OutboundPolicyAdapter, + /// Host stdout writer used for container stdout and console output. + pub(super) host_stdout: &'a mut HostStdout, + /// Host stderr writer used for container stderr output. + pub(super) host_stderr: &'a mut HostStderr, +} + +/// Drain the container output stream through the outbound policy adapter. +pub(super) async fn run_output_loop_with_adapter( + container_id: &str, + output: &mut Pin> + Send>>, + io: &mut AdapterOutputIo<'_, HostStdout, HostStderr>, +) -> Result<(), PodbotError> +where + HostStdout: AsyncWrite + Unpin, + HostStderr: AsyncWrite + Unpin, +{ + while let Some(chunk_result) = output.next().await { + let chunk = chunk_result + .map_err(|error| exec_failed(container_id, format!("exec stream failed: {error}")))?; + match chunk { + LogOutput::StdOut { message } | LogOutput::Console { message } => { + io.adapter + .handle_chunk(message.as_ref(), io.host_stdout) + .await + .map_err(|error| { + exec_failed( + container_id, + format!("failed writing stdout output: {error}"), + ) + })?; + } + LogOutput::StdErr { message } => { + write_output_chunk(container_id, io.host_stderr, message.as_ref(), "stderr") + .await?; + } + LogOutput::StdIn { .. } => {} + } + } + Ok(()) +} + +/// Drain the container output stream, routing each chunk to host stdout or +/// stderr. +pub(super) async fn run_output_loop_async( + container_id: &str, + output: &mut Pin> + Send>>, + host_stdout: &mut HostStdout, + host_stderr: &mut HostStderr, +) -> Result<(), PodbotError> +where + HostStdout: AsyncWrite + Unpin, + HostStderr: AsyncWrite + Unpin, +{ + while let Some(chunk_result) = output.next().await { + let chunk = chunk_result + .map_err(|error| exec_failed(container_id, format!("exec stream failed: {error}")))?; + handle_log_output_chunk(container_id, chunk, host_stdout, host_stderr).await?; + } + + Ok(()) +} + +/// Route a single container log-output chunk to the appropriate host stream. +async fn handle_log_output_chunk( + container_id: &str, + chunk: LogOutput, + host_stdout: &mut HostStdout, + host_stderr: &mut HostStderr, +) -> Result<(), PodbotError> +where + HostStdout: AsyncWrite + Unpin, + HostStderr: AsyncWrite + Unpin, +{ + match chunk { + LogOutput::StdOut { message } | LogOutput::Console { message } => { + write_output_chunk(container_id, host_stdout, message.as_ref(), "stdout").await + } + LogOutput::StdErr { message } => { + write_output_chunk(container_id, host_stderr, message.as_ref(), "stderr").await + } + LogOutput::StdIn { .. } => Ok(()), + } +} + +/// Write and flush a byte slice to `writer`, mapping I/O failures to +/// `PodbotError`. +async fn write_output_chunk( + container_id: &str, + writer: &mut Writer, + bytes: &[u8], + stream_name: &str, +) -> Result<(), PodbotError> +where + Writer: AsyncWrite + Unpin, +{ + writer.write_all(bytes).await.map_err(|error| { + exec_failed( + container_id, + format!("failed writing {stream_name} output: {error}"), + ) + })?; + writer.flush().await.map_err(|error| { + exec_failed( + container_id, + format!("failed flushing {stream_name} output: {error}"), + ) + })?; + Ok(()) +} diff --git a/src/engine/connection/exec/protocol_stdin.rs b/src/engine/connection/exec/protocol_stdin.rs new file mode 100644 index 00000000..d5d8cc1d --- /dev/null +++ b/src/engine/connection/exec/protocol_stdin.rs @@ -0,0 +1,179 @@ +//! Host-stdin forwarding and shutdown settlement for protocol exec sessions. +//! +//! These helpers own the stdin side of the protocol byte proxy: copying host +//! stdin into the container exec input (optionally rewriting the first ACP +//! `initialize` frame), pumping stdin chunks into the container-stdin sink +//! channel under runtime enforcement, and settling or aborting the forwarding +//! task during session shutdown. + +use std::io; +use std::pin::Pin; +use std::time::Duration; + +use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt}; +use tokio::task::JoinHandle; +use tokio::time::timeout; + +use super::super::acp_helpers; +use super::super::acp_runtime::WriteCmd; +use super::super::host_io::stdin_forwarding_disabled_for_tests; +use super::super::runtime_helpers::exec_failed; +use super::{ProtocolSessionOptions, STDIN_BUFFER_CAPACITY}; +use crate::error::PodbotError; + +/// Allow a short grace period for EOF- and flush-driven completion paths to +/// finish before treating stdin forwarding as stalled. `50ms` matches typical +/// local pipe and socket flush timings in this code path without adding a +/// noticeable shutdown delay; if future benchmarks or transport changes show +/// different behaviour, adjust `STDIN_SETTLE_TIMEOUT` and extend the proxy +/// tests that exercise EOF and non-EOF stdin shutdown cases. +const STDIN_SETTLE_TIMEOUT: Duration = Duration::from_millis(50); + +/// Reads the first newline-delimited ACP frame from `buffered_stdin`, +/// applies capability masking, and forwards the resulting bytes to +/// `sender` as a [`WriteCmd::Forward`]. +/// +/// Returns `Ok(true)` when the caller should continue pumping host +/// stdin (either because the masker produced no bytes to forward, or +/// because the send to the sink succeeded). Returns `Ok(false)` when +/// the sink channel has closed and the caller should return cleanly. +async fn send_masked_initialize_frame( + buffered_stdin: &mut tokio::io::BufReader, + sender: &tokio::sync::mpsc::Sender, +) -> io::Result +where + R: AsyncRead + Unpin, +{ + let bytes = acp_helpers::read_and_mask_initial_acp_frame(buffered_stdin).await?; + if bytes.is_empty() { + return Ok(true); + } + Ok(sender.send(WriteCmd::Forward(bytes)).await.is_ok()) +} + +/// Pumps the remainder of host stdin into the container-stdin sink as +/// a sequence of [`WriteCmd::Forward`] chunks bounded by +/// [`STDIN_BUFFER_CAPACITY`]. Stops on EOF or when the sink channel +/// closes. +async fn pump_raw_frames( + buffered_stdin: &mut tokio::io::BufReader, + sender: &tokio::sync::mpsc::Sender, +) -> io::Result<()> +where + R: AsyncRead + Unpin, +{ + use tokio::io::AsyncReadExt; + + let mut buf = vec![0u8; STDIN_BUFFER_CAPACITY]; + loop { + let bytes_read = buffered_stdin.read(&mut buf).await?; + if bytes_read == 0 { + break; + } + let chunk = buf + .get(..bytes_read) + .map(<[u8]>::to_vec) + .unwrap_or_default(); + if sender.send(WriteCmd::Forward(chunk)).await.is_err() { + break; + } + } + Ok(()) +} + +/// Copy host stdin into the container-stdin sink channel, optionally masking +/// the first ACP `initialize` frame before the raw copy begins. +pub(super) async fn forward_host_stdin_to_channel( + host_stdin: HostStdin, + sender: tokio::sync::mpsc::Sender, + rewrite_acp_initialize: bool, +) -> io::Result<()> +where + HostStdin: AsyncRead + Unpin, +{ + let mut buffered_stdin = tokio::io::BufReader::with_capacity(STDIN_BUFFER_CAPACITY, host_stdin); + + if rewrite_acp_initialize && !send_masked_initialize_frame(&mut buffered_stdin, &sender).await? + { + return Ok(()); + } + + pump_raw_frames(&mut buffered_stdin, &sender).await +} + +/// Wait for the stdin forwarding task to complete within a short grace +/// period. +pub(super) async fn settle_stdin_forwarding_task( + container_id: &str, + mut stdin_task: JoinHandle>, + options: ProtocolSessionOptions, +) -> Result<(), PodbotError> { + let Ok(join_result) = timeout(STDIN_SETTLE_TIMEOUT, &mut stdin_task).await else { + // The container output path has already completed, so stdin can be + // cancelled instead of waiting indefinitely on a live host reader. A + // timeout here still indicates that stdin forwarding did not complete + // cleanly before shutdown, so protocol mode must surface that failure + // instead of reporting success with potentially truncated input. + abort_stdin_forwarding_task(stdin_task); + if options.disable_stdin_forwarding || stdin_forwarding_disabled_for_tests() { + return Ok(()); + } + return Err(exec_failed( + container_id, + "stdin forwarding did not complete before protocol session shutdown", + )); + }; + + classify_stdin_forwarding_task_result(container_id, join_result) +} + +/// Map a join result from the stdin forwarding task to a `PodbotError`. +fn classify_stdin_forwarding_task_result( + container_id: &str, + join_result: Result, tokio::task::JoinError>, +) -> Result<(), PodbotError> { + match join_result { + Ok(Ok(())) => Ok(()), + Ok(Err(error)) => Err(exec_failed( + container_id, + format!("failed forwarding stdin to exec input: {error}"), + )), + Err(error) if error.is_cancelled() => Ok(()), + Err(error) => Err(exec_failed( + container_id, + format!("stdin forwarding task failed: {error}"), + )), + } +} + +/// Abort and drop the stdin forwarding task without awaiting it. +fn abort_stdin_forwarding_task(stdin_task: JoinHandle>) { + if !stdin_task.is_finished() { + stdin_task.abort(); + // Avoid awaiting the aborted task here because host stdin may be + // blocked in a non-cancellable read. Dropping the handle mirrors the + // attached-session shutdown path and keeps teardown bounded. + drop(stdin_task); + } +} + +/// Copy host stdin to the container exec input, optionally rewriting the +/// first ACP `initialize` frame before the raw copy begins. +pub(super) async fn forward_host_stdin_to_exec_async( + host_stdin: HostStdin, + mut input: Pin>, + rewrite_acp_initialize: bool, +) -> io::Result<()> +where + HostStdin: AsyncRead + Unpin, +{ + let mut buffered_stdin = tokio::io::BufReader::with_capacity(STDIN_BUFFER_CAPACITY, host_stdin); + + if rewrite_acp_initialize { + acp_helpers::forward_initial_acp_frame_async(&mut buffered_stdin, &mut input).await?; + } + tokio::io::copy(&mut buffered_stdin, &mut input).await?; + + input.flush().await?; + input.shutdown().await +} diff --git a/src/engine/connection/exec/runtime_helpers.rs b/src/engine/connection/exec/runtime_helpers.rs index 77125bab..39decb20 100644 --- a/src/engine/connection/exec/runtime_helpers.rs +++ b/src/engine/connection/exec/runtime_helpers.rs @@ -53,6 +53,10 @@ pub(super) fn exec_failed(container_id: &str, message: impl Into) -> Pod #[cfg(test)] mod tests { + //! Unit tests for exec runtime helper utilities. + + use std::io; + use super::*; use rstest::{fixture, rstest}; @@ -62,11 +66,10 @@ mod tests { } #[fixture] - fn current_thread_runtime() -> tokio::runtime::Runtime { + fn current_thread_runtime() -> io::Result { tokio::runtime::Builder::new_current_thread() .enable_all() .build() - .expect("runtime should be created") } #[test] @@ -105,10 +108,11 @@ mod tests { #[case::ok(OutsideTokioOutcome::Ok)] #[case::err(OutsideTokioOutcome::Err)] fn block_on_runtime_maps_outcomes_outside_tokio( - current_thread_runtime: tokio::runtime::Runtime, + current_thread_runtime: io::Result, #[case] outcome: OutsideTokioOutcome, ) { - let handle = current_thread_runtime.handle().clone(); + let rt = current_thread_runtime.expect("runtime should be created"); + let handle = rt.handle().clone(); match outcome { OutsideTokioOutcome::Ok => { diff --git a/src/engine/connection/exec/session.rs b/src/engine/connection/exec/session.rs index ddb4650e..f26fb359 100644 --- a/src/engine/connection/exec/session.rs +++ b/src/engine/connection/exec/session.rs @@ -70,6 +70,8 @@ pub(super) const fn protocol_session_options( #[cfg(test)] mod tests { + //! Unit tests for exec session capability policies. + use super::*; #[test] diff --git a/src/engine/connection/exec/terminal.rs b/src/engine/connection/exec/terminal.rs index 2251316b..5d9497a0 100644 --- a/src/engine/connection/exec/terminal.rs +++ b/src/engine/connection/exec/terminal.rs @@ -105,6 +105,8 @@ pub(super) async fn wait_for_sigwinch(signal: &mut Option( #[cfg(test)] mod tests { + //! Unit tests for the container git-identity configurator. + + use std::io; + use super::*; use crate::engine::connection::git_identity::host_reader::HostGitIdentity; use crate::engine::{CreateExecFuture, InspectExecFuture, ResizeExecFuture, StartExecFuture}; @@ -218,10 +222,10 @@ mod tests { } } - fn make_runtime() -> (tokio::runtime::Runtime, tokio::runtime::Handle) { - let rt = tokio::runtime::Runtime::new().expect("test requires a Tokio runtime"); + fn make_runtime() -> io::Result<(tokio::runtime::Runtime, tokio::runtime::Handle)> { + let rt = tokio::runtime::Runtime::new()?; let handle = rt.handle().clone(); - (rt, handle) + Ok((rt, handle)) } fn make_exec_client(exit_code: i64) -> MockExecClient { @@ -250,7 +254,7 @@ mod tests { #[test] fn returns_none_configured_when_both_fields_absent() { - let (_rt, handle) = make_runtime(); + let (_rt, handle) = make_runtime().expect("test requires a Tokio runtime"); // No exec expectations — no container commands should be issued. let client = MockExecClient::new(); let identity = HostGitIdentity { @@ -273,7 +277,7 @@ mod tests { #[test] fn returns_configured_when_both_fields_present() { - let (_rt, handle) = make_runtime(); + let (_rt, handle) = make_runtime().expect("test requires a Tokio runtime"); let client = make_exec_client(0); let identity = HostGitIdentity { name: Some(String::from("Alice")), @@ -301,7 +305,7 @@ mod tests { #[case] email: Option<&str>, #[case] missing_warning: &str, ) { - let (_rt, handle) = make_runtime(); + let (_rt, handle) = make_runtime().expect("test requires a Tokio runtime"); let client = make_exec_client(0); let identity = HostGitIdentity { name: name.map(String::from), @@ -332,7 +336,7 @@ mod tests { #[test] fn propagates_exec_failure_as_error() { - let (_rt, handle) = make_runtime(); + let (_rt, handle) = make_runtime().expect("test requires a Tokio runtime"); let client = make_exec_client(1); let identity = HostGitIdentity { name: Some(String::from("Alice")), diff --git a/src/engine/connection/repository_clone/mod.rs b/src/engine/connection/repository_clone/mod.rs index 41e90760..a4500459 100644 --- a/src/engine/connection/repository_clone/mod.rs +++ b/src/engine/connection/repository_clone/mod.rs @@ -144,6 +144,10 @@ fn github_remote(request: &RepositoryCloneRequest<'_>) -> String { #[cfg(test)] mod tests { + //! Unit tests for container repository-clone command construction. + + use std::io; + use super::*; use crate::engine::{CreateExecFuture, InspectExecFuture, ResizeExecFuture, StartExecFuture}; use bollard::exec::{CreateExecOptions, ResizeExecOptions, StartExecOptions}; @@ -171,18 +175,19 @@ mod tests { } } - fn runtime() -> (tokio::runtime::Runtime, tokio::runtime::Handle) { - let rt = tokio::runtime::Runtime::new().expect("test requires a Tokio runtime"); + fn runtime() -> io::Result<(tokio::runtime::Runtime, tokio::runtime::Handle)> { + let rt = tokio::runtime::Runtime::new()?; let handle = rt.handle().clone(); - (rt, handle) + Ok((rt, handle)) } - fn typed_request_values(branch: &str) -> (RepositoryRef, BranchName, WorkspacePath) { - ( - RepositoryRef::parse("leynos/podbot").expect("test repository should parse"), - BranchName::parse(branch).expect("test branch should parse"), - WorkspacePath::parse("/work").expect("test workspace should parse"), - ) + fn typed_request_values( + branch: &str, + ) -> Result<(RepositoryRef, BranchName, WorkspacePath), PodbotError> { + let repository = RepositoryRef::parse("leynos/podbot")?; + let branch_name = BranchName::parse(branch)?; + let workspace = WorkspacePath::parse("/work")?; + Ok((repository, branch_name, workspace)) } fn request<'a>( @@ -200,8 +205,8 @@ mod tests { } } - fn typed_askpass() -> AskpassPath { - AskpassPath::parse("/usr/local/bin/git-askpass").expect("test askpass should parse") + fn typed_askpass() -> Result { + AskpassPath::parse("/usr/local/bin/git-askpass") } fn expect_exec(client: &mut MockExecClient, command: Vec<&'static str>, exit_code: i64) { @@ -267,10 +272,11 @@ mod tests { #[test] fn clones_repository_and_verifies_branch() { - let (_rt, handle) = runtime(); + let (_rt, handle) = runtime().expect("test requires a Tokio runtime"); let mut client = MockExecClient::new(); - let (repository, branch, workspace) = typed_request_values("main"); - let askpass = typed_askpass(); + let (repository, branch, workspace) = + typed_request_values("main").expect("test request values should parse"); + let askpass = typed_askpass().expect("test askpass should parse"); let clone_request = request(&repository, &branch, &workspace, &askpass); arrange_successful_clone(&mut client); expect_exec( @@ -295,10 +301,11 @@ mod tests { #[test] fn clone_failure_returns_exec_error() { - let (_rt, handle) = runtime(); + let (_rt, handle) = runtime().expect("test requires a Tokio runtime"); let mut client = MockExecClient::new(); - let (repository, branch, workspace) = typed_request_values("main"); - let askpass = typed_askpass(); + let (repository, branch, workspace) = + typed_request_values("main").expect("test request values should parse"); + let askpass = typed_askpass().expect("test askpass should parse"); let clone_request = request(&repository, &branch, &workspace, &askpass); expect_exec( &mut client, @@ -324,10 +331,11 @@ mod tests { #[test] fn branch_verification_failure_returns_exec_error() { - let (_rt, handle) = runtime(); + let (_rt, handle) = runtime().expect("test requires a Tokio runtime"); let mut client = MockExecClient::new(); - let (repository, branch, workspace) = typed_request_values("main"); - let askpass = typed_askpass(); + let (repository, branch, workspace) = + typed_request_values("main").expect("test request values should parse"); + let askpass = typed_askpass().expect("test askpass should parse"); let clone_request = request(&repository, &branch, &workspace, &askpass); arrange_successful_clone(&mut client); // Branch verification fails (exit code 1). diff --git a/src/engine/connection/tests.rs b/src/engine/connection/tests.rs index 4998b83a..164b09da 100644 --- a/src/engine/connection/tests.rs +++ b/src/engine/connection/tests.rs @@ -4,6 +4,8 @@ //! covering environment variable resolution, fallback behaviour, and //! connection establishment for various socket types. +use std::io; + use mockable::MockEnv; use rstest::{fixture, rstest}; @@ -50,8 +52,8 @@ fn all_empty_env_vars() -> MockEnv { /// Fixture providing a tokio runtime for async tests. #[fixture] -fn runtime() -> tokio::runtime::Runtime { - tokio::runtime::Runtime::new().expect("runtime creation should succeed") +fn runtime() -> io::Result { + tokio::runtime::Runtime::new() } /// Helper function to create a `MockEnv` with `DOCKER_HOST` set to the @@ -237,10 +239,11 @@ mod tcp; // ============================================================================= #[rstest] -fn connect_and_verify_propagates_connection_errors(runtime: tokio::runtime::Runtime) { +fn connect_and_verify_propagates_connection_errors(runtime: io::Result) { // Using a non-existent Unix socket to trigger a connection error. // The actual error occurs during the connect phase, not the health check. - let result = runtime.block_on(async { + let rt = runtime.expect("runtime creation should succeed"); + let result = rt.block_on(async { EngineConnector::connect_and_verify_async("unix:///nonexistent/socket.sock").await }); @@ -258,10 +261,13 @@ fn connect_and_verify_propagates_connection_errors(runtime: tokio::runtime::Runt #[rstest] #[cfg(unix)] -fn connect_and_verify_classifies_bare_path_socket_not_found(runtime: tokio::runtime::Runtime) { +fn connect_and_verify_classifies_bare_path_socket_not_found( + runtime: io::Result, +) { // Bare paths are normalized to unix:// URIs before connecting, and // classification should use that normalized URI to extract the path. - let result = runtime.block_on(async { + let rt = runtime.expect("runtime creation should succeed"); + let result = rt.block_on(async { EngineConnector::connect_and_verify_async("/nonexistent/socket.sock").await }); @@ -279,7 +285,7 @@ fn connect_and_verify_classifies_bare_path_socket_not_found(runtime: tokio::runt #[rstest] fn connect_with_fallback_and_verify_uses_resolved_socket( empty_env: MockEnv, - runtime: tokio::runtime::Runtime, + runtime: io::Result, ) { // Verify that connect_with_fallback_and_verify resolves the socket correctly // before attempting connection. We use an explicit socket that will fail @@ -287,7 +293,8 @@ fn connect_with_fallback_and_verify_uses_resolved_socket( // path appears in the error message. let resolver = SocketResolver::new(&empty_env); - let result = runtime.block_on(async { + let rt = runtime.expect("runtime creation should succeed"); + let result = rt.block_on(async { EngineConnector::connect_with_fallback_and_verify_async( Some("unix:///nonexistent/test.sock"), &resolver, @@ -313,12 +320,15 @@ fn connect_with_fallback_and_verify_uses_resolved_socket( } #[rstest] -fn connect_with_fallback_and_verify_falls_back_to_env(runtime: tokio::runtime::Runtime) { +fn connect_with_fallback_and_verify_falls_back_to_env( + runtime: io::Result, +) { // Verify that when config is None, the resolver's environment fallback is used. let env = env_with_docker_host("unix:///env/docker.sock"); let resolver = SocketResolver::new(&env); - let result = runtime.block_on(async { + let rt = runtime.expect("runtime creation should succeed"); + let result = rt.block_on(async { EngineConnector::connect_with_fallback_and_verify_async(None::<&str>, &resolver).await }); diff --git a/src/engine/connection/upload_credentials/archive.rs b/src/engine/connection/upload_credentials/archive.rs index 3d0dbc56..d4bba4ed 100644 --- a/src/engine/connection/upload_credentials/archive.rs +++ b/src/engine/connection/upload_credentials/archive.rs @@ -135,35 +135,49 @@ fn append_non_directory_entry( ) -> io::Result<()> { let path = normalize_archive_path(relative_path); match entry.entry_kind { - EntryKind::File => { - let metadata = parent_dir.metadata(&entry.file_name)?; - let mut file = parent_dir.open(&entry.file_name)?; - let mut header = new_entry_header( - EntryType::Regular, - metadata.len(), - metadata_mode(&metadata, DEFAULT_FILE_MODE), - ); - - builder.append_data(&mut header, path, &mut file) - } - EntryKind::Symlink => { - let metadata = parent_dir.symlink_metadata(&entry.file_name)?; - let target = parent_dir.read_link_contents(&entry.file_name)?; - let mut header = new_entry_header( - EntryType::Symlink, - 0, - metadata_mode(&metadata, DEFAULT_FILE_MODE), - ); - - let normalized_target = normalize_archive_path(target.as_path()); - builder.append_link(&mut header, path, normalized_target) - } + EntryKind::File => append_file_entry(builder, parent_dir, entry, path), + EntryKind::Symlink => append_symlink_entry(builder, parent_dir, entry, path), EntryKind::Directory | EntryKind::Other => Err(io::Error::other( "non-directory entry helper received invalid entry kind", )), } } +fn append_file_entry( + builder: &mut Builder>, + parent_dir: &Dir, + entry: &SortedEntry, + path: String, +) -> io::Result<()> { + let metadata = parent_dir.metadata(&entry.file_name)?; + let mut file = parent_dir.open(&entry.file_name)?; + let mut header = new_entry_header( + EntryType::Regular, + metadata.len(), + metadata_mode(&metadata, DEFAULT_FILE_MODE), + ); + + builder.append_data(&mut header, path, &mut file) +} + +fn append_symlink_entry( + builder: &mut Builder>, + parent_dir: &Dir, + entry: &SortedEntry, + path: String, +) -> io::Result<()> { + let metadata = parent_dir.symlink_metadata(&entry.file_name)?; + let target = parent_dir.read_link_contents(&entry.file_name)?; + let mut header = new_entry_header( + EntryType::Symlink, + 0, + metadata_mode(&metadata, DEFAULT_FILE_MODE), + ); + + let normalized_target = normalize_archive_path(target.as_path()); + builder.append_link(&mut header, path, normalized_target) +} + fn new_entry_header(entry_type: EntryType, size: u64, mode: u32) -> Header { let mut header = Header::new_gnu(); header.set_entry_type(entry_type); diff --git a/src/error.rs b/src/error.rs index 602b0e05..92474c03 100644 --- a/src/error.rs +++ b/src/error.rs @@ -238,249 +238,5 @@ pub enum PodbotError { pub type Result = std::result::Result; #[cfg(test)] -mod tests { - //! Unit tests for error type display formatting and conversion behaviour. - - use super::*; - use eyre::Report; - use rstest::{fixture, rstest}; - - /// Fixture providing a sample configuration file path. - #[fixture] - fn config_path() -> PathBuf { - PathBuf::from("/etc/podbot/config.toml") - } - - /// Fixture providing a sample container socket path. - #[fixture] - fn socket_path() -> PathBuf { - PathBuf::from("/run/podman/podman.sock") - } - - /// Fixture providing a sample container ID. - #[fixture] - fn container_id() -> String { - String::from("abc123") - } - - #[rstest] - fn config_error_file_not_found_displays_correctly(config_path: PathBuf) { - let error = ConfigError::FileNotFound { path: config_path }; - assert_eq!( - error.to_string(), - "configuration file not found: /etc/podbot/config.toml" - ); - } - - #[rstest] - #[case( - "port", - "must be a positive integer", - "invalid configuration value for 'port': must be a positive integer" - )] - #[case( - "image", - "cannot be empty", - "invalid configuration value for 'image': cannot be empty" - )] - fn config_error_invalid_value_displays_correctly( - #[case] field: &str, - #[case] reason: &str, - #[case] expected: &str, - ) { - let error = ConfigError::InvalidValue { - field: String::from(field), - reason: String::from(reason), - }; - assert_eq!(error.to_string(), expected); - } - - #[rstest] - fn config_error_parse_error_displays_message() { - let error = ConfigError::ParseError { - message: String::from("unexpected token"), - }; - assert_eq!( - error.to_string(), - "failed to parse configuration file: unexpected token" - ); - } - - #[rstest] - fn config_error_ortho_config_displays_correctly() { - let ortho_error = ortho_config::OrthoError::Validation { - key: String::from("github.app_id"), - message: String::from("must be a positive integer"), - }; - let error = ConfigError::OrthoConfig(Arc::new(ortho_error)); - assert_eq!( - error.to_string(), - "configuration loading failed: Validation failed for 'github.app_id': must be a positive integer" - ); - } - - #[rstest] - fn container_error_permission_denied_displays_correctly(socket_path: PathBuf) { - let error = ContainerError::PermissionDenied { path: socket_path }; - let msg = error.to_string(); - assert!( - msg.starts_with( - "permission denied accessing container socket: /run/podman/podman.sock" - ), - "error message should start with socket path, got: {msg}" - ); - assert!( - msg.contains("Hint:"), - "error message should contain remediation hint, got: {msg}" - ); - } - - #[rstest] - fn container_error_start_failed_includes_container_id(container_id: String) { - let error = ContainerError::StartFailed { - container_id, - message: String::from("image not found"), - }; - assert_eq!( - error.to_string(), - "failed to start container 'abc123': image not found" - ); - } - - #[rstest] - #[case::health_check_failed( - ContainerError::HealthCheckFailed { message: String::from("ping failed") }, - "container engine health check failed: ping failed" - )] - #[case::health_check_timeout( - ContainerError::HealthCheckTimeout { seconds: 10 }, - "container engine health check timed out after 10 seconds" - )] - #[case::runtime_creation_failed( - ContainerError::RuntimeCreationFailed { message: String::from("cannot create reactor") }, - "failed to create async runtime for health check: cannot create reactor" - )] - fn container_error_health_check_displays_correctly( - #[case] error: ContainerError, - #[case] expected: &str, - ) { - assert_eq!(error.to_string(), expected); - } - - #[rstest] - fn github_error_token_expired_displays_correctly() { - let error = GitHubError::TokenExpired; - assert_eq!(error.to_string(), "installation token expired"); - } - - #[rstest] - fn github_error_auth_failed_displays_message() { - let error = GitHubError::AuthenticationFailed { - message: String::from("invalid signature"), - }; - assert_eq!( - error.to_string(), - "GitHub App authentication failed: invalid signature" - ); - } - - #[rstest] - fn filesystem_error_io_error_displays_message(config_path: PathBuf) { - let error = FilesystemError::IoError { - path: config_path, - message: String::from("disk full"), - }; - assert_eq!( - error.to_string(), - "I/O error at '/etc/podbot/config.toml': disk full" - ); - } - - #[rstest] - fn podbot_error_wraps_config_error() { - let config_error = ConfigError::MissingRequired { - field: String::from("github.app_id"), - }; - let podbot_error: PodbotError = config_error.into(); - assert_eq!( - podbot_error.to_string(), - "missing required configuration: github.app_id" - ); - } - - #[rstest] - fn podbot_error_wraps_container_error(container_id: String) { - let container_error = ContainerError::ExecFailed { - container_id, - message: String::from("command not found"), - }; - let podbot_error: PodbotError = container_error.into(); - assert_eq!( - podbot_error.to_string(), - "failed to execute command in container 'abc123': command not found" - ); - } - - #[rstest] - fn podbot_error_wraps_github_error() { - let github_error = GitHubError::TokenRefreshFailed { - message: String::from("rate limited"), - }; - let podbot_error: PodbotError = github_error.into(); - assert_eq!( - podbot_error.to_string(), - "failed to refresh installation token: rate limited" - ); - } - - #[rstest] - fn podbot_error_wraps_filesystem_error(config_path: PathBuf) { - let fs_error = FilesystemError::NotFound { path: config_path }; - let podbot_error: PodbotError = fs_error.into(); - assert_eq!( - podbot_error.to_string(), - "path not found: /etc/podbot/config.toml" - ); - } - - #[rstest] - #[case( - PodbotError::from(ConfigError::MissingRequired { - field: String::from("github.app_id"), - }), - "missing required configuration: github.app_id" - )] - #[case( - PodbotError::from(ContainerError::StartFailed { - container_id: String::from("abc123"), - message: String::from("image missing"), - }), - "failed to start container 'abc123': image missing" - )] - #[case( - PodbotError::from(GitHubError::TokenExpired), - "installation token expired" - )] - fn eyre_report_preserves_error_messages(#[case] error: PodbotError, #[case] expected: &str) { - let report = Report::from(error); - assert_eq!(report.to_string(), expected); - } - - /// Verify that `PodbotError` satisfies the trait bounds required for use - /// in async contexts and across thread boundaries. - #[rstest] - fn podbot_error_implements_std_error_send_sync() { - fn assert_bounds() {} - assert_bounds::(); - } - - /// Verify that each domain error enum implements `std::error::Error`. - #[rstest] - fn domain_errors_implement_std_error() { - fn assert_error() {} - assert_error::(); - assert_error::(); - assert_error::(); - assert_error::(); - } -} +#[path = "error_tests.rs"] +mod tests; diff --git a/src/error_tests.rs b/src/error_tests.rs new file mode 100644 index 00000000..1e7e9811 --- /dev/null +++ b/src/error_tests.rs @@ -0,0 +1,242 @@ +//! Unit tests for error type display formatting and conversion behaviour. + +use super::*; +use eyre::Report; +use rstest::{fixture, rstest}; + +/// Fixture providing a sample configuration file path. +#[fixture] +fn config_path() -> PathBuf { + PathBuf::from("/etc/podbot/config.toml") +} + +/// Fixture providing a sample container socket path. +#[fixture] +fn socket_path() -> PathBuf { + PathBuf::from("/run/podman/podman.sock") +} + +/// Fixture providing a sample container ID. +#[fixture] +fn container_id() -> String { + String::from("abc123") +} + +#[rstest] +fn config_error_file_not_found_displays_correctly(config_path: PathBuf) { + let error = ConfigError::FileNotFound { path: config_path }; + assert_eq!( + error.to_string(), + "configuration file not found: /etc/podbot/config.toml" + ); +} + +#[rstest] +#[case( + "port", + "must be a positive integer", + "invalid configuration value for 'port': must be a positive integer" +)] +#[case( + "image", + "cannot be empty", + "invalid configuration value for 'image': cannot be empty" +)] +fn config_error_invalid_value_displays_correctly( + #[case] field: &str, + #[case] reason: &str, + #[case] expected: &str, +) { + let error = ConfigError::InvalidValue { + field: String::from(field), + reason: String::from(reason), + }; + assert_eq!(error.to_string(), expected); +} + +#[rstest] +fn config_error_parse_error_displays_message() { + let error = ConfigError::ParseError { + message: String::from("unexpected token"), + }; + assert_eq!( + error.to_string(), + "failed to parse configuration file: unexpected token" + ); +} + +#[rstest] +fn config_error_ortho_config_displays_correctly() { + let ortho_error = ortho_config::OrthoError::Validation { + key: String::from("github.app_id"), + message: String::from("must be a positive integer"), + }; + let error = ConfigError::OrthoConfig(Arc::new(ortho_error)); + assert_eq!( + error.to_string(), + "configuration loading failed: Validation failed for 'github.app_id': must be a positive integer" + ); +} + +#[rstest] +fn container_error_permission_denied_displays_correctly(socket_path: PathBuf) { + let error = ContainerError::PermissionDenied { path: socket_path }; + let msg = error.to_string(); + assert!( + msg.starts_with("permission denied accessing container socket: /run/podman/podman.sock"), + "error message should start with socket path, got: {msg}" + ); + assert!( + msg.contains("Hint:"), + "error message should contain remediation hint, got: {msg}" + ); +} + +#[rstest] +fn container_error_start_failed_includes_container_id(container_id: String) { + let error = ContainerError::StartFailed { + container_id, + message: String::from("image not found"), + }; + assert_eq!( + error.to_string(), + "failed to start container 'abc123': image not found" + ); +} + +#[rstest] +#[case::health_check_failed( + ContainerError::HealthCheckFailed { message: String::from("ping failed") }, + "container engine health check failed: ping failed" +)] +#[case::health_check_timeout( + ContainerError::HealthCheckTimeout { seconds: 10 }, + "container engine health check timed out after 10 seconds" +)] +#[case::runtime_creation_failed( + ContainerError::RuntimeCreationFailed { message: String::from("cannot create reactor") }, + "failed to create async runtime for health check: cannot create reactor" +)] +fn container_error_health_check_displays_correctly( + #[case] error: ContainerError, + #[case] expected: &str, +) { + assert_eq!(error.to_string(), expected); +} + +#[rstest] +fn github_error_token_expired_displays_correctly() { + let error = GitHubError::TokenExpired; + assert_eq!(error.to_string(), "installation token expired"); +} + +#[rstest] +fn github_error_auth_failed_displays_message() { + let error = GitHubError::AuthenticationFailed { + message: String::from("invalid signature"), + }; + assert_eq!( + error.to_string(), + "GitHub App authentication failed: invalid signature" + ); +} + +#[rstest] +fn filesystem_error_io_error_displays_message(config_path: PathBuf) { + let error = FilesystemError::IoError { + path: config_path, + message: String::from("disk full"), + }; + assert_eq!( + error.to_string(), + "I/O error at '/etc/podbot/config.toml': disk full" + ); +} + +#[rstest] +fn podbot_error_wraps_config_error() { + let config_error = ConfigError::MissingRequired { + field: String::from("github.app_id"), + }; + let podbot_error: PodbotError = config_error.into(); + assert_eq!( + podbot_error.to_string(), + "missing required configuration: github.app_id" + ); +} + +#[rstest] +fn podbot_error_wraps_container_error(container_id: String) { + let container_error = ContainerError::ExecFailed { + container_id, + message: String::from("command not found"), + }; + let podbot_error: PodbotError = container_error.into(); + assert_eq!( + podbot_error.to_string(), + "failed to execute command in container 'abc123': command not found" + ); +} + +#[rstest] +fn podbot_error_wraps_github_error() { + let github_error = GitHubError::TokenRefreshFailed { + message: String::from("rate limited"), + }; + let podbot_error: PodbotError = github_error.into(); + assert_eq!( + podbot_error.to_string(), + "failed to refresh installation token: rate limited" + ); +} + +#[rstest] +fn podbot_error_wraps_filesystem_error(config_path: PathBuf) { + let fs_error = FilesystemError::NotFound { path: config_path }; + let podbot_error: PodbotError = fs_error.into(); + assert_eq!( + podbot_error.to_string(), + "path not found: /etc/podbot/config.toml" + ); +} + +#[rstest] +#[case( + PodbotError::from(ConfigError::MissingRequired { + field: String::from("github.app_id"), + }), + "missing required configuration: github.app_id" +)] +#[case( + PodbotError::from(ContainerError::StartFailed { + container_id: String::from("abc123"), + message: String::from("image missing"), + }), + "failed to start container 'abc123': image missing" +)] +#[case( + PodbotError::from(GitHubError::TokenExpired), + "installation token expired" +)] +fn eyre_report_preserves_error_messages(#[case] error: PodbotError, #[case] expected: &str) { + let report = Report::from(error); + assert_eq!(report.to_string(), expected); +} + +/// Verify that `PodbotError` satisfies the trait bounds required for use +/// in async contexts and across thread boundaries. +#[rstest] +fn podbot_error_implements_std_error_send_sync() { + fn assert_bounds() {} + assert_bounds::(); +} + +/// Verify that each domain error enum implements `std::error::Error`. +#[rstest] +fn domain_errors_implement_std_error() { + fn assert_error() {} + assert_error::(); + assert_error::(); + assert_error::(); + assert_error::(); +} diff --git a/src/github/classify.rs b/src/github/classify.rs index 412037e9..f8752a7d 100644 --- a/src/github/classify.rs +++ b/src/github/classify.rs @@ -117,6 +117,8 @@ fn is_rate_limited(message: &str) -> bool { #[cfg(test)] mod tests { + //! Unit tests for GitHub error classification. + use super::*; #[test] diff --git a/src/github/client_tests.rs b/src/github/client_tests.rs new file mode 100644 index 00000000..74d858e1 --- /dev/null +++ b/src/github/client_tests.rs @@ -0,0 +1,187 @@ +//! Unit tests for Octocrab App client construction, credential validation, +//! and retry metric status classification. + +use std::io; + +use cap_std::fs_utf8::Dir as Utf8Dir; +use rstest::rstest; +use tempfile::TempDir; + +use super::super::retry_metrics::github_status_class; +use super::super::*; +use super::{ec_pem, temp_key_dir, valid_rsa_pem}; + +#[rstest] +fn build_app_client_with_valid_key_succeeds( + valid_rsa_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); + dir.write("key.pem", &valid_rsa_pem) + .expect("should write key"); + let path = Utf8Path::new("/display/key.pem"); + let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); + // Octocrab's build() spawns a Tower buffer task requiring a Tokio runtime. + let rt = tokio::runtime::Runtime::new().expect("should create tokio runtime"); + let _guard = rt.enter(); + let result = build_app_client(12345, key); + assert!(result.is_ok(), "expected Ok, got: {result:?}"); +} + +#[rstest] +fn build_app_client_with_zero_app_id_succeeds( + valid_rsa_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); + dir.write("key.pem", &valid_rsa_pem) + .expect("should write key"); + let path = Utf8Path::new("/display/key.pem"); + let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); + // Builder does not validate app_id; GitHub validates at token time. + // Octocrab's build() spawns a Tower buffer task requiring a Tokio runtime. + let rt = tokio::runtime::Runtime::new().expect("should create tokio runtime"); + let _guard = rt.enter(); + let result = build_app_client(0, key); + assert!( + result.is_ok(), + "expected Ok even with zero app_id, got: {result:?}" + ); +} + +#[rstest] +fn build_app_client_without_runtime_returns_error( + valid_rsa_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); + dir.write("key.pem", &valid_rsa_pem) + .expect("should write key"); + let path = Utf8Path::new("/display/key.pem"); + let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); + // Call without entering a Tokio runtime — should return Err, not panic. + let result = build_app_client(42, key); + assert!(result.is_err(), "expected Err without runtime, got Ok"); + let message = result.err().map(|e| e.to_string()).unwrap_or_default(); + assert!( + message.contains("no Tokio runtime context"), + "error should mention missing runtime: {message}" + ); +} + +#[rstest] +#[case::client_error(http::StatusCode::TOO_MANY_REQUESTS, "4xx")] +#[case::server_error(http::StatusCode::INTERNAL_SERVER_ERROR, "5xx")] +#[case::redirect(http::StatusCode::TEMPORARY_REDIRECT, "3xx")] +fn github_status_class_groups_status_codes( + #[case] status_code: http::StatusCode, + #[case] expected_class: &str, +) { + assert_eq!(github_status_class(status_code), expected_class); +} + +#[rstest] +#[case::builder_context( + "failed to build GitHub App client: test error", + "failed to build GitHub App client" +)] +#[case::validation_context( + "failed to validate GitHub App credentials: test error", + "failed to validate GitHub App credentials" +)] +fn authentication_failed_error_includes_context( + #[case] message: &str, + #[case] expected_context: &str, +) { + let error = GitHubError::AuthenticationFailed { + message: String::from(message), + }; + let display = error.to_string(); + assert!( + display.contains(expected_context), + "error should include context: {display}" + ); + assert!( + display.contains("test error"), + "error should include cause: {display}" + ); +} + +#[rstest] +#[tokio::test] +async fn validate_app_credentials_with_missing_key_returns_error( + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (temp_dir, _dir) = temp_key_dir.expect("should create temp key dir"); + let key_path = Utf8Path::from_path(temp_dir.path()) + .expect("temp dir path should be UTF-8") + .join("key.pem"); + let result = validate_app_credentials(12345, &key_path).await; + assert!(result.is_err(), "expected Err for missing key file"); + match result { + Err(GitHubError::PrivateKeyLoadFailed { ref path, .. }) => { + assert!( + path.to_string_lossy().contains("key.pem"), + "error path should reference the missing file" + ); + } + other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), + } +} + +#[rstest] +#[tokio::test] +async fn validate_app_credentials_with_invalid_pem_returns_error( + ec_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (tmp, dir) = temp_key_dir.expect("should create temp key dir"); + dir.write("ec.pem", &ec_pem).expect("should write EC key"); + let full_path = tmp.path().join("ec.pem"); + let utf8_path = Utf8Path::from_path(&full_path).expect("temp path should be UTF-8"); + + let result = validate_app_credentials(12345, utf8_path).await; + assert!(result.is_err(), "expected Err for ECDSA key"); + match result { + Err(GitHubError::PrivateKeyLoadFailed { message, .. }) => { + assert!( + message.contains("ECDSA"), + "error should mention ECDSA: {message}" + ); + } + other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), + } +} + +#[rstest] +#[tokio::test] +async fn validate_with_client_propagates_mock_success() { + let mut mock = MockGitHubAppClient::new(); + mock.expect_validate_credentials() + .times(1) + .returning(|| Box::pin(async { Ok(()) })); + + let result = validate_with_client(&mock).await; + assert!(result.is_ok(), "expected Ok from mock client"); +} + +#[rstest] +#[tokio::test] +async fn validate_with_client_propagates_mock_error() { + let mut mock = MockGitHubAppClient::new(); + mock.expect_validate_credentials().times(1).returning(|| { + Box::pin(async { + Err(GitHubError::AuthenticationFailed { + message: String::from("mock authentication failure"), + }) + }) + }); + + let result = validate_with_client(&mock).await; + assert!(result.is_err(), "expected Err from mock client"); + let message = result.err().map(|e| e.to_string()).unwrap_or_default(); + assert!( + message.contains("mock authentication failure"), + "error should propagate mock message: {message}" + ); +} diff --git a/src/github/property_tests.rs b/src/github/property_tests.rs new file mode 100644 index 00000000..7e55e578 --- /dev/null +++ b/src/github/property_tests.rs @@ -0,0 +1,29 @@ +//! Property-based tests for GitHub retry metric status classification. + +use proptest::prelude::*; + +use super::super::retry_metrics::github_status_class; + +proptest! { + #[test] + fn github_status_class_range_invariant(raw_code in 100_u16..=999_u16) { + let status = http::StatusCode::from_u16(raw_code) + .expect("codes 100–999 are valid HTTP status codes"); + let class = github_status_class(status); + let expected = match raw_code { + 100..=199 => "1xx", + 200..=299 => "2xx", + 300..=399 => "3xx", + 400..=499 => "4xx", + 500..=599 => "5xx", + _ => "other", + }; + prop_assert_eq!( + class, + expected, + "github_status_class({}) should map to {}", + raw_code, + expected, + ); + } +} diff --git a/src/github/tests.rs b/src/github/tests.rs index a86bc83d..8eae75d7 100644 --- a/src/github/tests.rs +++ b/src/github/tests.rs @@ -1,15 +1,23 @@ //! Unit tests for GitHub App private key loading and client construction. //! //! Covers happy paths (valid RSA key), unhappy paths (missing, empty, invalid), -//! key type rejection (ECDSA, Ed25519, public keys, certificates), encrypted -//! key detection, and Octocrab App client building. +//! key type rejection (ECDSA, Ed25519, public keys, certificates), and +//! encrypted key detection. Client construction and credential validation +//! live in [`client_tests`]; property-based coverage lives in +//! [`property_tests`]. + +use std::io; -use super::retry_metrics::github_status_class; use super::*; use cap_std::fs_utf8::Dir as Utf8Dir; use rstest::{fixture, rstest}; use tempfile::TempDir; +#[path = "client_tests.rs"] +mod client_tests; +#[path = "property_tests.rs"] +mod property_tests; + /// Fixture providing valid RSA PEM content (PKCS#1 format). #[fixture] fn valid_rsa_pem() -> String { @@ -30,20 +38,22 @@ fn ed25519_pem() -> String { /// Fixture providing a temporary directory opened as a `Dir` capability. #[fixture] -fn temp_key_dir() -> (TempDir, Utf8Dir) { - let temp_dir = tempfile::tempdir().expect("should create temp dir"); +fn temp_key_dir() -> io::Result<(TempDir, Utf8Dir)> { + let temp_dir = tempfile::tempdir()?; let path_str = temp_dir .path() .to_str() - .expect("temp dir path should be UTF-8"); - let dir = - Utf8Dir::open_ambient_dir(path_str, ambient_authority()).expect("should open temp dir"); - (temp_dir, dir) + .ok_or_else(|| io::Error::other("temp dir path should be UTF-8"))?; + let dir = Utf8Dir::open_ambient_dir(path_str, ambient_authority())?; + Ok((temp_dir, dir)) } #[rstest] -fn load_valid_rsa_key_succeeds(valid_rsa_pem: String, temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn load_valid_rsa_key_succeeds( + valid_rsa_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write("key.pem", &valid_rsa_pem) .expect("should write key"); let path = Utf8Path::new("/display/key.pem"); @@ -52,8 +62,8 @@ fn load_valid_rsa_key_succeeds(valid_rsa_pem: String, temp_key_dir: (TempDir, Ut } #[rstest] -fn load_missing_file_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn load_missing_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); let path = Utf8Path::new("/config/missing.pem"); let result = load_private_key_from_dir(&dir, "missing.pem", path); assert!(result.is_err(), "expected Err for missing file"); @@ -66,8 +76,8 @@ fn load_missing_file_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { } #[rstest] -fn load_empty_file_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write("empty.pem", "").expect("should write empty file"); let path = Utf8Path::new("/config/empty.pem"); let result = load_private_key_from_dir(&dir, "empty.pem", path); @@ -80,8 +90,8 @@ fn load_empty_file_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { } #[rstest] -fn load_invalid_pem_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn load_invalid_pem_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write("garbage.pem", "this is not a PEM file at all") .expect("should write garbage file"); let path = Utf8Path::new("/config/garbage.pem"); @@ -95,8 +105,8 @@ fn load_invalid_pem_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { } #[rstest] -fn load_ec_key_returns_clear_error(ec_pem: String, temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn load_ec_key_returns_clear_error(ec_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write("ec.pem", &ec_pem).expect("should write EC key"); let path = Utf8Path::new("/config/ec.pem"); let result = load_private_key_from_dir(&dir, "ec.pem", path); @@ -113,8 +123,11 @@ fn load_ec_key_returns_clear_error(ec_pem: String, temp_key_dir: (TempDir, Utf8D } #[rstest] -fn load_ed25519_key_returns_clear_error(ed25519_pem: String, temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn load_ed25519_key_returns_clear_error( + ed25519_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write("ed25519.pem", &ed25519_pem) .expect("should write Ed25519 key"); let path = Utf8Path::new("/config/ed25519.pem"); @@ -128,8 +141,8 @@ fn load_ed25519_key_returns_clear_error(ed25519_pem: String, temp_key_dir: (Temp } #[rstest] -fn error_includes_file_path(temp_key_dir: (TempDir, Utf8Dir)) { - let (_tmp, dir) = temp_key_dir; +fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); let display = Utf8Path::new("/home/user/.config/podbot/app.pem"); let result = load_private_key_from_dir(&dir, "nonexistent.pem", display); match result { @@ -145,8 +158,11 @@ fn error_includes_file_path(temp_key_dir: (TempDir, Utf8Dir)) { } #[rstest] -fn load_private_key_resolves_full_path(valid_rsa_pem: String, temp_key_dir: (TempDir, Utf8Dir)) { - let (tmp, dir) = temp_key_dir; +fn load_private_key_resolves_full_path( + valid_rsa_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) { + let (tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write("github-app.pem", &valid_rsa_pem) .expect("should write key"); let full_path = tmp.path().join("github-app.pem"); @@ -217,12 +233,12 @@ fn load_private_key_missing_parent_returns_error() { "encrypted" )] fn load_invalid_key_types_return_clear_error( - temp_key_dir: (TempDir, Utf8Dir), + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, #[case] file_name: &str, #[case] pem_content: &str, #[case] expected_keyword: &str, ) { - let (_tmp, dir) = temp_key_dir; + let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); dir.write(file_name, pem_content) .expect("should write key file"); let display = format!("/config/{file_name}"); @@ -235,206 +251,3 @@ fn load_invalid_key_types_return_clear_error( "error for {file_name} should mention '{expected_keyword}': {message}" ); } - -#[rstest] -fn build_app_client_with_valid_key_succeeds( - valid_rsa_pem: String, - temp_key_dir: (TempDir, Utf8Dir), -) { - let (_tmp, dir) = temp_key_dir; - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); - let path = Utf8Path::new("/display/key.pem"); - let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); - // Octocrab's build() spawns a Tower buffer task requiring a Tokio runtime. - let rt = tokio::runtime::Runtime::new().expect("should create tokio runtime"); - let _guard = rt.enter(); - let result = build_app_client(12345, key); - assert!(result.is_ok(), "expected Ok, got: {result:?}"); -} - -#[rstest] -fn build_app_client_with_zero_app_id_succeeds( - valid_rsa_pem: String, - temp_key_dir: (TempDir, Utf8Dir), -) { - let (_tmp, dir) = temp_key_dir; - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); - let path = Utf8Path::new("/display/key.pem"); - let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); - // Builder does not validate app_id; GitHub validates at token time. - // Octocrab's build() spawns a Tower buffer task requiring a Tokio runtime. - let rt = tokio::runtime::Runtime::new().expect("should create tokio runtime"); - let _guard = rt.enter(); - let result = build_app_client(0, key); - assert!( - result.is_ok(), - "expected Ok even with zero app_id, got: {result:?}" - ); -} - -#[rstest] -fn build_app_client_without_runtime_returns_error( - valid_rsa_pem: String, - temp_key_dir: (TempDir, Utf8Dir), -) { - let (_tmp, dir) = temp_key_dir; - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); - let path = Utf8Path::new("/display/key.pem"); - let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); - // Call without entering a Tokio runtime — should return Err, not panic. - let result = build_app_client(42, key); - assert!(result.is_err(), "expected Err without runtime, got Ok"); - let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( - message.contains("no Tokio runtime context"), - "error should mention missing runtime: {message}" - ); -} - -#[rstest] -#[case::client_error(http::StatusCode::TOO_MANY_REQUESTS, "4xx")] -#[case::server_error(http::StatusCode::INTERNAL_SERVER_ERROR, "5xx")] -#[case::redirect(http::StatusCode::TEMPORARY_REDIRECT, "3xx")] -fn github_status_class_groups_status_codes( - #[case] status_code: http::StatusCode, - #[case] expected_class: &str, -) { - assert_eq!(github_status_class(status_code), expected_class); -} - -#[rstest] -#[case::builder_context( - "failed to build GitHub App client: test error", - "failed to build GitHub App client" -)] -#[case::validation_context( - "failed to validate GitHub App credentials: test error", - "failed to validate GitHub App credentials" -)] -fn authentication_failed_error_includes_context( - #[case] message: &str, - #[case] expected_context: &str, -) { - let error = GitHubError::AuthenticationFailed { - message: String::from(message), - }; - let display = error.to_string(); - assert!( - display.contains(expected_context), - "error should include context: {display}" - ); - assert!( - display.contains("test error"), - "error should include cause: {display}" - ); -} - -#[rstest] -#[tokio::test] -async fn validate_app_credentials_with_missing_key_returns_error(temp_key_dir: (TempDir, Utf8Dir)) { - let (temp_dir, _dir) = temp_key_dir; - let key_path = Utf8Path::from_path(temp_dir.path()) - .expect("temp dir path should be UTF-8") - .join("key.pem"); - let result = validate_app_credentials(12345, &key_path).await; - assert!(result.is_err(), "expected Err for missing key file"); - match result { - Err(GitHubError::PrivateKeyLoadFailed { ref path, .. }) => { - assert!( - path.to_string_lossy().contains("key.pem"), - "error path should reference the missing file" - ); - } - other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), - } -} - -#[rstest] -#[tokio::test] -async fn validate_app_credentials_with_invalid_pem_returns_error( - ec_pem: String, - temp_key_dir: (TempDir, Utf8Dir), -) { - let (tmp, dir) = temp_key_dir; - dir.write("ec.pem", &ec_pem).expect("should write EC key"); - let full_path = tmp.path().join("ec.pem"); - let utf8_path = Utf8Path::from_path(&full_path).expect("temp path should be UTF-8"); - - let result = validate_app_credentials(12345, utf8_path).await; - assert!(result.is_err(), "expected Err for ECDSA key"); - match result { - Err(GitHubError::PrivateKeyLoadFailed { message, .. }) => { - assert!( - message.contains("ECDSA"), - "error should mention ECDSA: {message}" - ); - } - other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), - } -} - -#[rstest] -#[tokio::test] -async fn validate_with_client_propagates_mock_success() { - let mut mock = MockGitHubAppClient::new(); - mock.expect_validate_credentials() - .times(1) - .returning(|| Box::pin(async { Ok(()) })); - - let result = validate_with_client(&mock).await; - assert!(result.is_ok(), "expected Ok from mock client"); -} - -#[rstest] -#[tokio::test] -async fn validate_with_client_propagates_mock_error() { - let mut mock = MockGitHubAppClient::new(); - mock.expect_validate_credentials().times(1).returning(|| { - Box::pin(async { - Err(GitHubError::AuthenticationFailed { - message: String::from("mock authentication failure"), - }) - }) - }); - - let result = validate_with_client(&mock).await; - assert!(result.is_err(), "expected Err from mock client"); - let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( - message.contains("mock authentication failure"), - "error should propagate mock message: {message}" - ); -} - -mod property_tests { - use proptest::prelude::*; - - use super::super::retry_metrics::github_status_class; - - proptest! { - #[test] - fn github_status_class_range_invariant(raw_code in 100_u16..=999_u16) { - let status = http::StatusCode::from_u16(raw_code) - .expect("codes 100–999 are valid HTTP status codes"); - let class = github_status_class(status); - let expected = match raw_code { - 100..=199 => "1xx", - 200..=299 => "2xx", - 300..=399 => "3xx", - 400..=499 => "4xx", - 500..=599 => "5xx", - _ => "other", - }; - prop_assert_eq!( - class, - expected, - "github_status_class({}) should map to {}", - raw_code, - expected, - ); - } - } -} diff --git a/tests/bdd_config_helpers.rs b/tests/bdd_config_helpers.rs index 4e62de93..9d0a9f1e 100644 --- a/tests/bdd_config_helpers.rs +++ b/tests/bdd_config_helpers.rs @@ -3,15 +3,13 @@ #![cfg(feature = "internal")] use camino::Utf8PathBuf; -use ortho_config::MergeComposer; -use ortho_config::serde_json::json; use podbot::config::{ AgentKind, AgentMode, AppConfig, GitHubConfig, SandboxConfig, WorkspaceConfig, WorkspaceSource, }; use podbot::error::{ConfigError, PodbotError}; use rstest::fixture; use rstest_bdd::Slot; -use rstest_bdd_macros::{ScenarioState, given, then, when}; +use rstest_bdd_macros::{ScenarioState, given, then}; // Helper functions to reduce duplication in step definitions @@ -364,166 +362,7 @@ fn dev_fuse_mounting_disabled(config_state: &ConfigState) { ); } -// Layer precedence step definitions - -/// Recursively merges two JSON values, combining nested objects field-by-field. -/// -/// For object values, fields are merged recursively. For non-objects, the new -/// value completely overwrites the existing one. This mirrors how `OrthoConfig` -/// merges nested configuration structures. -fn merge_json_values( - existing: &ortho_config::serde_json::Value, - new_value: &ortho_config::serde_json::Value, -) -> ortho_config::serde_json::Value { - use ortho_config::serde_json::Value; - - match (existing, new_value) { - (Value::Object(existing_obj), Value::Object(new_obj)) => { - let mut merged = existing_obj.clone(); - for (key, new_child) in new_obj { - if let Some(existing_child) = merged.get(key) { - // Recursively merge nested objects; for non-objects the new value wins. - merged.insert(key.clone(), merge_json_values(existing_child, new_child)); - } else { - merged.insert(key.clone(), new_child.clone()); - } - } - Value::Object(merged) - } - // For non-object values, the new value completely overwrites the existing one. - _ => new_value.clone(), - } -} - -/// Merges a new value into an existing layer slot (if present). -fn merge_layer( - existing: Option, - new_value: ortho_config::serde_json::Value, -) -> ortho_config::serde_json::Value { - if let Some(existing_value) = existing { - merge_json_values(&existing_value, &new_value) - } else { - new_value - } -} - -/// Merges the current file layer with a new value (combining fields). -fn merge_file_layer( - config_state: &ConfigState, - new_value: ortho_config::serde_json::Value, -) -> ortho_config::serde_json::Value { - merge_layer(config_state.file_layer.get(), new_value) -} - -/// Merges the current env layer with a new value (combining fields). -fn merge_env_layer( - config_state: &ConfigState, - new_value: ortho_config::serde_json::Value, -) -> ortho_config::serde_json::Value { - merge_layer(config_state.env_layer.get(), new_value) -} - -#[given("defaults provide engine_socket as nil")] -fn defaults_provide_engine_socket_nil(config_state: &ConfigState) { - // Defaults already have engine_socket as None, nothing to do. - // Access config_state to satisfy clippy (rstest_bdd requires the parameter). - drop(config_state.file_layer.get()); -} - -#[given("a file layer provides engine_socket as {socket}")] -fn file_layer_provides_engine_socket(config_state: &ConfigState, socket: String) { - let merged = merge_file_layer(config_state, json!({ "engine_socket": socket })); - config_state.file_layer.set(merged); -} - -#[given("an environment layer provides engine_socket as {socket}")] -fn env_layer_provides_engine_socket(config_state: &ConfigState, socket: String) { - let merged = merge_env_layer(config_state, json!({ "engine_socket": socket })); - config_state.env_layer.set(merged); -} - -#[given("a CLI layer provides engine_socket as {socket}")] -fn cli_layer_provides_engine_socket(config_state: &ConfigState, socket: String) { - config_state - .cli_layer - .set(json!({ "engine_socket": socket })); -} - -#[given("a file layer provides image as {image}")] -fn file_layer_provides_image(config_state: &ConfigState, image: String) { - let merged = merge_file_layer(config_state, json!({ "image": image })); - config_state.file_layer.set(merged); -} - -#[given("a file layer provides sandbox.privileged as {value}")] -fn file_layer_provides_sandbox_privileged(config_state: &ConfigState, value: bool) { - let merged = merge_file_layer(config_state, json!({ "sandbox": { "privileged": value } })); - config_state.file_layer.set(merged); -} - -#[given("a file layer provides sandbox.mount_dev_fuse as {value}")] -fn file_layer_provides_sandbox_mount_dev_fuse(config_state: &ConfigState, value: bool) { - let merged = merge_file_layer( - config_state, - json!({ "sandbox": { "mount_dev_fuse": value } }), - ); - config_state.file_layer.set(merged); -} - -#[given("an environment layer provides sandbox.privileged as {value}")] -fn env_layer_provides_sandbox_privileged(config_state: &ConfigState, value: bool) { - let merged = merge_env_layer(config_state, json!({ "sandbox": { "privileged": value } })); - config_state.env_layer.set(merged); -} - -#[when("configuration is merged")] -#[expect(clippy::expect_used, reason = "test step - panics are acceptable")] -fn configuration_is_merged(config_state: &ConfigState) { - let mut composer = MergeComposer::new(); - - // Layer 1: Defaults - let defaults = ortho_config::serde_json::to_value(AppConfig::default()) - .expect("serialization should succeed"); - composer.push_defaults(defaults); - - // Layer 2: File - if let Some(file_layer) = config_state.file_layer.get() { - composer.push_file(file_layer, None); - } - - // Layer 3: Environment - if let Some(env_layer) = config_state.env_layer.get() { - composer.push_environment(env_layer); - } - - // Layer 4: CLI - if let Some(cli_layer) = config_state.cli_layer.get() { - composer.push_cli(cli_layer); - } - - let config: AppConfig = podbot::config::merge_from_layers_for_tests(composer.layers()) - .expect("merge should succeed"); - config_state.config.set(config); -} - -#[then("the engine socket is {socket}")] -fn engine_socket_is(config_state: &ConfigState, socket: String) { - let config = get_config(config_state); - assert_eq!( - config.engine_socket.as_deref(), - Some(socket.as_str()), - "Expected engine socket to be {}", - socket - ); -} - -#[then("the image is {image}")] -fn image_is(config_state: &ConfigState, image: String) { - let config = get_config(config_state); - assert_eq!( - config.image.as_deref(), - Some(image.as_str()), - "Expected image to be {}", - image - ); -} +// Layer precedence step definitions live in a sibling submodule to keep this +// module within the 400-line limit. +#[path = "bdd_config_steps/layer.rs"] +mod layer_steps; diff --git a/tests/bdd_config_loader_helpers.rs b/tests/bdd_config_loader_helpers.rs index aa6c438b..31c2c41f 100644 --- a/tests/bdd_config_loader_helpers.rs +++ b/tests/bdd_config_loader_helpers.rs @@ -112,9 +112,11 @@ fn write_config_file( ) -> StepResult { let temp_dir = tempfile::TempDir::new().map_err(|e| e.to_string())?; let temp_dir_arc = Arc::new(temp_dir); - let raw_path = temp_dir_arc.path().join(file_name); - std::fs::write(&raw_path, content).map_err(|e| e.to_string())?; + let dir = cap_std::fs::Dir::open_ambient_dir(temp_dir_arc.path(), cap_std::ambient_authority()) + .map_err(|e| e.to_string())?; + dir.write(file_name, content).map_err(|e| e.to_string())?; + let raw_path = temp_dir_arc.path().join(file_name); let path = Utf8PathBuf::try_from(raw_path).map_err(|e| e.to_string())?; config_loader_state.temp_dir.set(temp_dir_arc); Ok(path) diff --git a/tests/bdd_config_steps/layer.rs b/tests/bdd_config_steps/layer.rs new file mode 100644 index 00000000..a98d5c80 --- /dev/null +++ b/tests/bdd_config_steps/layer.rs @@ -0,0 +1,171 @@ +//! Layer-precedence step definitions for podbot configuration behaviour +//! tests: merging file, environment, and CLI layers over defaults. + +use ortho_config::MergeComposer; +use ortho_config::serde_json::json; +use podbot::config::AppConfig; +use rstest_bdd_macros::{given, then, when}; + +use super::{ConfigState, get_config}; + +/// Recursively merges two JSON values, combining nested objects field-by-field. +/// +/// For object values, fields are merged recursively. For non-objects, the new +/// value completely overwrites the existing one. This mirrors how `OrthoConfig` +/// merges nested configuration structures. +fn merge_json_values( + existing: &ortho_config::serde_json::Value, + new_value: &ortho_config::serde_json::Value, +) -> ortho_config::serde_json::Value { + use ortho_config::serde_json::Value; + + match (existing, new_value) { + (Value::Object(existing_obj), Value::Object(new_obj)) => { + let mut merged = existing_obj.clone(); + for (key, new_child) in new_obj { + if let Some(existing_child) = merged.get(key) { + // Recursively merge nested objects; for non-objects the new value wins. + merged.insert(key.clone(), merge_json_values(existing_child, new_child)); + } else { + merged.insert(key.clone(), new_child.clone()); + } + } + Value::Object(merged) + } + // For non-object values, the new value completely overwrites the existing one. + _ => new_value.clone(), + } +} + +/// Merges a new value into an existing layer slot (if present). +fn merge_layer( + existing: Option, + new_value: ortho_config::serde_json::Value, +) -> ortho_config::serde_json::Value { + if let Some(existing_value) = existing { + merge_json_values(&existing_value, &new_value) + } else { + new_value + } +} + +/// Merges the current file layer with a new value (combining fields). +fn merge_file_layer( + config_state: &ConfigState, + new_value: ortho_config::serde_json::Value, +) -> ortho_config::serde_json::Value { + merge_layer(config_state.file_layer.get(), new_value) +} + +/// Merges the current env layer with a new value (combining fields). +fn merge_env_layer( + config_state: &ConfigState, + new_value: ortho_config::serde_json::Value, +) -> ortho_config::serde_json::Value { + merge_layer(config_state.env_layer.get(), new_value) +} + +#[given("defaults provide engine_socket as nil")] +fn defaults_provide_engine_socket_nil(config_state: &ConfigState) { + // Defaults already have engine_socket as None, nothing to do. + // Access config_state to satisfy clippy (rstest_bdd requires the parameter). + drop(config_state.file_layer.get()); +} + +#[given("a file layer provides engine_socket as {socket}")] +fn file_layer_provides_engine_socket(config_state: &ConfigState, socket: String) { + let merged = merge_file_layer(config_state, json!({ "engine_socket": socket })); + config_state.file_layer.set(merged); +} + +#[given("an environment layer provides engine_socket as {socket}")] +fn env_layer_provides_engine_socket(config_state: &ConfigState, socket: String) { + let merged = merge_env_layer(config_state, json!({ "engine_socket": socket })); + config_state.env_layer.set(merged); +} + +#[given("a CLI layer provides engine_socket as {socket}")] +fn cli_layer_provides_engine_socket(config_state: &ConfigState, socket: String) { + config_state + .cli_layer + .set(json!({ "engine_socket": socket })); +} + +#[given("a file layer provides image as {image}")] +fn file_layer_provides_image(config_state: &ConfigState, image: String) { + let merged = merge_file_layer(config_state, json!({ "image": image })); + config_state.file_layer.set(merged); +} + +#[given("a file layer provides sandbox.privileged as {value}")] +fn file_layer_provides_sandbox_privileged(config_state: &ConfigState, value: bool) { + let merged = merge_file_layer(config_state, json!({ "sandbox": { "privileged": value } })); + config_state.file_layer.set(merged); +} + +#[given("a file layer provides sandbox.mount_dev_fuse as {value}")] +fn file_layer_provides_sandbox_mount_dev_fuse(config_state: &ConfigState, value: bool) { + let merged = merge_file_layer( + config_state, + json!({ "sandbox": { "mount_dev_fuse": value } }), + ); + config_state.file_layer.set(merged); +} + +#[given("an environment layer provides sandbox.privileged as {value}")] +fn env_layer_provides_sandbox_privileged(config_state: &ConfigState, value: bool) { + let merged = merge_env_layer(config_state, json!({ "sandbox": { "privileged": value } })); + config_state.env_layer.set(merged); +} + +#[when("configuration is merged")] +#[expect(clippy::expect_used, reason = "test step - panics are acceptable")] +fn configuration_is_merged(config_state: &ConfigState) { + let mut composer = MergeComposer::new(); + + // Layer 1: Defaults + let defaults = ortho_config::serde_json::to_value(AppConfig::default()) + .expect("serialization should succeed"); + composer.push_defaults(defaults); + + // Layer 2: File + if let Some(file_layer) = config_state.file_layer.get() { + composer.push_file(file_layer, None); + } + + // Layer 3: Environment + if let Some(env_layer) = config_state.env_layer.get() { + composer.push_environment(env_layer); + } + + // Layer 4: CLI + if let Some(cli_layer) = config_state.cli_layer.get() { + composer.push_cli(cli_layer); + } + + let config: AppConfig = podbot::config::merge_from_layers_for_tests(composer.layers()) + .expect("merge should succeed"); + config_state.config.set(config); +} + +#[then("the engine socket is {socket}")] +fn engine_socket_is(config_state: &ConfigState, socket: String) { + let config = get_config(config_state); + assert_eq!( + config.engine_socket.as_deref(), + Some(socket.as_str()), + "Expected engine socket to be {}", + socket + ); +} + +#[then("the image is {image}")] +fn image_is(config_state: &ConfigState, image: String) { + let config = get_config(config_state); + assert_eq!( + config.image.as_deref(), + Some(image.as_str()), + "Expected image to be {}", + image + ); +} diff --git a/tests/bdd_hosting_config_loader_helpers.rs b/tests/bdd_hosting_config_loader_helpers.rs index ebb0deea..cb720745 100644 --- a/tests/bdd_hosting_config_loader_helpers.rs +++ b/tests/bdd_hosting_config_loader_helpers.rs @@ -203,8 +203,11 @@ fn write_config_file( ) -> StepResult { let temp_dir = tempfile::TempDir::new().map_err(|error| error.to_string())?; let temp_dir_arc = Arc::new(temp_dir); + let dir = cap_std::fs::Dir::open_ambient_dir(temp_dir_arc.path(), cap_std::ambient_authority()) + .map_err(|error| error.to_string())?; + dir.write("config.toml", content) + .map_err(|error| error.to_string())?; let raw_path = temp_dir_arc.path().join("config.toml"); - std::fs::write(&raw_path, content).map_err(|error| error.to_string())?; let path = Utf8PathBuf::try_from(raw_path).map_err(|error| error.to_string())?; hosting_config_loader_state.temp_dir.set(temp_dir_arc); Ok(path) diff --git a/tests/make_audit_target.rs b/tests/make_audit_target.rs index 9ca7dd74..7e6407e8 100644 --- a/tests/make_audit_target.rs +++ b/tests/make_audit_target.rs @@ -15,34 +15,44 @@ //! See also: `Makefile` (`rust-audit` target), `docs/developers-guide.md` //! (§ 2 Quality gates, § 2.1 Security audit ignores). -use std::fs; -#[cfg(unix)] -use std::os::unix::fs::PermissionsExt; use std::path::Path; use std::process::Command; +use cap_std::fs::Dir; use rstest::rstest; use tempfile::TempDir; type TestResult = Result>; -fn write_file(path: &Path, contents: &str) -> TestResult { - if let Some(parent) = path.parent() { - fs::create_dir_all(parent)?; +/// Open a capability handle on the temporary workspace root. +fn open_workspace_dir(workspace: &Path) -> TestResult { + Ok(Dir::open_ambient_dir( + workspace, + cap_std::ambient_authority(), + )?) +} + +fn write_file(workspace_dir: &Dir, relative_path: &str, contents: &str) -> TestResult { + if let Some(parent) = Path::new(relative_path).parent() + && !parent.as_os_str().is_empty() + { + workspace_dir.create_dir_all(parent)?; } - fs::write(path, contents)?; + workspace_dir.write(relative_path, contents)?; Ok(()) } +/// Writes the fake `cargo` script at `bin/cargo` inside the workspace; the +/// caller derives its absolute path with `workspace.join("bin/cargo")`. fn write_fake_cargo( - bin_dir: &Path, + workspace_dir: &Dir, log_path: &Path, exit_status: i32, metadata_status: i32, -) -> TestResult { - let cargo_path = bin_dir.join("cargo"); +) -> TestResult { write_file( - &cargo_path, + workspace_dir, + "bin/cargo", &format!( "#!/usr/bin/env sh\nif [ \"$1\" = metadata ]; then\nprintf '%s\\n' \"$PODBOT_FAKE_CARGO_METADATA\"\nexit {metadata_status}\nfi\nprintf '%s|%s\\n' \"$PWD\" \"$*\" >> '{}'\nexit {exit_status}\n", log_path.display() @@ -50,11 +60,12 @@ fn write_fake_cargo( )?; #[cfg(unix)] { - let mut permissions = fs::metadata(&cargo_path)?.permissions(); - permissions.set_mode(0o755); - fs::set_permissions(&cargo_path, permissions)?; + use cap_std::fs::PermissionsExt; + + let permissions = cap_std::fs::Permissions::from_mode(0o755); + workspace_dir.set_permissions("bin/cargo", permissions)?; } - Ok(cargo_path) + Ok(()) } fn run_rust_audit( @@ -72,8 +83,12 @@ fn run_rust_audit( .output()?) } -fn create_manifest(path: &Path) -> TestResult { - write_file(path, "[package]\nname = \"fixture\"\nversion = \"0.0.0\"\n") +fn create_manifest(workspace_dir: &Dir, relative_path: &str) -> TestResult { + write_file( + workspace_dir, + relative_path, + "[package]\nname = \"fixture\"\nversion = \"0.0.0\"\n", + ) } fn cargo_metadata_for(workspace: &Path, manifests: &[&Path]) -> String { @@ -104,20 +119,22 @@ fn cargo_metadata_for(workspace: &Path, manifests: &[&Path]) -> String { fn rust_audit_invokes_cargo_audit_once_at_workspace_root() { let temp = TempDir::new().expect("temporary workspace should be created"); let workspace = temp.path(); + let workspace_dir = open_workspace_dir(workspace).expect("workspace dir should open"); let root_manifest = workspace.join("Cargo.toml"); let member_manifest = workspace.join("crates/agent/Cargo.toml"); - create_manifest(&root_manifest).expect("root manifest should be created"); - create_manifest(&member_manifest).expect("nested manifest should be created"); - create_manifest(&workspace.join("target/ignored/Cargo.toml")) + create_manifest(&workspace_dir, "Cargo.toml").expect("root manifest should be created"); + create_manifest(&workspace_dir, "crates/agent/Cargo.toml") + .expect("nested manifest should be created"); + create_manifest(&workspace_dir, "target/ignored/Cargo.toml") .expect("target manifest fixture should be created"); - create_manifest(&workspace.join("node_modules/ignored/Cargo.toml")) + create_manifest(&workspace_dir, "node_modules/ignored/Cargo.toml") .expect("node_modules manifest fixture should be created"); - create_manifest(&workspace.join(".venv/ignored/Cargo.toml")) + create_manifest(&workspace_dir, ".venv/ignored/Cargo.toml") .expect("virtualenv manifest fixture should be created"); let log_path = workspace.join("cargo-audit.log"); - let fake_cargo = write_fake_cargo(&workspace.join("bin"), &log_path, 0, 0) - .expect("fake cargo should be built"); + write_fake_cargo(&workspace_dir, &log_path, 0, 0).expect("fake cargo should be built"); + let fake_cargo = workspace.join("bin/cargo"); let metadata = cargo_metadata_for(workspace, &[&root_manifest, &member_manifest]); let output = @@ -129,7 +146,9 @@ fn rust_audit_invokes_cargo_audit_once_at_workspace_root() { String::from_utf8_lossy(&output.stdout), String::from_utf8_lossy(&output.stderr) ); - let log = fs::read_to_string(log_path).expect("fake cargo log should be readable"); + let log = workspace_dir + .read_to_string("cargo-audit.log") + .expect("fake cargo log should be readable"); assert_eq!( log, format!("{}|audit\n", workspace.display()), @@ -160,16 +179,18 @@ fn rust_audit_propagates_failure( ) { let temp = TempDir::new().expect("temporary workspace should be created"); let workspace = temp.path(); - create_manifest(&workspace.join("Cargo.toml")).expect("root manifest should be created"); + let workspace_dir = open_workspace_dir(workspace).expect("workspace dir should open"); + create_manifest(&workspace_dir, "Cargo.toml").expect("root manifest should be created"); let log_path = workspace.join("cargo-audit.log"); - let fake_cargo = write_fake_cargo( - &workspace.join("bin"), + write_fake_cargo( + &workspace_dir, &log_path, audit_exit_status, metadata_exit_status, ) .expect("fake cargo should be built"); + let fake_cargo = workspace.join("bin/cargo"); let metadata = cargo_metadata_for(workspace, &[&workspace.join("Cargo.toml")]); let output = @@ -182,7 +203,9 @@ fn rust_audit_propagates_failure( String::from_utf8_lossy(&output.stderr), ); if should_audit_run { - let log = fs::read_to_string(&log_path).expect("fake cargo log should be readable"); + let log = workspace_dir + .read_to_string("cargo-audit.log") + .expect("fake cargo log should be readable"); assert_eq!( log, format!("{}|audit\n", workspace.display()), @@ -190,7 +213,7 @@ fn rust_audit_propagates_failure( ); } else { assert!( - !log_path.exists(), + !workspace_dir.exists("cargo-audit.log"), "cargo audit should not run after metadata failure" ); } diff --git a/tests/test_utils.rs b/tests/test_utils.rs index ed71c6d6..71535ec8 100644 --- a/tests/test_utils.rs +++ b/tests/test_utils.rs @@ -200,6 +200,8 @@ pub fn clean_env() -> EnvGuard<'static> { #[cfg(test)] mod tests { + //! Unit tests for the shared integration-test utilities. + use bollard::exec::{CreateExecOptions, CreateExecResults, StartExecOptions, StartExecResults}; use bollard::models::ExecInspectResponse; use mockall::mock; From 6fa2ec4cb576ebc97e3082067171287fb8f24ca2 Mon Sep 17 00:00:00 2001 From: leynos Date: Thu, 9 Jul 2026 18:45:26 +0200 Subject: [PATCH 2/5] Adopt the Whitaker Dylint suite in the lint gate and CI Wire the Whitaker Dylint suite into the standard quality gates as part of the estate-wide rollout (see leynos/netsuke#410): - Makefile: add a WHITAKER tool variable and run the suite after Clippy in the lint target, with warnings denied. - CI: pin whitaker-installer to the 0.2.5 crates.io release via a job-level WHITAKER_INSTALLER_VERSION environment variable, cache the installer binary and the cargo-binstall download cache keyed by runner OS, architecture, and installer version, and install with cargo binstall --locked, falling back to building from crates.io on runners without binstall. --- .github/workflows/ci.yml | 19 +++++++++++++++++++ Makefile | 4 +++- 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 2d0f210e..ea47ee22 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -13,6 +13,7 @@ jobs: env: CARGO_TERM_COLOR: always BUILD_PROFILE: debug + WHITAKER_INSTALLER_VERSION: '0.2.5' steps: - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 - name: Setup Rust @@ -41,6 +42,24 @@ jobs: - name: Audit dependencies if: github.actor != 'dependabot[bot]' run: make audit + - name: Cache Whitaker installer + uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 + with: + path: | + ~/.cargo/bin/whitaker-installer + ~/.cache/cargo-binstall + key: whitaker-installer-${{ runner.os }}-${{ runner.arch }}-${{ env.WHITAKER_INSTALLER_VERSION }} + - name: Install Whitaker + run: | + if ! command -v whitaker-installer >/dev/null 2>&1; then + if cargo binstall --version >/dev/null 2>&1; then + cargo binstall --no-confirm --locked "whitaker-installer@${WHITAKER_INSTALLER_VERSION}" + else + echo "cargo-binstall unavailable; building whitaker-installer from crates.io" + cargo install --locked whitaker-installer --version "${WHITAKER_INSTALLER_VERSION}" + fi + fi + whitaker-installer - name: Lint run: make lint - name: No-CLI compile check diff --git a/Makefile b/Makefile index 5d2391e9..1fba0a03 100644 --- a/Makefile +++ b/Makefile @@ -11,6 +11,7 @@ CARGO_FLAGS ?= --all-targets --all-features CLIPPY_FLAGS ?= $(CARGO_FLAGS) -- $(RUST_FLAGS) TEST_FLAGS ?= $(CARGO_FLAGS) MDLINT ?= $(shell command -v markdownlint-cli2 2>/dev/null || printf '%s' "$$HOME/.bun/bin/markdownlint-cli2") +WHITAKER ?= whitaker NIXIE ?= nixie build: target/debug/$(TARGET) ## Build debug binary @@ -27,9 +28,10 @@ test: ## Run tests with warnings treated as errors target/%/$(TARGET): ## Build binary in debug or release mode $(CARGO) build $(BUILD_JOBS) $(if $(findstring release,$(@)),--release) --bin $(TARGET) -lint: ## Run Clippy with warnings denied +lint: ## Run Clippy and the Whitaker Dylint suite with warnings denied RUSTDOCFLAGS="$(RUSTDOC_FLAGS)" $(CARGO) doc --no-deps $(CARGO) clippy $(CLIPPY_FLAGS) + RUSTFLAGS="$(RUST_FLAGS)" $(WHITAKER) --all -- $(CARGO_FLAGS) typecheck: ## Type-check the workspace RUSTFLAGS="$(RUST_FLAGS)" $(CARGO) check $(CARGO_FLAGS) $(BUILD_JOBS) From fff38edbab7eebb438ead8b7a97b325cebeaa07e Mon Sep 17 00:00:00 2001 From: leynos Date: Thu, 9 Jul 2026 19:21:19 +0200 Subject: [PATCH 3/5] Address review feedback on test structure and archive metadata Consolidate the ACP test doubles and frame builders into a shared acp_test_support module: one poison-tolerant RecordingWriter replaces the three near-identical recorders, and one jsonrpc_frame builder replaces the per-module frame constructors, resolving the CodeScene code-duplication findings on acp_frame_tests, protocol_acp_policy_integration_tests, and protocol_acp_tests. The two structurally similar non-enforcing policy tests are now a single parameterized rstest case, and the two initialize-frame builders share a blocked_capabilities helper. Also, per reviewer feedback: - read tar entry metadata from the open file handle instead of a prior directory call, so the header size cannot drift from the streamed bytes if the file changes between the two calls; - consume fallible rstest fixtures by returning Result from the tests and propagating with `?` rather than unwrapping inline; - pass the frame value to serialize_frame by reference, removing the explicit drop; - return the mapped ConfigError directly instead of wrapping with Ok(...?); - name the fake cargo script path once via a constant. --- src/config/loader.rs | 9 +- src/engine/connection/exec/acp_frame_tests.rs | 19 +--- .../connection/exec/acp_runtime_bdd_tests.rs | 9 +- .../connection/exec/acp_runtime_tests.rs | 83 +------------- .../connection/exec/acp_test_support.rs | 104 +++++++++++++++++ src/engine/connection/exec/mod.rs | 2 + .../exec/protocol_acp_forwarding_tests.rs | 7 +- .../protocol_acp_policy_integration_tests.rs | 105 ++++-------------- .../exec/protocol_acp_routing_tests.rs | 28 ++--- .../connection/exec/protocol_acp_tests.rs | 88 ++++----------- src/engine/connection/exec/runtime_helpers.rs | 5 +- .../connection/upload_credentials/archive.rs | 5 +- src/github/client_tests.rs | 40 +++---- src/github/tests.rs | 67 ++++++----- tests/make_audit_target.rs | 13 ++- 15 files changed, 248 insertions(+), 336 deletions(-) create mode 100644 src/engine/connection/exec/acp_test_support.rs diff --git a/src/config/loader.rs b/src/config/loader.rs index f32a6406..168c1737 100644 --- a/src/config/loader.rs +++ b/src/config/loader.rs @@ -249,9 +249,10 @@ fn serialize_agent_override( field: &str, value: T, ) -> Result { - Ok( - serde_json::to_value(value).map_err(|error| ConfigError::ParseError { + serde_json::to_value(value).map_err(|error| { + ConfigError::ParseError { message: format!("failed to serialize agent {field} override: {error}"), - })?, - ) + } + .into() + }) } diff --git a/src/engine/connection/exec/acp_frame_tests.rs b/src/engine/connection/exec/acp_frame_tests.rs index b4b8dbd0..7d760059 100644 --- a/src/engine/connection/exec/acp_frame_tests.rs +++ b/src/engine/connection/exec/acp_frame_tests.rs @@ -13,27 +13,14 @@ use super::{ OutboundFrameAssembler, }; use crate::engine::connection::exec::acp_policy::MethodDenylist; +use crate::engine::connection::exec::acp_test_support::jsonrpc_frame; fn permitted_frame(method: &str, line_ending: &[u8]) -> Result, serde_json::Error> { - let mut bytes = serde_json::to_vec(&serde_json::json!({ - "jsonrpc": "2.0", - "id": 1, - "method": method, - "params": {}, - }))?; - bytes.extend_from_slice(line_ending); - Ok(bytes) + jsonrpc_frame(Some(&serde_json::json!(1)), method, line_ending) } fn blocked_request_frame(id: &Value, method: &str) -> Result, serde_json::Error> { - let mut bytes = serde_json::to_vec(&serde_json::json!({ - "jsonrpc": "2.0", - "id": id, - "method": method, - "params": {}, - }))?; - bytes.push(b'\n'); - Ok(bytes) + jsonrpc_frame(Some(id), method, b"\n") } fn assembler() -> OutboundFrameAssembler { diff --git a/src/engine/connection/exec/acp_runtime_bdd_tests.rs b/src/engine/connection/exec/acp_runtime_bdd_tests.rs index 15fa4f2f..5b35512c 100644 --- a/src/engine/connection/exec/acp_runtime_bdd_tests.rs +++ b/src/engine/connection/exec/acp_runtime_bdd_tests.rs @@ -68,16 +68,15 @@ fn denylist_state() -> DenylistState { /// Serializes `value` to compact JSON bytes and appends a newline terminator, /// producing a well-formed ACP frame suitable for test input. -fn serialize_frame(value: serde_json::Value) -> StepResult> { +fn serialize_frame(value: &serde_json::Value) -> StepResult> { let mut bytes = - serde_json::to_vec(&value).map_err(|err| format!("frame serialization failed: {err}"))?; - drop(value); + serde_json::to_vec(value).map_err(|err| format!("frame serialization failed: {err}"))?; bytes.push(b'\n'); Ok(bytes) } fn make_request_frame(method: &str, id: i64) -> StepResult> { - serialize_frame(serde_json::json!({ + serialize_frame(&serde_json::json!({ "jsonrpc": "2.0", "id": id, "method": method, @@ -86,7 +85,7 @@ fn make_request_frame(method: &str, id: i64) -> StepResult> { } fn make_notification_frame(method: &str) -> StepResult> { - serialize_frame(serde_json::json!({ + serialize_frame(&serde_json::json!({ "jsonrpc": "2.0", "method": method, "params": {}, diff --git a/src/engine/connection/exec/acp_runtime_tests.rs b/src/engine/connection/exec/acp_runtime_tests.rs index 18eff121..b433bd3a 100644 --- a/src/engine/connection/exec/acp_runtime_tests.rs +++ b/src/engine/connection/exec/acp_runtime_tests.rs @@ -3,7 +3,6 @@ use std::io; use std::pin::Pin; -use std::sync::{Arc, Mutex, PoisonError}; use std::task::{Context, Poll}; use ortho_config::serde_json::{self, Value}; @@ -16,55 +15,7 @@ use super::{ run_container_stdin_sink, }; use crate::engine::connection::exec::acp_policy::MethodDenylist; - -/// Recording host-stdout writer that captures every byte and tracks shutdown. -#[derive(Clone, Default)] -struct RecordingWriter { - bytes: Arc>>, - shutdown_called: Arc>, -} - -impl RecordingWriter { - fn snapshot(&self) -> Vec { - self.bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .clone() - } - - fn shutdown_observed(&self) -> bool { - *self - .shutdown_called - .lock() - .unwrap_or_else(PoisonError::into_inner) - } -} - -impl AsyncWrite for RecordingWriter { - fn poll_write( - self: Pin<&mut Self>, - _cx: &mut Context<'_>, - buf: &[u8], - ) -> Poll> { - self.bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .extend_from_slice(buf); - Poll::Ready(Ok(buf.len())) - } - - fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - Poll::Ready(Ok(())) - } - - fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - *self - .shutdown_called - .lock() - .unwrap_or_else(PoisonError::into_inner) = true; - Poll::Ready(Ok(())) - } -} +use crate::engine::connection::exec::acp_test_support::{RecordingWriter, jsonrpc_frame}; /// Writer that always returns `BrokenPipe` on writes, used to simulate the /// agent having already exited. @@ -91,42 +42,16 @@ impl AsyncWrite for BrokenPipeWriter { } } -/// Builds a newline-terminated JSON-RPC 2.0 frame. -/// -/// Pass `id = Some(…)` for requests; `id = None` for notifications. -fn make_jsonrpc_frame(method: &str, id: Option<&Value>) -> Result, serde_json::Error> { - let value = id.map_or_else( - || { - serde_json::json!({ - "jsonrpc": "2.0", - "method": method, - "params": {}, - }) - }, - |request_id| { - serde_json::json!({ - "jsonrpc": "2.0", - "id": request_id, - "method": method, - "params": {}, - }) - }, - ); - let mut bytes = serde_json::to_vec(&value)?; - bytes.push(b'\n'); - Ok(bytes) -} - fn permitted_frame() -> Result, serde_json::Error> { - make_jsonrpc_frame("session/new", Some(&serde_json::json!(1))) + jsonrpc_frame(Some(&serde_json::json!(1)), "session/new", b"\n") } fn blocked_request_frame(id: &Value) -> Result, serde_json::Error> { - make_jsonrpc_frame("terminal/create", Some(id)) + jsonrpc_frame(Some(id), "terminal/create", b"\n") } fn blocked_notification_frame() -> Result, serde_json::Error> { - make_jsonrpc_frame("fs/changed", None) + jsonrpc_frame(None, "fs/changed", b"\n") } fn build_adapter() -> (OutboundPolicyAdapter, mpsc::Receiver) { diff --git a/src/engine/connection/exec/acp_test_support.rs b/src/engine/connection/exec/acp_test_support.rs new file mode 100644 index 00000000..f34c917a --- /dev/null +++ b/src/engine/connection/exec/acp_test_support.rs @@ -0,0 +1,104 @@ +//! Shared test doubles and frame builders for the ACP test modules. +//! +//! Consolidates the recording writer used to capture host or container +//! output in tests, together with the newline-terminated JSON-RPC frame +//! builder, so the individual ACP test modules do not duplicate them. + +use std::io; +use std::pin::Pin; +use std::sync::{Arc, Mutex, PoisonError}; +use std::task::{Context, Poll}; + +use serde_json::Value; +use tokio::io::AsyncWrite; + +/// Recording writer that captures every byte written to it and tracks +/// whether `poll_shutdown` was observed. +/// +/// Clones share the same underlying buffers, so a test can clone the +/// writer before moving it into the code under test and query the clone +/// afterwards. +#[derive(Clone, Default)] +pub(super) struct RecordingWriter { + bytes: Arc>>, + shutdown_called: Arc>, +} + +impl RecordingWriter { + /// Create a fresh recorder with empty buffers. + pub(super) fn new() -> Self { + Self::default() + } + + /// Return a copy of the bytes captured so far. + pub(super) fn snapshot(&self) -> Vec { + self.bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .clone() + } + + /// Return `true` when `poll_shutdown` has been called on any clone. + pub(super) fn shutdown_observed(&self) -> bool { + *self + .shutdown_called + .lock() + .unwrap_or_else(PoisonError::into_inner) + } +} + +impl AsyncWrite for RecordingWriter { + fn poll_write( + self: Pin<&mut Self>, + _cx: &mut Context<'_>, + buf: &[u8], + ) -> Poll> { + self.bytes + .lock() + .unwrap_or_else(PoisonError::into_inner) + .extend_from_slice(buf); + Poll::Ready(Ok(buf.len())) + } + + fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + *self + .shutdown_called + .lock() + .unwrap_or_else(PoisonError::into_inner) = true; + Poll::Ready(Ok(())) + } +} + +/// Build a serialised JSON-RPC 2.0 frame terminated by `line_ending`. +/// +/// Pass `id = Some(…)` for requests and `id = None` for notifications. +pub(super) fn jsonrpc_frame( + id: Option<&Value>, + method: &str, + line_ending: &[u8], +) -> Result, serde_json::Error> { + let payload = id.map_or_else( + || { + serde_json::json!({ + "jsonrpc": "2.0", + "method": method, + "params": {}, + }) + }, + |request_id| { + serde_json::json!({ + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": {}, + }) + }, + ); + let mut bytes = serde_json::to_vec(&payload)?; + bytes.extend_from_slice(line_ending); + Ok(bytes) +} diff --git a/src/engine/connection/exec/mod.rs b/src/engine/connection/exec/mod.rs index 4a8dce15..a7f0a602 100644 --- a/src/engine/connection/exec/mod.rs +++ b/src/engine/connection/exec/mod.rs @@ -7,6 +7,8 @@ mod acp_frame; mod acp_helpers; mod acp_policy; mod acp_runtime; +#[cfg(test)] +mod acp_test_support; mod attached; mod helpers; mod host_io; diff --git a/src/engine/connection/exec/protocol_acp_forwarding_tests.rs b/src/engine/connection/exec/protocol_acp_forwarding_tests.rs index e8bf20e8..cd1c9ca0 100644 --- a/src/engine/connection/exec/protocol_acp_forwarding_tests.rs +++ b/src/engine/connection/exec/protocol_acp_forwarding_tests.rs @@ -62,7 +62,7 @@ fn forwarding_does_not_wait_indefinitely_for_oversized_initial_frame() { let mut buffered_stdin = tokio::io::BufReader::with_capacity(STDIN_BUFFER_CAPACITY, host_reader); let recording_input = RecordingInputWriter::new(); - let forwarded_bytes = recording_input.bytes.clone(); + let recorder = recording_input.clone(); let mut container_input: Pin> = Box::pin(recording_input); runtime @@ -77,10 +77,7 @@ fn forwarding_does_not_wait_indefinitely_for_oversized_initial_frame() { .expect("initial forwarding should succeed"); assert_eq!( - forwarded_bytes - .lock() - .expect("writer mutex should not poison") - .len(), + recorder.snapshot().len(), MAX_FIRST_FRAME_BYTES, "only the bounded first-frame buffer should be held before streaming resumes" ); diff --git a/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs b/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs index 6ef99780..f7276358 100644 --- a/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs +++ b/src/engine/connection/exec/protocol_acp_policy_integration_tests.rs @@ -1,55 +1,17 @@ //! Integration-style ACP policy-selection tests for protocol sessions. use std::io; -use std::pin::Pin; -use std::sync::{Arc, Mutex, PoisonError}; -use std::task::{Context, Poll}; use bollard::container::LogOutput; use futures_util::stream; -use tokio::io::AsyncWrite; +use rstest::rstest; use super::super::{ProtocolProxyIo, ProtocolSessionOptions, run_protocol_session_with_io_async}; +use crate::engine::connection::exec::acp_test_support::{RecordingWriter, jsonrpc_frame}; use crate::engine::connection::exec::session::CapabilityPolicy; use crate::engine::connection::exec::{ExecMode, ExecRequest}; use crate::error::PodbotError; -#[derive(Clone, Default)] -struct RecordingOutput { - bytes: Arc>>, -} - -impl RecordingOutput { - fn snapshot(&self) -> Vec { - self.bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .clone() - } -} - -impl AsyncWrite for RecordingOutput { - fn poll_write( - self: Pin<&mut Self>, - _cx: &mut Context<'_>, - buf: &[u8], - ) -> Poll> { - self.bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .extend_from_slice(buf); - Poll::Ready(Ok(buf.len())) - } - - fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - Poll::Ready(Ok(())) - } - - fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - Poll::Ready(Ok(())) - } -} - fn protocol_request() -> Result { ExecRequest::new( "policy-selection-sandbox", @@ -59,26 +21,19 @@ fn protocol_request() -> Result { } fn blocked_request_frame(id: i64) -> Result, serde_json::Error> { - let mut bytes = serde_json::to_vec(&serde_json::json!({ - "jsonrpc": "2.0", - "id": id, - "method": "terminal/create", - "params": {}, - }))?; - bytes.push(b'\n'); - Ok(bytes) + jsonrpc_frame(Some(&serde_json::json!(id)), "terminal/create", b"\n") } fn drive_policy_session(policy: CapabilityPolicy) -> io::Result<(Vec, Vec, Vec)> { let runtime = tokio::runtime::Runtime::new()?; let initialize = super::initialize_frame("\n").map_err(io::Error::other)?; let host_stdin = runtime.block_on(super::build_host_stdin(&initialize))?; - let host_stdout = RecordingOutput::default(); + let host_stdout = RecordingWriter::default(); let host_stdout_handle = host_stdout.clone(); - let host_stderr = RecordingOutput::default(); + let host_stderr = RecordingWriter::default(); let host_stderr_handle = host_stderr.clone(); - let container_input = super::RecordingInputWriter::new(); - let container_stdin = container_input.bytes.clone(); + let container_input = RecordingWriter::new(); + let container_recorder = container_input.clone(); let blocked = blocked_request_frame(7).map_err(io::Error::other)?; let output = stream::iter([Ok(LogOutput::StdOut { message: blocked.into(), @@ -96,12 +51,8 @@ fn drive_policy_session(policy: CapabilityPolicy) -> io::Result<(Vec, Vec Result { } fn blocked_terminal_create_frame() -> Result, serde_json::Error> { - let mut bytes = serde_json::to_vec(&serde_json::json!({ - "jsonrpc": "2.0", - "id": 7, - "method": "terminal/create", - "params": {}, - }))?; - bytes.push(b'\n'); - Ok(bytes) + jsonrpc_frame(Some(&serde_json::json!(7)), "terminal/create", b"\n") } fn run_policy_output_frame( @@ -43,10 +36,10 @@ fn run_policy_output_frame( let request = protocol_request().map_err(io::Error::other)?; let host_stdin = runtime.block_on(build_host_stdin(&[]))?; let host_stdout = RecordingInputWriter::new(); - let host_stdout_bytes = host_stdout.bytes.clone(); + let host_stdout_recorder = host_stdout.clone(); let host_stderr = RecordingInputWriter::new(); let container_input = RecordingInputWriter::new(); - let container_stdin_bytes = container_input.bytes.clone(); + let container_stdin_recorder = container_input.clone(); let output = stream::iter([Ok(LogOutput::StdOut { message: frame.to_vec().into(), })]); @@ -62,15 +55,10 @@ fn run_policy_output_frame( )) .map_err(io::Error::other)?; - let recorded_stdout = host_stdout_bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .clone(); - let recorded_container_stdin = container_stdin_bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .clone(); - Ok((recorded_stdout, recorded_container_stdin)) + Ok(( + host_stdout_recorder.snapshot(), + container_stdin_recorder.snapshot(), + )) } /// Pure query: returns `true` when `bytes` parses as the synthesized denial diff --git a/src/engine/connection/exec/protocol_acp_tests.rs b/src/engine/connection/exec/protocol_acp_tests.rs index 3989fe84..7a8d844d 100644 --- a/src/engine/connection/exec/protocol_acp_tests.rs +++ b/src/engine/connection/exec/protocol_acp_tests.rs @@ -5,57 +5,15 @@ //! individual test groups to sibling submodules. use std::io; -use std::pin::Pin; -use std::sync::{Arc, Mutex, PoisonError}; -use std::task::{Context, Poll}; -use tokio::io::{AsyncWrite, AsyncWriteExt, DuplexStream}; +use tokio::io::{AsyncWriteExt, DuplexStream}; use super::*; use crate::engine::connection::exec::acp_helpers::{ ACP_FILE_SYSTEM_CAPABILITY, ACP_TERMINAL_CAPABILITY, MAX_FIRST_FRAME_BYTES, forward_initial_acp_frame_async, mask_acp_initialize_frame, split_frame_line_ending, }; - -struct RecordingInputWriter { - bytes: Arc>>, - shutdown_called: Arc>, -} - -impl RecordingInputWriter { - fn new() -> Self { - Self { - bytes: Arc::new(Mutex::new(Vec::new())), - shutdown_called: Arc::new(Mutex::new(false)), - } - } -} - -impl AsyncWrite for RecordingInputWriter { - fn poll_write( - self: Pin<&mut Self>, - _cx: &mut Context<'_>, - buf: &[u8], - ) -> Poll> { - self.bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .extend_from_slice(buf); - Poll::Ready(Ok(buf.len())) - } - - fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - Poll::Ready(Ok(())) - } - - fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { - *self - .shutdown_called - .lock() - .unwrap_or_else(PoisonError::into_inner) = true; - Poll::Ready(Ok(())) - } -} +use crate::engine::connection::exec::acp_test_support::RecordingWriter as RecordingInputWriter; fn initialize_frame_with_capabilities( capabilities: &serde_json::Value, @@ -79,15 +37,21 @@ fn initialize_frame_with_capabilities( Ok(frame) } +/// Builds the blocked `fs`/`terminal` capability object, optionally with an +/// unrelated `_meta` entry that masking must preserve. +fn blocked_capabilities(include_meta: bool) -> serde_json::Value { + let mut capabilities = serde_json::json!({ + "fs": { "readTextFile": true, "writeTextFile": true }, + "terminal": true + }); + if include_meta && let Some(object) = capabilities.as_object_mut() { + object.insert(String::from("_meta"), serde_json::json!({ "custom": true })); + } + capabilities +} + fn initialize_frame(line_ending: &str) -> Result, serde_json::Error> { - initialize_frame_with_capabilities( - &serde_json::json!({ - "fs": { "readTextFile": true, "writeTextFile": true }, - "terminal": true, - "_meta": { "custom": true } - }), - line_ending, - ) + initialize_frame_with_capabilities(&blocked_capabilities(true), line_ending) } /// Builds a serialised ACP `initialize` frame whose `clientCapabilities` @@ -115,13 +79,7 @@ pub(super) fn initialize_without_blocked_capabilities() -> Result, serde fn initialize_with_only_blocked_capabilities( line_ending: &str, ) -> Result, serde_json::Error> { - initialize_frame_with_capabilities( - &serde_json::json!({ - "fs": { "readTextFile": true, "writeTextFile": true }, - "terminal": true - }), - line_ending, - ) + initialize_frame_with_capabilities(&blocked_capabilities(false), line_ending) } fn session_new_bytes() -> Vec { @@ -194,8 +152,7 @@ fn run_forwarding_with_rewrite( let runtime = tokio::runtime::Runtime::new()?; let host_stdin = runtime.block_on(build_host_stdin(host_stdin_bytes))?; let container_input = RecordingInputWriter::new(); - let forwarded_bytes = container_input.bytes.clone(); - let shutdown_called = container_input.shutdown_called.clone(); + let recorder = container_input.clone(); runtime.block_on(forward_host_stdin_to_exec_async( host_stdin, @@ -203,14 +160,7 @@ fn run_forwarding_with_rewrite( rewrite_acp_initialize, ))?; - let forwarded = forwarded_bytes - .lock() - .unwrap_or_else(PoisonError::into_inner) - .clone(); - let shutdown = *shutdown_called - .lock() - .unwrap_or_else(PoisonError::into_inner); - Ok((forwarded, shutdown)) + Ok((recorder.snapshot(), recorder.shutdown_observed())) } /// Constructs a host-stdin byte sequence containing a masked `initialize` diff --git a/src/engine/connection/exec/runtime_helpers.rs b/src/engine/connection/exec/runtime_helpers.rs index 39decb20..2e9d0544 100644 --- a/src/engine/connection/exec/runtime_helpers.rs +++ b/src/engine/connection/exec/runtime_helpers.rs @@ -110,8 +110,8 @@ mod tests { fn block_on_runtime_maps_outcomes_outside_tokio( current_thread_runtime: io::Result, #[case] outcome: OutsideTokioOutcome, - ) { - let rt = current_thread_runtime.expect("runtime should be created"); + ) -> io::Result<()> { + let rt = current_thread_runtime?; let handle = rt.handle().clone(); match outcome { @@ -137,6 +137,7 @@ mod tests { ); } } + Ok(()) } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] diff --git a/src/engine/connection/upload_credentials/archive.rs b/src/engine/connection/upload_credentials/archive.rs index d4bba4ed..d621f0e4 100644 --- a/src/engine/connection/upload_credentials/archive.rs +++ b/src/engine/connection/upload_credentials/archive.rs @@ -149,8 +149,11 @@ fn append_file_entry( entry: &SortedEntry, path: String, ) -> io::Result<()> { - let metadata = parent_dir.metadata(&entry.file_name)?; let mut file = parent_dir.open(&entry.file_name)?; + // Read the metadata from the open handle so the header size cannot drift + // from the bytes actually streamed if the file changes on disk between + // the two calls. + let metadata = file.metadata()?; let mut header = new_entry_header( EntryType::Regular, metadata.len(), diff --git a/src/github/client_tests.rs b/src/github/client_tests.rs index 74d858e1..4aeeff2a 100644 --- a/src/github/client_tests.rs +++ b/src/github/client_tests.rs @@ -15,48 +15,47 @@ use super::{ec_pem, temp_key_dir, valid_rsa_pem}; fn build_app_client_with_valid_key_succeeds( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); // Octocrab's build() spawns a Tower buffer task requiring a Tokio runtime. - let rt = tokio::runtime::Runtime::new().expect("should create tokio runtime"); + let rt = tokio::runtime::Runtime::new()?; let _guard = rt.enter(); let result = build_app_client(12345, key); assert!(result.is_ok(), "expected Ok, got: {result:?}"); + Ok(()) } #[rstest] fn build_app_client_with_zero_app_id_succeeds( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); // Builder does not validate app_id; GitHub validates at token time. // Octocrab's build() spawns a Tower buffer task requiring a Tokio runtime. - let rt = tokio::runtime::Runtime::new().expect("should create tokio runtime"); + let rt = tokio::runtime::Runtime::new()?; let _guard = rt.enter(); let result = build_app_client(0, key); assert!( result.is_ok(), "expected Ok even with zero app_id, got: {result:?}" ); + Ok(()) } #[rstest] fn build_app_client_without_runtime_returns_error( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); // Call without entering a Tokio runtime — should return Err, not panic. @@ -67,6 +66,7 @@ fn build_app_client_without_runtime_returns_error( message.contains("no Tokio runtime context"), "error should mention missing runtime: {message}" ); + Ok(()) } #[rstest] @@ -111,8 +111,8 @@ fn authentication_failed_error_includes_context( #[tokio::test] async fn validate_app_credentials_with_missing_key_returns_error( temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (temp_dir, _dir) = temp_key_dir.expect("should create temp key dir"); +) -> io::Result<()> { + let (temp_dir, _dir) = temp_key_dir?; let key_path = Utf8Path::from_path(temp_dir.path()) .expect("temp dir path should be UTF-8") .join("key.pem"); @@ -127,6 +127,7 @@ async fn validate_app_credentials_with_missing_key_returns_error( } other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), } + Ok(()) } #[rstest] @@ -134,9 +135,9 @@ async fn validate_app_credentials_with_missing_key_returns_error( async fn validate_app_credentials_with_invalid_pem_returns_error( ec_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("ec.pem", &ec_pem).expect("should write EC key"); +) -> io::Result<()> { + let (tmp, dir) = temp_key_dir?; + dir.write("ec.pem", &ec_pem)?; let full_path = tmp.path().join("ec.pem"); let utf8_path = Utf8Path::from_path(&full_path).expect("temp path should be UTF-8"); @@ -151,6 +152,7 @@ async fn validate_app_credentials_with_invalid_pem_returns_error( } other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), } + Ok(()) } #[rstest] diff --git a/src/github/tests.rs b/src/github/tests.rs index 8eae75d7..1bdca44c 100644 --- a/src/github/tests.rs +++ b/src/github/tests.rs @@ -52,18 +52,18 @@ fn temp_key_dir() -> io::Result<(TempDir, Utf8Dir)> { fn load_valid_rsa_key_succeeds( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("key.pem", &valid_rsa_pem) - .expect("should write key"); +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); let result = load_private_key_from_dir(&dir, "key.pem", path); assert!(result.is_ok(), "expected Ok, got: {result:?}"); + Ok(()) } #[rstest] -fn load_missing_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); +fn load_missing_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; let path = Utf8Path::new("/config/missing.pem"); let result = load_private_key_from_dir(&dir, "missing.pem", path); assert!(result.is_err(), "expected Err for missing file"); @@ -73,12 +73,13 @@ fn load_missing_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) message.contains("failed to read file"), "error should mention file read failure: {message}" ); + Ok(()) } #[rstest] -fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("empty.pem", "").expect("should write empty file"); +fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("empty.pem", "")?; let path = Utf8Path::new("/config/empty.pem"); let result = load_private_key_from_dir(&dir, "empty.pem", path); assert!(result.is_err(), "expected Err for empty file"); @@ -87,13 +88,13 @@ fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { message.contains("empty"), "error should mention empty file: {message}" ); + Ok(()) } #[rstest] -fn load_invalid_pem_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("garbage.pem", "this is not a PEM file at all") - .expect("should write garbage file"); +fn load_invalid_pem_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("garbage.pem", "this is not a PEM file at all")?; let path = Utf8Path::new("/config/garbage.pem"); let result = load_private_key_from_dir(&dir, "garbage.pem", path); assert!(result.is_err(), "expected Err for invalid PEM"); @@ -102,12 +103,16 @@ fn load_invalid_pem_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) message.contains("invalid RSA private key"), "error should mention invalid RSA key: {message}" ); + Ok(()) } #[rstest] -fn load_ec_key_returns_clear_error(ec_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("ec.pem", &ec_pem).expect("should write EC key"); +fn load_ec_key_returns_clear_error( + ec_pem: String, + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("ec.pem", &ec_pem)?; let path = Utf8Path::new("/config/ec.pem"); let result = load_private_key_from_dir(&dir, "ec.pem", path); assert!(result.is_err(), "expected Err for ECDSA key"); @@ -120,16 +125,16 @@ fn load_ec_key_returns_clear_error(ec_pem: String, temp_key_dir: io::Result<(Tem message.contains("RSA"), "error should mention RSA requirement: {message}" ); + Ok(()) } #[rstest] fn load_ed25519_key_returns_clear_error( ed25519_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("ed25519.pem", &ed25519_pem) - .expect("should write Ed25519 key"); +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write("ed25519.pem", &ed25519_pem)?; let path = Utf8Path::new("/config/ed25519.pem"); let result = load_private_key_from_dir(&dir, "ed25519.pem", path); assert!(result.is_err(), "expected Err for Ed25519 key"); @@ -138,11 +143,12 @@ fn load_ed25519_key_returns_clear_error( message.contains("invalid RSA private key"), "Ed25519 PKCS#8 should fail RSA parse: {message}" ); + Ok(()) } #[rstest] -fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); +fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; let display = Utf8Path::new("/home/user/.config/podbot/app.pem"); let result = load_private_key_from_dir(&dir, "nonexistent.pem", display); match result { @@ -155,20 +161,21 @@ fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) { } other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), } + Ok(()) } #[rstest] fn load_private_key_resolves_full_path( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) { - let (tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write("github-app.pem", &valid_rsa_pem) - .expect("should write key"); +) -> io::Result<()> { + let (tmp, dir) = temp_key_dir?; + dir.write("github-app.pem", &valid_rsa_pem)?; let full_path = tmp.path().join("github-app.pem"); let utf8_path = Utf8Path::from_path(&full_path).expect("temp path should be UTF-8"); let result = load_private_key(utf8_path); assert!(result.is_ok(), "expected Ok, got: {result:?}"); + Ok(()) } #[rstest] @@ -237,10 +244,9 @@ fn load_invalid_key_types_return_clear_error( #[case] file_name: &str, #[case] pem_content: &str, #[case] expected_keyword: &str, -) { - let (_tmp, dir) = temp_key_dir.expect("should create temp key dir"); - dir.write(file_name, pem_content) - .expect("should write key file"); +) -> io::Result<()> { + let (_tmp, dir) = temp_key_dir?; + dir.write(file_name, pem_content)?; let display = format!("/config/{file_name}"); let path = Utf8Path::new(&display); let result = load_private_key_from_dir(&dir, file_name, path); @@ -250,4 +256,5 @@ fn load_invalid_key_types_return_clear_error( message.contains(expected_keyword), "error for {file_name} should mention '{expected_keyword}': {message}" ); + Ok(()) } diff --git a/tests/make_audit_target.rs b/tests/make_audit_target.rs index 7e6407e8..9315f618 100644 --- a/tests/make_audit_target.rs +++ b/tests/make_audit_target.rs @@ -24,6 +24,9 @@ use tempfile::TempDir; type TestResult = Result>; +/// Workspace-relative location of the generated fake `cargo` script. +const FAKE_CARGO_RELATIVE_PATH: &str = "bin/cargo"; + /// Open a capability handle on the temporary workspace root. fn open_workspace_dir(workspace: &Path) -> TestResult { Ok(Dir::open_ambient_dir( @@ -43,7 +46,7 @@ fn write_file(workspace_dir: &Dir, relative_path: &str, contents: &str) -> TestR } /// Writes the fake `cargo` script at `bin/cargo` inside the workspace; the -/// caller derives its absolute path with `workspace.join("bin/cargo")`. +/// caller derives its absolute path with `workspace.join(FAKE_CARGO_RELATIVE_PATH)`. fn write_fake_cargo( workspace_dir: &Dir, log_path: &Path, @@ -52,7 +55,7 @@ fn write_fake_cargo( ) -> TestResult { write_file( workspace_dir, - "bin/cargo", + FAKE_CARGO_RELATIVE_PATH, &format!( "#!/usr/bin/env sh\nif [ \"$1\" = metadata ]; then\nprintf '%s\\n' \"$PODBOT_FAKE_CARGO_METADATA\"\nexit {metadata_status}\nfi\nprintf '%s|%s\\n' \"$PWD\" \"$*\" >> '{}'\nexit {exit_status}\n", log_path.display() @@ -63,7 +66,7 @@ fn write_fake_cargo( use cap_std::fs::PermissionsExt; let permissions = cap_std::fs::Permissions::from_mode(0o755); - workspace_dir.set_permissions("bin/cargo", permissions)?; + workspace_dir.set_permissions(FAKE_CARGO_RELATIVE_PATH, permissions)?; } Ok(()) } @@ -134,7 +137,7 @@ fn rust_audit_invokes_cargo_audit_once_at_workspace_root() { let log_path = workspace.join("cargo-audit.log"); write_fake_cargo(&workspace_dir, &log_path, 0, 0).expect("fake cargo should be built"); - let fake_cargo = workspace.join("bin/cargo"); + let fake_cargo = workspace.join(FAKE_CARGO_RELATIVE_PATH); let metadata = cargo_metadata_for(workspace, &[&root_manifest, &member_manifest]); let output = @@ -190,7 +193,7 @@ fn rust_audit_propagates_failure( metadata_exit_status, ) .expect("fake cargo should be built"); - let fake_cargo = workspace.join("bin/cargo"); + let fake_cargo = workspace.join(FAKE_CARGO_RELATIVE_PATH); let metadata = cargo_metadata_for(workspace, &[&workspace.join("Cargo.toml")]); let output = From bc42f124c1a6f5f58b2f5c94129b672d1dab8238 Mon Sep 17 00:00:00 2001 From: leynos Date: Thu, 9 Jul 2026 19:25:03 +0200 Subject: [PATCH 4/5] Merge the two similar single-frame routing tests CodeScene flagged the permitted-frame and blocked-notification tests as structurally duplicated. They are now a single parameterized rstest case that drives one frame through the adapter and asserts the expected host-stdout routing per frame kind. --- .../connection/exec/acp_runtime_tests.rs | 52 +++++++++---------- 1 file changed, 26 insertions(+), 26 deletions(-) diff --git a/src/engine/connection/exec/acp_runtime_tests.rs b/src/engine/connection/exec/acp_runtime_tests.rs index b433bd3a..144bb254 100644 --- a/src/engine/connection/exec/acp_runtime_tests.rs +++ b/src/engine/connection/exec/acp_runtime_tests.rs @@ -70,22 +70,44 @@ async fn drain_channel(mut rx: mpsc::Receiver) -> Vec { received } +/// The frame variants exercised by the single-frame routing test. +#[derive(Debug, Clone, Copy)] +enum SingleFrame { + Permitted, + BlockedNotification, +} + +#[rstest] +#[case::permitted_reaches_host_stdout(SingleFrame::Permitted)] +#[case::blocked_notification_drops_silently(SingleFrame::BlockedNotification)] #[tokio::test] -async fn permitted_frame_writes_to_host_stdout_only() { +async fn single_frame_routing_matches_policy(#[case] kind: SingleFrame) { let (mut adapter, rx) = build_adapter(); let host_stdout = RecordingWriter::default(); let recorder = host_stdout.clone(); let mut writer: Pin> = Box::pin(host_stdout); - let frame = permitted_frame().expect("frame should serialize"); + let frame = match kind { + SingleFrame::Permitted => permitted_frame(), + SingleFrame::BlockedNotification => blocked_notification_frame(), + } + .expect("frame should serialize"); adapter .handle_chunk(&frame, &mut writer) .await - .expect("permitted frame writes succeed"); + .expect("frame handles cleanly"); adapter.finish(); drop(adapter); - assert_eq!(recorder.snapshot(), frame); + let expected_stdout: &[u8] = match kind { + SingleFrame::Permitted => &frame, + SingleFrame::BlockedNotification => b"", + }; + assert_eq!( + recorder.snapshot(), + expected_stdout, + "{kind:?} should route to host stdout accordingly", + ); let received = drain_channel(rx).await; assert!(received.is_empty(), "no commands should reach the sink"); } @@ -127,28 +149,6 @@ async fn blocked_request_skips_host_stdout_and_queues_synthesized_response() { ); } -#[tokio::test] -async fn blocked_notification_drops_silently_without_sink_command() { - let (mut adapter, rx) = build_adapter(); - let host_stdout = RecordingWriter::default(); - let recorder = host_stdout.clone(); - let mut writer: Pin> = Box::pin(host_stdout); - let frame = blocked_notification_frame().expect("frame should serialize"); - - adapter - .handle_chunk(&frame, &mut writer) - .await - .expect("blocked notification handles cleanly"); - drop(adapter); - - assert!(recorder.snapshot().is_empty()); - let received = drain_channel(rx).await; - assert!( - received.is_empty(), - "notifications must not generate a response" - ); -} - #[tokio::test] async fn permitted_frame_after_blocked_frame_still_reaches_host_stdout() { let (mut adapter, rx) = build_adapter(); From 2ffc497cead5dca440c3685ca91af6548435e53e Mon Sep 17 00:00:00 2001 From: leynos Date: Thu, 9 Jul 2026 19:28:34 +0200 Subject: [PATCH 5/5] Use eyre::ensure in Result-returning tests The fixture-consumption rework made several tests return Result so fixtures propagate with `?`, but the crate denies clippy::panic_in_result_fn, which rejects assert!/assert_eq!/panic! in any Result-returning function, including tests. Swap those assertions for eyre::ensure! and the panicking match arms for eyre::bail!, and return eyre::Result so both io::Error and PodbotError propagate. --- src/engine/connection/exec/runtime_helpers.rs | 11 ++-- src/github/client_tests.rs | 30 ++++----- src/github/tests.rs | 63 ++++++++++--------- 3 files changed, 56 insertions(+), 48 deletions(-) diff --git a/src/engine/connection/exec/runtime_helpers.rs b/src/engine/connection/exec/runtime_helpers.rs index 2e9d0544..5915b93a 100644 --- a/src/engine/connection/exec/runtime_helpers.rs +++ b/src/engine/connection/exec/runtime_helpers.rs @@ -58,6 +58,7 @@ mod tests { use std::io; use super::*; + use eyre::{bail, ensure}; use rstest::{fixture, rstest}; enum OutsideTokioOutcome { @@ -110,7 +111,7 @@ mod tests { fn block_on_runtime_maps_outcomes_outside_tokio( current_thread_runtime: io::Result, #[case] outcome: OutsideTokioOutcome, - ) -> io::Result<()> { + ) -> eyre::Result<()> { let rt = current_thread_runtime?; let handle = rt.handle().clone(); @@ -119,14 +120,16 @@ mod tests { let result: Result = block_on_runtime(&handle, async { Ok(42_u32) }); - assert_eq!(result.expect("future should resolve to Ok(42)"), 42); + ensure!(result? == 42, "future should resolve to Ok(42)"); } OutsideTokioOutcome::Err => { let result: Result<(), crate::error::PodbotError> = block_on_runtime(&handle, async { Err(exec_failed("c", "injected error")) }); - let err = result.expect_err("future should resolve to Err"); - assert!( + let Err(err) = result else { + bail!("future should resolve to Err"); + }; + ensure!( matches!( err, crate::error::PodbotError::Container( diff --git a/src/github/client_tests.rs b/src/github/client_tests.rs index 4aeeff2a..edd3b4be 100644 --- a/src/github/client_tests.rs +++ b/src/github/client_tests.rs @@ -3,6 +3,8 @@ use std::io; +use eyre::{bail, ensure}; + use cap_std::fs_utf8::Dir as Utf8Dir; use rstest::rstest; use tempfile::TempDir; @@ -15,7 +17,7 @@ use super::{ec_pem, temp_key_dir, valid_rsa_pem}; fn build_app_client_with_valid_key_succeeds( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); @@ -24,7 +26,7 @@ fn build_app_client_with_valid_key_succeeds( let rt = tokio::runtime::Runtime::new()?; let _guard = rt.enter(); let result = build_app_client(12345, key); - assert!(result.is_ok(), "expected Ok, got: {result:?}"); + ensure!(result.is_ok(), "expected Ok, got: {result:?}"); Ok(()) } @@ -32,7 +34,7 @@ fn build_app_client_with_valid_key_succeeds( fn build_app_client_with_zero_app_id_succeeds( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); @@ -42,7 +44,7 @@ fn build_app_client_with_zero_app_id_succeeds( let rt = tokio::runtime::Runtime::new()?; let _guard = rt.enter(); let result = build_app_client(0, key); - assert!( + ensure!( result.is_ok(), "expected Ok even with zero app_id, got: {result:?}" ); @@ -53,16 +55,16 @@ fn build_app_client_with_zero_app_id_succeeds( fn build_app_client_without_runtime_returns_error( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); let key = load_private_key_from_dir(&dir, "key.pem", path).expect("should load valid key"); // Call without entering a Tokio runtime — should return Err, not panic. let result = build_app_client(42, key); - assert!(result.is_err(), "expected Err without runtime, got Ok"); + ensure!(result.is_err(), "expected Err without runtime, got Ok"); let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( + ensure!( message.contains("no Tokio runtime context"), "error should mention missing runtime: {message}" ); @@ -111,21 +113,20 @@ fn authentication_failed_error_includes_context( #[tokio::test] async fn validate_app_credentials_with_missing_key_returns_error( temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (temp_dir, _dir) = temp_key_dir?; let key_path = Utf8Path::from_path(temp_dir.path()) .expect("temp dir path should be UTF-8") .join("key.pem"); let result = validate_app_credentials(12345, &key_path).await; - assert!(result.is_err(), "expected Err for missing key file"); match result { Err(GitHubError::PrivateKeyLoadFailed { ref path, .. }) => { - assert!( + ensure!( path.to_string_lossy().contains("key.pem"), "error path should reference the missing file" ); } - other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), + other => bail!("expected PrivateKeyLoadFailed, got: {other:?}"), } Ok(()) } @@ -135,22 +136,21 @@ async fn validate_app_credentials_with_missing_key_returns_error( async fn validate_app_credentials_with_invalid_pem_returns_error( ec_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (tmp, dir) = temp_key_dir?; dir.write("ec.pem", &ec_pem)?; let full_path = tmp.path().join("ec.pem"); let utf8_path = Utf8Path::from_path(&full_path).expect("temp path should be UTF-8"); let result = validate_app_credentials(12345, utf8_path).await; - assert!(result.is_err(), "expected Err for ECDSA key"); match result { Err(GitHubError::PrivateKeyLoadFailed { message, .. }) => { - assert!( + ensure!( message.contains("ECDSA"), "error should mention ECDSA: {message}" ); } - other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), + other => bail!("expected PrivateKeyLoadFailed, got: {other:?}"), } Ok(()) } diff --git a/src/github/tests.rs b/src/github/tests.rs index 1bdca44c..32e53632 100644 --- a/src/github/tests.rs +++ b/src/github/tests.rs @@ -8,6 +8,8 @@ use std::io; +use eyre::{bail, ensure}; + use super::*; use cap_std::fs_utf8::Dir as Utf8Dir; use rstest::{fixture, rstest}; @@ -52,24 +54,26 @@ fn temp_key_dir() -> io::Result<(TempDir, Utf8Dir)> { fn load_valid_rsa_key_succeeds( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("key.pem", &valid_rsa_pem)?; let path = Utf8Path::new("/display/key.pem"); let result = load_private_key_from_dir(&dir, "key.pem", path); - assert!(result.is_ok(), "expected Ok, got: {result:?}"); + ensure!(result.is_ok(), "expected Ok, got: {result:?}"); Ok(()) } #[rstest] -fn load_missing_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { +fn load_missing_file_returns_error( + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; let path = Utf8Path::new("/config/missing.pem"); let result = load_private_key_from_dir(&dir, "missing.pem", path); - assert!(result.is_err(), "expected Err for missing file"); + ensure!(result.is_err(), "expected Err for missing file"); let error = result.as_ref().err(); let message = format!("{error:?}"); - assert!( + ensure!( message.contains("failed to read file"), "error should mention file read failure: {message}" ); @@ -77,14 +81,14 @@ fn load_missing_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) } #[rstest] -fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { +fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("empty.pem", "")?; let path = Utf8Path::new("/config/empty.pem"); let result = load_private_key_from_dir(&dir, "empty.pem", path); - assert!(result.is_err(), "expected Err for empty file"); + ensure!(result.is_err(), "expected Err for empty file"); let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( + ensure!( message.contains("empty"), "error should mention empty file: {message}" ); @@ -92,14 +96,16 @@ fn load_empty_file_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) - } #[rstest] -fn load_invalid_pem_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { +fn load_invalid_pem_returns_error( + temp_key_dir: io::Result<(TempDir, Utf8Dir)>, +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("garbage.pem", "this is not a PEM file at all")?; let path = Utf8Path::new("/config/garbage.pem"); let result = load_private_key_from_dir(&dir, "garbage.pem", path); - assert!(result.is_err(), "expected Err for invalid PEM"); + ensure!(result.is_err(), "expected Err for invalid PEM"); let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( + ensure!( message.contains("invalid RSA private key"), "error should mention invalid RSA key: {message}" ); @@ -110,18 +116,18 @@ fn load_invalid_pem_returns_error(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) fn load_ec_key_returns_clear_error( ec_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("ec.pem", &ec_pem)?; let path = Utf8Path::new("/config/ec.pem"); let result = load_private_key_from_dir(&dir, "ec.pem", path); - assert!(result.is_err(), "expected Err for ECDSA key"); + ensure!(result.is_err(), "expected Err for ECDSA key"); let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( + ensure!( message.contains("ECDSA"), "error should mention ECDSA: {message}" ); - assert!( + ensure!( message.contains("RSA"), "error should mention RSA requirement: {message}" ); @@ -132,14 +138,14 @@ fn load_ec_key_returns_clear_error( fn load_ed25519_key_returns_clear_error( ed25519_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write("ed25519.pem", &ed25519_pem)?; let path = Utf8Path::new("/config/ed25519.pem"); let result = load_private_key_from_dir(&dir, "ed25519.pem", path); - assert!(result.is_err(), "expected Err for Ed25519 key"); + ensure!(result.is_err(), "expected Err for Ed25519 key"); let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( + ensure!( message.contains("invalid RSA private key"), "Ed25519 PKCS#8 should fail RSA parse: {message}" ); @@ -147,19 +153,18 @@ fn load_ed25519_key_returns_clear_error( } #[rstest] -fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io::Result<()> { +fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; let display = Utf8Path::new("/home/user/.config/podbot/app.pem"); let result = load_private_key_from_dir(&dir, "nonexistent.pem", display); match result { Err(GitHubError::PrivateKeyLoadFailed { ref path, .. }) => { - assert_eq!( - path.to_str(), - Some("/home/user/.config/podbot/app.pem"), - "error path should match display path" + ensure!( + path.to_str() == Some("/home/user/.config/podbot/app.pem"), + "error path should match display path: {path:?}" ); } - other => panic!("expected PrivateKeyLoadFailed, got: {other:?}"), + other => bail!("expected PrivateKeyLoadFailed, got: {other:?}"), } Ok(()) } @@ -168,13 +173,13 @@ fn error_includes_file_path(temp_key_dir: io::Result<(TempDir, Utf8Dir)>) -> io: fn load_private_key_resolves_full_path( valid_rsa_pem: String, temp_key_dir: io::Result<(TempDir, Utf8Dir)>, -) -> io::Result<()> { +) -> eyre::Result<()> { let (tmp, dir) = temp_key_dir?; dir.write("github-app.pem", &valid_rsa_pem)?; let full_path = tmp.path().join("github-app.pem"); let utf8_path = Utf8Path::from_path(&full_path).expect("temp path should be UTF-8"); let result = load_private_key(utf8_path); - assert!(result.is_ok(), "expected Ok, got: {result:?}"); + ensure!(result.is_ok(), "expected Ok, got: {result:?}"); Ok(()) } @@ -244,15 +249,15 @@ fn load_invalid_key_types_return_clear_error( #[case] file_name: &str, #[case] pem_content: &str, #[case] expected_keyword: &str, -) -> io::Result<()> { +) -> eyre::Result<()> { let (_tmp, dir) = temp_key_dir?; dir.write(file_name, pem_content)?; let display = format!("/config/{file_name}"); let path = Utf8Path::new(&display); let result = load_private_key_from_dir(&dir, file_name, path); - assert!(result.is_err(), "expected Err for {file_name}"); + ensure!(result.is_err(), "expected Err for {file_name}"); let message = result.err().map(|e| e.to_string()).unwrap_or_default(); - assert!( + ensure!( message.contains(expected_keyword), "error for {file_name} should mention '{expected_keyword}': {message}" );