diff --git a/architecture/gateway.md b/architecture/gateway.md index 769f57f6a0..df340f6cc0 100644 --- a/architecture/gateway.md +++ b/architecture/gateway.md @@ -509,6 +509,11 @@ validate its config. The effective sandbox config contains only the registered services required by that policy; supervisors invoke those services directly on the request path. +The effective sandbox config also carries the supervisor-wide HTTP response +whole-body timeout. The gateway reads this static value from +`[openshell.supervisor]`, defaults it to 120 seconds, and distributes it as +milliseconds. A zero value from an older gateway maps to the same default. + Provider credential expiry is enforced during gateway-to-sandbox credential resolution and again by the sandbox placeholder resolver. This keeps expired credentials from resolving even when a running sandbox still has retained diff --git a/architecture/sandbox.md b/architecture/sandbox.md index 055ef7e4a3..ad0ce61117 100644 --- a/architecture/sandbox.md +++ b/architecture/sandbox.md @@ -185,6 +185,16 @@ middleware registry validates implementation-owned config. The generic registry and chain runner live in `openshell-supervisor-middleware`; first-party implementations live in `openshell-supervisor-middleware-builtins`. +The same selected chain can inspect the matching final HTTP response before it +returns to the workload. Response stages select header-only, whole-body, or +streaming mode independently. The relay preserves upstream framing for a +header-only chain and owns normalized downstream framing only when body bytes +can change. Whole-body stages delay commitment and share one non-resetting, +supervisor-wide accumulation deadline. Body stages receive a final body result +and then one trailer exchange; trailer mutations can only change or remove +existing, non-protected names. Intentional blocks return the canonical 403 +before commitment and abort delivery without injected bytes after commitment. + The supervisor installs policy and middleware registry changes as one runtime generation and preserves the last-known-good generation if preparation fails. Policy-only updates reuse the connected registry, so an external middleware diff --git a/crates/openshell-core/src/grpc_client.rs b/crates/openshell-core/src/grpc_client.rs index 54f0db6902..4da38ae1ee 100644 --- a/crates/openshell-core/src/grpc_client.rs +++ b/crates/openshell-core/src/grpc_client.rs @@ -943,6 +943,8 @@ pub struct SettingsPollResult { pub policy_validation_failure_mode: crate::PolicyValidationFailureMode, /// Whether the gateway can mint authenticated extension credentials. pub extension_authentication_enabled: bool, + /// Supervisor-wide response whole-body accumulation timeout. + pub http_response_whole_body_timeout_ms: u64, } fn settings_poll_result(inner: crate::proto::GetSandboxConfigResponse) -> SettingsPollResult { @@ -963,6 +965,11 @@ fn settings_poll_result(inner: crate::proto::GetSandboxConfigResponse) -> Settin .parse() .unwrap_or_default(), extension_authentication_enabled: inner.extension_authentication_enabled, + http_response_whole_body_timeout_ms: if inner.http_response_whole_body_timeout_ms == 0 { + crate::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS + } else { + inner.http_response_whole_body_timeout_ms + }, } } @@ -984,6 +991,15 @@ mod settings_poll_tests { ); } + #[test] + fn zero_whole_body_timeout_uses_compatibility_default() { + let result = settings_poll_result(GetSandboxConfigResponse::default()); + assert_eq!( + result.http_response_whole_body_timeout_ms, + crate::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS + ); + } + #[test] fn unknown_validation_failure_mode_fails_closed() { let result = settings_poll_result(GetSandboxConfigResponse { diff --git a/crates/openshell-core/src/lib.rs b/crates/openshell-core/src/lib.rs index 7acb72dd6f..f38ce8d5b2 100644 --- a/crates/openshell-core/src/lib.rs +++ b/crates/openshell-core/src/lib.rs @@ -74,6 +74,9 @@ pub const VERSION: &str = match option_env!("OPENSHELL_GIT_VERSION") { None => env!("CARGO_PKG_VERSION"), }; +/// Default wall-clock bound for HTTP response whole-body middleware buffering. +pub const DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS: u64 = 120_000; + #[cfg(test)] #[path = "../build_version.rs"] mod build_version; diff --git a/crates/openshell-sandbox/src/lib.rs b/crates/openshell-sandbox/src/lib.rs index 7afae200b5..f92b5ba5fd 100644 --- a/crates/openshell-sandbox/src/lib.rs +++ b/crates/openshell-sandbox/src/lib.rs @@ -2380,6 +2380,9 @@ async fn load_policy( openshell_core::grpc_client::fetch_settings_snapshot(endpoint, id) }) .await?; + openshell_supervisor_network::set_http_response_whole_body_timeout(Duration::from_millis( + snapshot.http_response_whole_body_timeout_ms, + )); let mut proto_policy = if let Some(p) = snapshot.policy.clone() { p @@ -3781,6 +3784,9 @@ async fn run_policy_poll_loop_with_client( // reconciled below instead of being recorded as already applied. match client.poll_settings(&ctx.sandbox_id).await { Ok(result) => { + openshell_supervisor_network::set_http_response_whole_body_timeout( + Duration::from_millis(result.http_response_whole_body_timeout_ms), + ); let _ = ctx.workspace_tx.send(client.workspace()); match initial_poll_disposition(&ctx.loaded_policy_origin, &result) { InitialPollDisposition::Acknowledge(candidate) => { @@ -3867,6 +3873,10 @@ async fn run_policy_poll_loop_with_client( } }; + openshell_supervisor_network::set_http_response_whole_body_timeout(Duration::from_millis( + result.http_response_whole_body_timeout_ms, + )); + // Reuse installed per-service credentials, rotating only when one is // missing or due. Rotation happens on the existing gateway channel and // updates slots in place, so it is independent of config revision and @@ -4963,6 +4973,8 @@ network_policies: workspace: String::new(), policy_validation_failure_mode: PolicyValidationFailureMode::default(), extension_authentication_enabled: false, + http_response_whole_body_timeout_ms: + openshell_core::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS, } } diff --git a/crates/openshell-server/src/config_file.rs b/crates/openshell-server/src/config_file.rs index 64ac953624..6bfe0b7ae5 100644 --- a/crates/openshell-server/src/config_file.rs +++ b/crates/openshell-server/src/config_file.rs @@ -214,12 +214,40 @@ pub struct OtlpConfig { #[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(deny_unknown_fields)] pub struct SupervisorFileSection { + /// Wall-clock limit for accumulating and processing a response through + /// whole-body middleware. Accepts a positive integer followed by `ms`, + /// `s`, or `m`. + #[serde(default)] + pub http_response_whole_body_timeout: Option, + /// Statically registered supervisor middleware services. Registration is /// operator-owned and changes require a gateway restart. #[serde(default)] pub middleware: Vec, } +impl SupervisorFileSection { + /// Resolve the configured whole-body timeout to milliseconds. + #[must_use] + pub fn http_response_whole_body_timeout_ms(&self) -> u64 { + self.http_response_whole_body_timeout + .as_deref() + .and_then(parse_positive_duration_ms) + .unwrap_or(openshell_core::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS) + } +} + +fn parse_positive_duration_ms(value: &str) -> Option { + let value = value.trim(); + let (number, multiplier) = value + .strip_suffix("ms") + .map(|number| (number, 1)) + .or_else(|| value.strip_suffix('s').map(|number| (number, 1_000))) + .or_else(|| value.strip_suffix('m').map(|number| (number, 60_000)))?; + let number = number.parse::().ok()?; + (number > 0).then_some(number.checked_mul(multiplier)?) +} + /// One `[[openshell.supervisor.middleware]]` supervisor middleware registration. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(deny_unknown_fields)] @@ -415,6 +443,18 @@ pub fn load(path: &Path) -> Result { message: "omit the field to use default encrypted gateway credential storage, or specify exactly one external credential driver", }); } + if file + .openshell + .supervisor + .http_response_whole_body_timeout + .as_deref() + .is_some_and(|value| parse_positive_duration_ms(value).is_none()) + { + return Err(ConfigFileError::InvalidValue { + field: "openshell.supervisor.http_response_whole_body_timeout", + message: "expected a positive integer duration ending in ms, s, or m", + }); + } Ok(file) } @@ -604,6 +644,39 @@ service_name = "openshell-gateway-dev" assert_eq!(otlp.service_name.as_deref(), Some("openshell-gateway-dev")); } + #[test] + fn parses_http_response_whole_body_timeout() { + let tmp = write_tmp( + r#" +[openshell.supervisor] +http_response_whole_body_timeout = "2m" +"#, + ); + let file = load(tmp.path()).expect("valid supervisor timeout parses"); + assert_eq!( + file.openshell + .supervisor + .http_response_whole_body_timeout_ms(), + 120_000 + ); + } + + #[test] + fn rejects_invalid_http_response_whole_body_timeout() { + for value in ["0s", "120", "later", "18446744073709551615m"] { + let tmp = write_tmp(&format!( + "[openshell.supervisor]\nhttp_response_whole_body_timeout = \"{value}\"\n" + )); + let error = load(tmp.path()).expect_err("invalid timeout must be rejected"); + assert!( + error + .to_string() + .contains("http_response_whole_body_timeout"), + "{error}" + ); + } + } + #[test] fn otlp_config_requires_only_endpoint() { let toml = r#" diff --git a/crates/openshell-server/src/grpc/policy.rs b/crates/openshell-server/src/grpc/policy.rs index f1f57c2f5b..1f191a7beb 100644 --- a/crates/openshell-server/src/grpc/policy.rs +++ b/crates/openshell-server/src/grpc/policy.rs @@ -2524,6 +2524,7 @@ pub(super) async fn handle_get_sandbox_config( .as_str() .to_string(), extension_authentication_enabled: state.sandbox_jwt_issuer.is_some(), + http_response_whole_body_timeout_ms: state.http_response_whole_body_timeout_ms, })) } diff --git a/crates/openshell-server/src/lib.rs b/crates/openshell-server/src/lib.rs index a8c8afdf08..219d4badd8 100644 --- a/crates/openshell-server/src/lib.rs +++ b/crates/openshell-server/src/lib.rs @@ -304,6 +304,9 @@ pub struct ServerState { /// Validated built-in and operator-registered supervisor middleware. pub middleware_registry: Arc, + /// Supervisor-wide response whole-body accumulation timeout. + pub http_response_whole_body_timeout_ms: u64, + /// OIDC JWKS cache for JWT validation. `None` when OIDC is not configured. pub oidc_cache: Option>, @@ -419,6 +422,8 @@ impl ServerState { gateway_shutting_down: AtomicBool::new(false), extension_mint_limiter: auth::extension_mint_limit::ExtensionMintLimiter::default(), middleware_registry: Arc::new(MiddlewareRegistry::default()), + http_response_whole_body_timeout_ms: + openshell_core::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS, oidc_cache, sandbox_jwt_issuer: None, sandbox_jwt_authenticator: None, @@ -658,6 +663,14 @@ pub(crate) async fn run_server( oidc_cache, credentials, ); + state.http_response_whole_body_timeout_ms = config_file.as_ref().map_or( + openshell_core::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS, + |file| { + file.openshell + .supervisor + .http_response_whole_body_timeout_ms() + }, + ); state.middleware_registry = middleware_registry; state.gateway_interceptors = gateway_interceptors; state.provider_profile_sources = provider_profile_sources; diff --git a/crates/openshell-supervisor-middleware/src/lib.rs b/crates/openshell-supervisor-middleware/src/lib.rs index 1065c18faf..57d350cd4e 100644 --- a/crates/openshell-supervisor-middleware/src/lib.rs +++ b/crates/openshell-supervisor-middleware/src/lib.rs @@ -5,8 +5,15 @@ pub mod headers; mod remote; +mod response; mod websocket; +pub use response::{ + HttpResponseDiagnostics, HttpResponseFinish, HttpResponseInvocation, + HttpResponseInvocationOutcome, HttpResponseMiddlewareFailure, HttpResponsePreflightInput, + HttpResponsePreflightOutcome, HttpResponseSession, MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES, +}; + pub use websocket::{ WebSocketCoverage, WebSocketCoverageState, WebSocketInvocation, WebSocketInvocationOutcome, WebSocketMessageAdmission, WebSocketMessageOutcome, WebSocketMessageType, @@ -626,6 +633,16 @@ impl MiddlewareDispatch { Self::Grpc(service) => service.open_websocket_session(receiver).await, } } + + async fn open_http_response_pre_return( + &self, + receiver: tokio::sync::mpsc::Receiver, + ) -> std::result::Result { + match self { + Self::InProcess(service) => service.open_http_response_pre_return(receiver).await, + Self::Grpc(service) => service.open_http_response_pre_return(receiver).await, + } + } } struct MiddlewareServiceState { @@ -831,6 +848,7 @@ fn validate_payload_limit(source: &str, binding: &MiddlewareBinding) -> Result Result Err(miette!( - "{source} advertises HTTP_RESPONSE/PRE_RETURN, which is not yet supported" - )), + ) => Ok(SupportedBinding::HttpResponsePreReturn), ( Some(SupervisorMiddlewareOperation::WebsocketMessage), Some(SupervisorMiddlewarePhase::PreCredentials), @@ -3686,7 +3702,7 @@ mod tests { } #[test] - fn manifest_rejects_http_response_pre_return_binding_until_dispatch_is_available() { + fn manifest_accepts_http_response_pre_return_binding_when_dispatch_is_available() { let registration = external_registration(4096); let manifest = MiddlewareManifest { name: "example/response".into(), @@ -3700,13 +3716,8 @@ mod tests { expected_audience: String::new(), }; - let error = validate_external_manifest(®istration, &manifest, 4096, false) - .expect_err("HTTP response pre-return binding must remain unavailable"); - assert!( - error - .to_string() - .contains("HTTP_RESPONSE/PRE_RETURN, which is not yet supported") - ); + validate_external_manifest(®istration, &manifest, 4096, false) + .expect("HTTP response pre-return binding is supported"); } #[test] diff --git a/crates/openshell-supervisor-middleware/src/remote.rs b/crates/openshell-supervisor-middleware/src/remote.rs index 9443038100..80049b69dc 100644 --- a/crates/openshell-supervisor-middleware/src/remote.rs +++ b/crates/openshell-supervisor-middleware/src/remote.rs @@ -103,6 +103,14 @@ impl GrpcMiddlewareService { ) -> std::result::Result { self.service.open_websocket_session(receiver).await } + + /// Open a remote HTTP response pre-return stream through the gRPC adapter. + pub async fn open_http_response_pre_return( + &self, + receiver: tokio::sync::mpsc::Receiver, + ) -> std::result::Result { + self.service.open_http_response_pre_return(receiver).await + } } #[derive(Clone)] diff --git a/crates/openshell-supervisor-middleware/src/response.rs b/crates/openshell-supervisor-middleware/src/response.rs new file mode 100644 index 0000000000..5baa7f93a1 --- /dev/null +++ b/crates/openshell-supervisor-middleware/src/response.rs @@ -0,0 +1,2641 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! HTTP response pre-return middleware chain execution. + +use std::collections::BTreeMap; +use std::time::Duration; + +use futures::StreamExt as _; +use prost::Message as _; +use tokio::sync::mpsc; +use tokio::time::Instant; + +use openshell_core::proto::{ + Finding, HttpHeader, HttpRequestTarget, HttpResponseBodyMode, HttpResponseBodyPassThrough, + HttpResponseBodyUnit, HttpResponseEvent, HttpResponseEventResult, HttpResponsePreflight, + HttpResponseTrailers, MiddlewareSessionEnd, MiddlewareSessionEndReason, RequestContext, + http_response_body_result, http_response_body_skip_remaining, http_response_body_transform, + http_response_body_unit, http_response_event, http_response_event_result, + http_response_preflight_result, +}; + +use super::{ + ChainEntry, ChainRunner, DescribedChainEntry, MAX_MIDDLEWARE_CHAIN_TIMEOUT, + MAX_MIDDLEWARE_CONTEXT_BYTES, MAX_MIDDLEWARE_FINDING_BYTES, MAX_MIDDLEWARE_FINDINGS_PER_STAGE, + MAX_MIDDLEWARE_HEADER_BYTES, MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES, MAX_MIDDLEWARE_HEADERS, + MAX_MIDDLEWARE_METADATA_BYTES, MAX_MIDDLEWARE_METADATA_ENTRIES, MAX_MIDDLEWARE_REASON_BYTES, + MAX_MIDDLEWARE_REASON_CODE_BYTES, MAX_MIDDLEWARE_TARGET_BYTES, MiddlewareDiagnosticPolicy, + MiddlewareSessionAdmission, MiddlewareSessionPermit, NamespacedFinding, OnError, headers, + is_stable_reason_code, middleware_denial_reason, +}; + +const STREAM_CHANNEL_CAPACITY: usize = 4; +pub const MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES: usize = 64 * 1024; + +#[derive(Debug, Clone)] +pub struct HttpResponsePreflightInput { + pub context: RequestContext, + pub target: HttpRequestTarget, + pub status_code: u16, + /// Parsed upstream Content-Length when present and valid. + pub declared_body_length: Option, + /// Sanitized, lowercased final response headers in wire order. + pub headers: Vec, + /// Lowercased names nominated by the original response's `Connection` + /// fields. Their values are not exposed to middleware. + pub connection_nominated_headers: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum HttpResponseInvocationOutcome { + Skip, + BlockDelivery, + HeadersOnly, + WholeBody, + Stream, + Trailers, + PassThrough, + Transform, + SkipRemaining, + FailOpen, + FailClosed, +} + +#[derive(Debug, Clone)] +pub struct HttpResponseInvocation { + pub config_name: String, + pub implementation: String, + pub outcome: HttpResponseInvocationOutcome, + pub sequence: Option, + pub input_size: usize, + pub output_size: Option, + pub failed: bool, + pub stage_disabled: bool, + pub reason_code: Option, + pub failure_category: Option, +} + +pub struct HttpResponsePreflightOutcome { + pub allowed: bool, + pub reason: String, + pub denial: Option, + pub headers: Vec, + pub session: Option, + pub findings: Vec, + pub metadata: BTreeMap>, + pub invocations: Vec, + pub session_capacity_exhausted: bool, +} + +#[derive(Debug)] +pub struct HttpResponseMiddlewareFailure { + pub reason: String, + pub denial: Option, + /// Exchange diagnostics collected before a consuming operation failed. + pub diagnostics: HttpResponseDiagnostics, +} + +impl std::fmt::Display for HttpResponseMiddlewareFailure { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(&self.reason) + } +} + +impl std::error::Error for HttpResponseMiddlewareFailure {} + +impl HttpResponseMiddlewareFailure { + fn with_diagnostics(mut self, diagnostics: HttpResponseDiagnostics) -> Self { + self.diagnostics = diagnostics; + self + } +} + +#[derive(Debug)] +pub struct HttpResponseFinish { + /// Units released while whole-body stages were finalized. + pub body_units: Vec>, + pub trailers: Vec, + /// True when a whole-body stage transformed or deleted body bytes. The + /// caller must strip stale representation validators before commitment. + pub strip_stale_integrity_headers: bool, + pub findings: Vec, + pub metadata: BTreeMap>, + pub invocations: Vec, +} + +#[derive(Debug, Default)] +pub struct HttpResponseDiagnostics { + pub findings: Vec, + pub metadata: BTreeMap>, + pub invocations: Vec, +} + +struct HttpResponseStageTransport { + sender: mpsc::Sender, + responses: super::HttpResponseResultStream, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum StageMode { + HeadersOnly, + WholeBody, + Stream, +} + +struct HttpResponseStage { + entry: DescribedChainEntry, + transport: Option, + mode: StageMode, + next_sequence: u64, + whole_body: Vec, +} + +impl HttpResponseStage { + fn is_active(&self) -> bool { + self.transport.is_some() + } + + fn is_body_active(&self) -> bool { + self.is_active() && self.mode != StageMode::HeadersOnly + } + + async fn end(&mut self, reason: MiddlewareSessionEndReason) { + if let Some(transport) = self.transport.take() { + let _ = tokio::time::timeout( + Duration::from_millis(10), + transport.sender.send(session_end_event(reason)), + ) + .await; + } + } +} + +pub struct HttpResponseSession { + runner: ChainRunner, + stages: Vec, + findings: Vec, + metadata: BTreeMap>, + invocations: Vec, + session_admission: Option, + body_transformed: bool, + defer_output_until_finish: bool, + deferred_output: Vec>, + connection_nominated_headers: Vec, + whole_body_deadline: Option, +} + +impl HttpResponseSession { + pub fn take_diagnostics(&mut self) -> HttpResponseDiagnostics { + HttpResponseDiagnostics { + findings: std::mem::take(&mut self.findings), + metadata: std::mem::take(&mut self.metadata), + invocations: std::mem::take(&mut self.invocations), + } + } + + #[must_use] + pub fn requires_whole_body(&self) -> bool { + self.stages.iter().any(|stage| { + stage.is_active() && stage.mode == StageMode::WholeBody && stage.next_sequence == 1 + }) + } + + /// Start the platform-owned whole-body wall-clock deadline. + pub fn start_whole_body_deadline(&mut self, timeout: Duration) { + self.whole_body_deadline = self.requires_whole_body().then(|| Instant::now() + timeout); + } + + #[must_use] + pub fn whole_body_deadline(&self) -> Option { + self.requires_whole_body() + .then_some(self.whole_body_deadline) + .flatten() + } + + /// Fail each still-buffering whole-body stage in policy order. + /// + /// Fail-open stages release their retained input through the remaining + /// chain. A fail-closed stage stops the response with a typed failure. + pub async fn expire_whole_body_deadline( + &mut self, + ) -> Result>, HttpResponseMiddlewareFailure> { + self.whole_body_deadline = None; + let mut released = std::mem::take(&mut self.deferred_output); + let chain_deadline = Instant::now() + MAX_MIDDLEWARE_CHAIN_TIMEOUT; + for index in 0..self.stages.len() { + if !self.stages[index].is_active() + || self.stages[index].mode != StageMode::WholeBody + || self.stages[index].next_sequence != 1 + { + continue; + } + let original = std::mem::take(&mut self.stages[index].whole_body); + let output = self + .handle_stage_failure(index, "whole_body_accumulation_timeout", None, original) + .await?; + if !output.is_empty() { + released.extend( + self.process_units_from(index + 1, output, chain_deadline) + .await?, + ); + } + } + self.defer_output_until_finish = false; + Ok(released) + } + + #[must_use] + pub fn stream_unit_limit(&self) -> usize { + self.stages + .iter() + .filter(|stage| stage.is_active() && stage.mode == StageMode::Stream) + .map(|stage| { + stage + .entry + .max_payload_bytes + .clamp(1, MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) + }) + .min() + .unwrap_or(MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) + } + + /// Process one normalized body unit through the active chain. + /// + /// The caller must provide no more than [`Self::stream_unit_limit`] bytes. + /// A whole-body barrier retains output until [`Self::finish`] is called. + pub async fn push_body( + &mut self, + data: Vec, + ) -> Result>, HttpResponseMiddlewareFailure> { + if data.len() > self.stream_unit_limit() { + return Err(HttpResponseMiddlewareFailure { + reason: "response_stream_unit_over_capacity".into(), + denial: None, + diagnostics: HttpResponseDiagnostics::default(), + }); + } + let _work = self + .runner + .reserve_middleware_work_admission() + .await + .map_err(|error| HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {error}"), + denial: None, + diagnostics: HttpResponseDiagnostics::default(), + })?; + let deadline = Instant::now() + MAX_MIDDLEWARE_CHAIN_TIMEOUT; + let output = self.process_units_from(0, vec![data], deadline).await?; + if !self.defer_output_until_finish { + return Ok(output); + } + if self.requires_whole_body() { + self.deferred_output.extend(output); + return Ok(Vec::new()); + } + + self.defer_output_until_finish = false; + let mut released = std::mem::take(&mut self.deferred_output); + released.extend(output); + Ok(released) + } + + /// Finalize every body stage, preserve normalized trailers, and end streams. + pub async fn finish( + mut self, + mut trailers: Vec, + ) -> Result { + let _work = match self.runner.reserve_middleware_work_admission().await { + Ok(work) => work, + Err(error) => { + return Err(HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {error}"), + denial: None, + diagnostics: self.take_diagnostics(), + }); + } + }; + let deadline = Instant::now() + MAX_MIDDLEWARE_CHAIN_TIMEOUT; + let mut released = std::mem::take(&mut self.deferred_output); + for index in 0..self.stages.len() { + let stage_output = match self.finish_stage(index, deadline).await { + Ok(output) => output, + Err(failure) => { + self.end_all(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + return Err(failure.with_diagnostics(self.take_diagnostics())); + } + }; + if !stage_output.is_empty() { + let output = match self + .process_units_from(index + 1, stage_output, deadline) + .await + { + Ok(output) => output, + Err(failure) => { + self.end_all(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + return Err(failure.with_diagnostics(self.take_diagnostics())); + } + }; + released.extend(output); + } + } + + if self.body_transformed { + strip_stale_integrity(&mut trailers); + } + let trailers = match self.process_trailers(trailers, deadline).await { + Ok(trailers) => trailers, + Err(failure) => { + self.end_all(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + return Err(failure.with_diagnostics(self.take_diagnostics())); + } + }; + self.end_all(MiddlewareSessionEndReason::Normal).await; + self.session_admission.take(); + Ok(HttpResponseFinish { + body_units: released, + trailers, + strip_stale_integrity_headers: self.body_transformed, + findings: self.findings, + metadata: self.metadata, + invocations: self.invocations, + }) + } + + pub async fn end(mut self, reason: MiddlewareSessionEndReason) { + self.end_all(reason).await; + } + + async fn process_units_from( + &mut self, + start: usize, + mut units: Vec>, + deadline: Instant, + ) -> Result>, HttpResponseMiddlewareFailure> { + for index in start..self.stages.len() { + let mut next = Vec::new(); + for unit in units { + let chunk_limit = if self.stages[index].mode == StageMode::Stream { + self.stages[index] + .entry + .max_payload_bytes + .min(MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) + } else { + unit.len().max(1) + }; + if unit.is_empty() { + next.extend(self.process_stage_unit(index, unit, deadline).await?); + } else { + for chunk in unit.chunks(chunk_limit) { + next.extend( + self.process_stage_unit(index, chunk.to_vec(), deadline) + .await?, + ); + } + } + } + units = next; + if units.is_empty() + && self.stages[index + 1..] + .iter() + .all(|stage| stage.mode != StageMode::WholeBody) + { + break; + } + } + Ok(units) + } + + async fn process_stage_unit( + &mut self, + index: usize, + data: Vec, + deadline: Instant, + ) -> Result>, HttpResponseMiddlewareFailure> { + let deadline = self.exchange_deadline(deadline); + let stage = &mut self.stages[index]; + if !stage.is_active() || stage.mode == StageMode::HeadersOnly { + return Ok(vec![data]); + } + if stage.mode == StageMode::WholeBody { + if stage.whole_body.len().saturating_add(data.len()) > stage.entry.max_payload_bytes { + let mut original = std::mem::take(&mut stage.whole_body); + original.extend_from_slice(&data); + return self + .handle_stage_failure(index, "whole_body_over_capacity", None, original) + .await; + } + stage.whole_body.extend_from_slice(&data); + return Ok(Vec::new()); + } + + let sequence = stage.next_sequence; + stage.next_sequence += 1; + let event = body_event(sequence, data.clone(), false); + let result = match exchange(stage, event, deadline).await { + Ok(result) => result, + Err(reason) => { + let reason = self.classify_timeout_reason(reason); + return self + .handle_stage_failure(index, &reason, Some(sequence), data) + .await; + } + }; + self.apply_body_result(index, result, sequence, data).await + } + + async fn finish_stage( + &mut self, + index: usize, + deadline: Instant, + ) -> Result>, HttpResponseMiddlewareFailure> { + if !self.stages[index].is_body_active() { + return Ok(Vec::new()); + } + let deadline = self.exchange_deadline(deadline); + let mode = self.stages[index].mode; + let mut output = Vec::new(); + if mode == StageMode::WholeBody { + let data = std::mem::take(&mut self.stages[index].whole_body); + let sequence = 1; + self.stages[index].next_sequence = 2; + let result = match exchange( + &mut self.stages[index], + body_event(sequence, data.clone(), true), + deadline, + ) + .await + { + Ok(result) => result, + Err(reason) => { + let reason = self.classify_timeout_reason(reason); + return self + .handle_stage_failure(index, &reason, Some(sequence), data) + .await; + } + }; + output.extend( + self.apply_body_result(index, result, sequence, data) + .await?, + ); + } + + if mode == StageMode::Stream { + let sequence = self.stages[index].next_sequence; + self.stages[index].next_sequence += 1; + let result = match exchange( + &mut self.stages[index], + body_event(sequence, Vec::new(), true), + deadline, + ) + .await + { + Ok(result) => result, + Err(reason) => { + let reason = self.classify_timeout_reason(reason); + return self + .handle_stage_failure(index, &reason, Some(sequence), Vec::new()) + .await; + } + }; + output.extend( + self.apply_body_result(index, result, sequence, Vec::new()) + .await?, + ); + } + Ok(output) + } + + async fn apply_body_result( + &mut self, + index: usize, + result: HttpResponseEventResult, + sequence: u64, + original: Vec, + ) -> Result>, HttpResponseMiddlewareFailure> { + let max_payload_bytes = self.stages[index].entry.max_payload_bytes; + let decision = match validate_body_result(result, sequence, max_payload_bytes) { + Ok(decision) => decision, + Err(reason) => { + return self + .handle_stage_failure(index, reason, Some(sequence), original) + .await; + } + }; + let input_size = original.len(); + let stage = &mut self.stages[index]; + collect_diagnostics( + stage, + decision.findings, + decision.metadata, + &mut self.findings, + &mut self.metadata, + ); + let reason_code = (!decision.reason_code.is_empty()).then_some(decision.reason_code); + match decision.action { + BodyAction::PassThrough => { + let output_size = original.len(); + self.invocations.push(body_invocation_with_reason( + stage, + HttpResponseInvocationOutcome::PassThrough, + sequence, + input_size, + output_size, + reason_code, + )); + Ok((!original.is_empty()) + .then_some(original) + .into_iter() + .collect()) + } + BodyAction::Transform(replacement) => { + self.body_transformed = true; + self.invocations.push(body_invocation_with_reason( + stage, + HttpResponseInvocationOutcome::Transform, + sequence, + input_size, + replacement.len(), + reason_code, + )); + Ok((!replacement.is_empty()) + .then_some(replacement) + .into_iter() + .collect()) + } + BodyAction::SkipRemaining(action) => { + let output = match action { + CurrentBodyAction::PassThrough => original, + CurrentBodyAction::Transform(replacement) => { + self.body_transformed = true; + replacement + } + }; + stage.mode = StageMode::HeadersOnly; + self.invocations.push(body_invocation_with_reason( + stage, + HttpResponseInvocationOutcome::SkipRemaining, + sequence, + input_size, + output.len(), + reason_code, + )); + stage.end(MiddlewareSessionEndReason::Normal).await; + self.release_admission_if_idle(); + Ok((!output.is_empty()).then_some(output).into_iter().collect()) + } + BodyAction::BlockDelivery => { + let config_name = stage.entry.entry.name.clone(); + let denial_reason = middleware_denial_reason(&config_name, reason_code.as_deref()); + self.invocations.push(body_invocation_with_reason( + stage, + HttpResponseInvocationOutcome::BlockDelivery, + sequence, + input_size, + 0, + reason_code.clone(), + )); + self.end_all(MiddlewareSessionEndReason::MiddlewareDenial) + .await; + Err(HttpResponseMiddlewareFailure { + reason: denial_reason, + denial: Some(super::MiddlewareDenial { + config_name, + reason_code, + }), + diagnostics: HttpResponseDiagnostics::default(), + }) + } + } + } + + async fn handle_stage_failure( + &mut self, + index: usize, + reason: &str, + sequence: Option, + original: Vec, + ) -> Result>, HttpResponseMiddlewareFailure> { + let stage = &mut self.stages[index]; + let fail_open = stage.entry.on_error() == OnError::FailOpen; + let outcome = if fail_open { + HttpResponseInvocationOutcome::FailOpen + } else { + HttpResponseInvocationOutcome::FailClosed + }; + self.invocations.push(HttpResponseInvocation { + config_name: stage.entry.entry.name.clone(), + implementation: stage.entry.entry.implementation.clone(), + outcome, + sequence, + input_size: original.len(), + output_size: None, + failed: true, + stage_disabled: true, + reason_code: None, + failure_category: Some(response_failure_category(reason).into()), + }); + stage + .end(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + self.release_admission_if_idle(); + if fail_open { + if original.is_empty() { + Ok(Vec::new()) + } else { + Ok(vec![original]) + } + } else { + Err(HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {reason}"), + denial: None, + diagnostics: HttpResponseDiagnostics::default(), + }) + } + } + + async fn end_all(&mut self, reason: MiddlewareSessionEndReason) { + for stage in &mut self.stages { + stage.end(reason).await; + } + } + + fn release_admission_if_idle(&mut self) { + if self.stages.iter().all(|stage| !stage.is_active()) { + self.session_admission.take(); + } + } + + fn exchange_deadline(&self, chain_deadline: Instant) -> Instant { + self.whole_body_deadline() + .map_or(chain_deadline, |deadline| deadline.min(chain_deadline)) + } + + fn classify_timeout_reason(&self, reason: String) -> String { + if reason == "middleware_timeout" + && self + .whole_body_deadline + .is_some_and(|deadline| Instant::now() >= deadline) + { + "whole_body_accumulation_timeout".into() + } else { + reason + } + } + + async fn process_trailers( + &mut self, + mut trailers: Vec, + deadline: Instant, + ) -> Result, HttpResponseMiddlewareFailure> { + for index in 0..self.stages.len() { + if !self.stages[index].is_body_active() { + continue; + } + let event = HttpResponseEvent { + event: Some(http_response_event::Event::Trailers(HttpResponseTrailers { + headers: trailers.clone(), + })), + }; + let result = match exchange(&mut self.stages[index], event, deadline).await { + Ok(result) => result, + Err(reason) => { + trailers = self + .handle_trailer_failure(index, &reason, trailers) + .await?; + continue; + } + }; + let decision = match validate_trailers_result( + result, + &trailers, + &self.stages[index].entry, + &self.connection_nominated_headers, + ) { + Ok(decision) => decision, + Err(reason) => { + trailers = self + .handle_trailer_failure(index, &reason, trailers) + .await?; + continue; + } + }; + let input_size = encoded_header_bytes(&trailers); + trailers = decision.headers; + let output_size = encoded_header_bytes(&trailers); + let reason_code = (!decision.reason_code.is_empty()).then_some(decision.reason_code); + let stage = &mut self.stages[index]; + collect_diagnostics( + stage, + decision.findings, + decision.metadata, + &mut self.findings, + &mut self.metadata, + ); + self.invocations.push(HttpResponseInvocation { + config_name: stage.entry.entry.name.clone(), + implementation: stage.entry.entry.implementation.clone(), + outcome: HttpResponseInvocationOutcome::Trailers, + sequence: None, + input_size, + output_size: Some(output_size), + failed: false, + stage_disabled: false, + reason_code, + failure_category: None, + }); + } + Ok(trailers) + } + + async fn handle_trailer_failure( + &mut self, + index: usize, + reason: &str, + original: Vec, + ) -> Result, HttpResponseMiddlewareFailure> { + let stage = &mut self.stages[index]; + let fail_open = stage.entry.on_error() == OnError::FailOpen; + self.invocations.push(HttpResponseInvocation { + config_name: stage.entry.entry.name.clone(), + implementation: stage.entry.entry.implementation.clone(), + outcome: if fail_open { + HttpResponseInvocationOutcome::FailOpen + } else { + HttpResponseInvocationOutcome::FailClosed + }, + sequence: None, + input_size: encoded_header_bytes(&original), + output_size: None, + failed: true, + stage_disabled: true, + reason_code: None, + failure_category: Some(response_failure_category(reason).into()), + }); + stage + .end(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + self.release_admission_if_idle(); + if fail_open { + Ok(original) + } else { + Err(HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {reason}"), + denial: None, + diagnostics: HttpResponseDiagnostics::default(), + }) + } + } +} + +impl ChainRunner { + pub async fn preflight_http_response( + &self, + entries: &[ChainEntry], + input: HttpResponsePreflightInput, + ) -> miette::Result { + validate_preflight_input(&input)?; + let described = self.describe_http_response_chain(entries).await?; + if described.is_empty() { + return Ok(empty_preflight_outcome(input.headers)); + } + let session_admission = match self.try_reserve_middleware_session() { + MiddlewareSessionAdmission::Admitted(admission) => admission, + MiddlewareSessionAdmission::AtCapacity => { + return Ok(response_session_capacity_exhausted( + described, + input.headers, + )); + } + }; + let _work = self.reserve_middleware_work_admission().await?; + let original_restriction = body_restriction(&input); + let mut headers = input.headers.clone(); + let mut stages = Vec::new(); + let mut findings = Vec::new(); + let mut metadata = BTreeMap::new(); + let mut invocations = Vec::new(); + + for entry in described { + let Some(service) = entry.service.as_ref() else { + if let Some(reason) = + collect_preflight_failure(&entry, "binding_not_described", &mut invocations) + { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + }; + let (sender, receiver) = mpsc::channel(STREAM_CHANNEL_CAPACITY); + let preflight = HttpResponsePreflight { + context: Some(input.context.clone()), + target: Some(input.target.clone()), + status_code: u32::from(input.status_code), + headers: headers.clone(), + middleware_name: entry.entry.implementation.clone(), + config: Some(entry.entry.config.clone()), + max_payload_bytes: entry.max_payload_bytes as u64, + permitted_body_modes: permitted_body_modes( + &input, + &entry, + original_restriction.as_deref(), + ), + }; + let timeout = entry.timeout; + let opened = tokio::time::timeout(timeout, async { + let mut responses = service + .service + .open_http_response_pre_return(receiver) + .await?; + sender + .send(HttpResponseEvent { + event: Some(http_response_event::Event::Preflight(preflight)), + }) + .await + .map_err(|_| tonic::Status::unavailable("middleware request stream closed"))?; + let response = responses.next().await.ok_or_else(|| { + tonic::Status::unavailable("middleware result stream closed") + })??; + Ok::<_, tonic::Status>((responses, response)) + }) + .await; + let (responses, response) = match opened { + Ok(Ok(opened)) => opened, + Ok(Err(error)) => { + let reason = if error.code() == tonic::Code::DeadlineExceeded { + "middleware_timeout".to_string() + } else { + service.diagnostic_policy.error_reason(&error) + }; + if let Some(reason) = + collect_preflight_failure(&entry, &reason, &mut invocations) + { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + Err(_) => { + if let Some(reason) = + collect_preflight_failure(&entry, "middleware_timeout", &mut invocations) + { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + }; + let Some(http_response_event_result::Result::PreflightResult(decision)) = + response.result + else { + if let Some(reason) = collect_preflight_failure( + &entry, + "unexpected_response_result", + &mut invocations, + ) { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + }; + if let Err(reason) = validate_diagnostics( + &decision.reason, + &decision.reason_code, + &decision.findings, + &decision.metadata, + ) { + if let Some(reason) = collect_preflight_failure(&entry, reason, &mut invocations) { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + let reason_code = + (!decision.reason_code.is_empty()).then(|| decision.reason_code.clone()); + let decision_findings = decision.findings; + let decision_metadata = decision.metadata; + match decision.action { + Some(http_response_preflight_result::Action::Skip(_)) => { + collect_preflight_diagnostics( + &entry, + decision_findings, + decision_metadata, + &mut findings, + &mut metadata, + ); + invocations.push(HttpResponseInvocation { + config_name: entry.entry.name.clone(), + implementation: entry.entry.implementation.clone(), + outcome: HttpResponseInvocationOutcome::Skip, + sequence: None, + input_size: 0, + output_size: None, + failed: false, + stage_disabled: false, + reason_code, + failure_category: None, + }); + let mut skipped = HttpResponseStage { + entry, + transport: Some(HttpResponseStageTransport { sender, responses }), + mode: StageMode::HeadersOnly, + next_sequence: 1, + whole_body: Vec::new(), + }; + skipped.end(MiddlewareSessionEndReason::StageSkipped).await; + } + Some(http_response_preflight_result::Action::Inspect(inspect)) => { + let permitted_modes = + permitted_body_modes(&input, &entry, original_restriction.as_deref()); + let mode = match validate_inspect(&entry, &inspect, &permitted_modes) { + Ok(mode) => mode, + Err(reason) => { + if let Some(reason) = + collect_preflight_failure(&entry, &reason, &mut invocations) + { + end_stages( + &mut stages, + MiddlewareSessionEndReason::MiddlewareFailure, + ) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + }; + let updated = match headers::apply( + headers::HeaderAuthority::Response, + &headers, + &input.connection_nominated_headers, + &inspect.header_mutations, + ) { + Ok(updated) => updated, + Err(error) => { + let reason = service + .diagnostic_policy + .header_mutation_error_reason(&error); + if let Some(reason) = + collect_preflight_failure(&entry, &reason, &mut invocations) + { + end_stages( + &mut stages, + MiddlewareSessionEndReason::MiddlewareFailure, + ) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + }; + headers = updated; + if mode == StageMode::Stream { + strip_stale_integrity(&mut headers); + } + collect_preflight_diagnostics( + &entry, + decision_findings, + decision_metadata, + &mut findings, + &mut metadata, + ); + invocations.push(HttpResponseInvocation { + config_name: entry.entry.name.clone(), + implementation: entry.entry.implementation.clone(), + outcome: match mode { + StageMode::HeadersOnly => HttpResponseInvocationOutcome::HeadersOnly, + StageMode::WholeBody => HttpResponseInvocationOutcome::WholeBody, + StageMode::Stream => HttpResponseInvocationOutcome::Stream, + }, + sequence: None, + input_size: 0, + output_size: None, + failed: false, + stage_disabled: false, + reason_code, + failure_category: None, + }); + let mut stage = HttpResponseStage { + entry, + transport: Some(HttpResponseStageTransport { sender, responses }), + mode, + next_sequence: 1, + whole_body: Vec::new(), + }; + if mode == StageMode::HeadersOnly { + stage.end(MiddlewareSessionEndReason::Normal).await; + } else { + stages.push(stage); + } + } + Some(http_response_preflight_result::Action::BlockDelivery(_)) => { + collect_preflight_diagnostics( + &entry, + decision_findings, + decision_metadata, + &mut findings, + &mut metadata, + ); + invocations.push(HttpResponseInvocation { + config_name: entry.entry.name.clone(), + implementation: entry.entry.implementation.clone(), + outcome: HttpResponseInvocationOutcome::BlockDelivery, + sequence: None, + input_size: 0, + output_size: None, + failed: false, + stage_disabled: false, + reason_code: reason_code.clone(), + failure_category: None, + }); + stages.push(HttpResponseStage { + entry: entry.clone(), + transport: Some(HttpResponseStageTransport { sender, responses }), + mode: StageMode::HeadersOnly, + next_sequence: 1, + whole_body: Vec::new(), + }); + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareDenial).await; + return Ok(blocked_preflight_outcome( + headers, + super::MiddlewareDenial { + config_name: entry.entry.name.clone(), + reason_code, + }, + findings, + metadata, + invocations, + )); + } + None => { + if let Some(reason) = collect_preflight_failure( + &entry, + "invalid_preflight_decision", + &mut invocations, + ) { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + } + } + } + + if stages.is_empty() { + drop(session_admission); + return Ok(HttpResponsePreflightOutcome { + allowed: true, + reason: String::new(), + denial: None, + headers, + session: None, + findings, + metadata, + invocations, + session_capacity_exhausted: false, + }); + } + let defer_output_until_finish = stages + .iter() + .any(|stage| stage.is_active() && stage.mode == StageMode::WholeBody); + Ok(HttpResponsePreflightOutcome { + allowed: true, + reason: String::new(), + denial: None, + headers, + session: Some(HttpResponseSession { + runner: self.clone(), + stages, + findings: Vec::new(), + metadata: BTreeMap::new(), + invocations: Vec::new(), + session_admission: Some(session_admission), + body_transformed: false, + defer_output_until_finish, + deferred_output: Vec::new(), + connection_nominated_headers: input.connection_nominated_headers, + whole_body_deadline: None, + }), + findings, + metadata, + invocations, + session_capacity_exhausted: false, + }) + } +} + +enum BodyAction { + PassThrough, + Transform(Vec), + BlockDelivery, + SkipRemaining(CurrentBodyAction), +} + +enum CurrentBodyAction { + PassThrough, + Transform(Vec), +} + +struct BodyDecision { + action: BodyAction, + reason_code: String, + findings: Vec, + metadata: std::collections::HashMap, +} + +struct TrailersDecision { + headers: Vec, + reason_code: String, + findings: Vec, + metadata: std::collections::HashMap, +} + +fn validate_trailers_result( + result: HttpResponseEventResult, + trailers: &[HttpHeader], + entry: &DescribedChainEntry, + connection_nominated_headers: &[String], +) -> Result { + let Some(http_response_event_result::Result::TrailersResult(result)) = result.result else { + return Err("unexpected_response_result".into()); + }; + validate_diagnostics( + &result.reason, + &result.reason_code, + &result.findings, + &result.metadata, + ) + .map_err(str::to_string)?; + if result.trailer_mutations.len() > headers::MAX_HEADER_MUTATIONS { + return Err("header_mutation_count_over_capacity".into()); + } + let encoded_mutations = result + .trailer_mutations + .iter() + .fold(0usize, |total, mutation| { + total.saturating_add(mutation.encoded_len()) + }); + if encoded_mutations > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES { + return Err("header_mutation_bytes_over_capacity".into()); + } + let headers = headers::apply( + headers::HeaderAuthority::ResponseTrailers, + trailers, + connection_nominated_headers, + &result.trailer_mutations, + ) + .map_err(|error| { + entry.service.as_ref().map_or_else( + || error.to_string(), + |service| { + service + .diagnostic_policy + .header_mutation_error_reason(&error) + }, + ) + })?; + Ok(TrailersDecision { + headers, + reason_code: result.reason_code, + findings: result.findings, + metadata: result.metadata, + }) +} + +fn encoded_header_bytes(headers: &[HttpHeader]) -> usize { + headers.iter().fold(0usize, |total, header| { + total.saturating_add(header.encoded_len()) + }) +} + +fn validate_body_result( + result: HttpResponseEventResult, + sequence: u64, + max_payload_bytes: usize, +) -> Result { + let Some(http_response_event_result::Result::BodyResult(body)) = result.result else { + return Err("unexpected_response_result"); + }; + if body.sequence != sequence { + return Err("response_body_sequence_mismatch"); + } + validate_diagnostics( + &body.reason, + &body.reason_code, + &body.findings, + &body.metadata, + )?; + let action = match body.action { + Some(http_response_body_result::Action::PassThrough(HttpResponseBodyPassThrough {})) => { + BodyAction::PassThrough + } + Some(http_response_body_result::Action::Transform(transform)) => BodyAction::Transform( + validate_replacement(transform.replacement, max_payload_bytes)?, + ), + Some(http_response_body_result::Action::BlockDelivery(_)) => BodyAction::BlockDelivery, + Some(http_response_body_result::Action::SkipRemaining(skip)) => { + let current = match skip.current { + Some(http_response_body_skip_remaining::Current::PassThrough( + HttpResponseBodyPassThrough {}, + )) => CurrentBodyAction::PassThrough, + Some(http_response_body_skip_remaining::Current::Transform(transform)) => { + CurrentBodyAction::Transform(validate_replacement( + transform.replacement, + max_payload_bytes, + )?) + } + None => return Err("invalid_response_body_skip_remaining_action"), + }; + BodyAction::SkipRemaining(current) + } + None => return Err("invalid_response_body_decision"), + }; + Ok(BodyDecision { + action, + reason_code: body.reason_code, + findings: body.findings, + metadata: body.metadata, + }) +} + +fn validate_replacement( + replacement: Option, + max_payload_bytes: usize, +) -> Result, &'static str> { + let Some(http_response_body_transform::Replacement::Data(replacement)) = replacement else { + return Err("response_body_replacement_missing"); + }; + if replacement.len() > max_payload_bytes { + return Err("response_body_replacement_over_capacity"); + } + Ok(replacement) +} + +fn validate_inspect( + entry: &DescribedChainEntry, + inspect: &openshell_core::proto::HttpResponsePreflightInspect, + permitted_modes: &[i32], +) -> Result { + let mode = match HttpResponseBodyMode::try_from(inspect.body_mode) { + Ok(HttpResponseBodyMode::HeadersOnly) => StageMode::HeadersOnly, + Ok(HttpResponseBodyMode::WholeBodyBytes) => StageMode::WholeBody, + Ok(HttpResponseBodyMode::StreamBytes) => StageMode::Stream, + Ok(HttpResponseBodyMode::Unspecified) | Err(_) => { + return Err("invalid_response_body_mode".into()); + } + }; + if !permitted_modes.contains(&inspect.body_mode) { + return Err("response_body_mode_not_permitted".into()); + } + if inspect.header_mutations.len() > headers::MAX_HEADER_MUTATIONS { + return Err("header_mutation_count_over_capacity".into()); + } + let encoded_mutations = inspect + .header_mutations + .iter() + .fold(0usize, |total, mutation| { + total.saturating_add(mutation.encoded_len()) + }); + if encoded_mutations > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES { + return Err("header_mutation_bytes_over_capacity".into()); + } + if entry.max_payload_bytes == 0 && mode != StageMode::HeadersOnly { + return Err("response_payload_limit_invalid".into()); + } + Ok(mode) +} + +fn validate_preflight_input(input: &HttpResponsePreflightInput) -> miette::Result<()> { + if input.context.encoded_len() > MAX_MIDDLEWARE_CONTEXT_BYTES { + return Err(miette::miette!("response context exceeds platform limit")); + } + if input.target.encoded_len() > MAX_MIDDLEWARE_TARGET_BYTES { + return Err(miette::miette!("response target exceeds platform limit")); + } + if input.headers.len() > MAX_MIDDLEWARE_HEADERS { + return Err(miette::miette!( + "response header count exceeds platform limit" + )); + } + if input.headers.iter().fold(0usize, |total, header| { + total.saturating_add(header.encoded_len()) + }) > MAX_MIDDLEWARE_HEADER_BYTES + { + return Err(miette::miette!("response headers exceed platform limit")); + } + Ok(()) +} + +fn validate_diagnostics( + reason: &str, + reason_code: &str, + findings: &[Finding], + metadata: &std::collections::HashMap, +) -> Result<(), &'static str> { + if reason.len() > MAX_MIDDLEWARE_REASON_BYTES { + return Err("response_reason_over_capacity"); + } + if !reason_code.is_empty() + && (reason_code.len() > MAX_MIDDLEWARE_REASON_CODE_BYTES + || !is_stable_reason_code(reason_code)) + { + return Err("response_reason_code_invalid"); + } + if findings.len() > MAX_MIDDLEWARE_FINDINGS_PER_STAGE { + return Err("response_findings_over_capacity"); + } + if findings + .iter() + .any(|finding| finding.encoded_len() > MAX_MIDDLEWARE_FINDING_BYTES) + { + return Err("response_finding_over_capacity"); + } + if metadata.len() > MAX_MIDDLEWARE_METADATA_ENTRIES { + return Err("response_metadata_count_over_capacity"); + } + if metadata.iter().fold(0usize, |total, (key, value)| { + total.saturating_add(key.len()).saturating_add(value.len()) + }) > MAX_MIDDLEWARE_METADATA_BYTES + { + return Err("response_metadata_bytes_over_capacity"); + } + Ok(()) +} + +fn body_restriction(input: &HttpResponsePreflightInput) -> Option { + if input.target.method.eq_ignore_ascii_case("HEAD") + || input.status_code == 204 + || input.status_code == 304 + { + return Some("bodyless_response".into()); + } + if input.status_code == 206 + || input + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("content-range")) + || input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("content-type") + && header + .value + .split(';') + .next() + .is_some_and(|value| value.trim().eq_ignore_ascii_case("multipart/byteranges")) + }) + { + return Some("unsupported_partial_response".into()); + } + if input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("cache-control") + && header.value.split(',').any(|directive| { + directive + .split('=') + .next() + .is_some_and(|name| name.trim().eq_ignore_ascii_case("no-transform")) + }) + }) { + return Some("response_no_transform".into()); + } + if input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("content-encoding") + && header + .value + .split(',') + .any(|coding| !coding.trim().eq_ignore_ascii_case("identity")) + }) { + return Some("unsupported_content_encoding".into()); + } + None +} + +fn permitted_body_modes( + input: &HttpResponsePreflightInput, + entry: &DescribedChainEntry, + body_restriction: Option<&str>, +) -> Vec { + let mut modes = vec![HttpResponseBodyMode::HeadersOnly as i32]; + if body_restriction.is_some() { + return modes; + } + if input + .declared_body_length + .is_none_or(|length| length <= entry.max_payload_bytes as u64) + && !is_open_ended_response(input) + { + modes.push(HttpResponseBodyMode::WholeBodyBytes as i32); + } + if entry.max_payload_bytes >= 2 { + modes.push(HttpResponseBodyMode::StreamBytes as i32); + } + modes +} + +fn is_open_ended_response(input: &HttpResponsePreflightInput) -> bool { + input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("content-type") + && matches!( + header.value.split(';').next().map(str::trim), + Some(value) + if value.eq_ignore_ascii_case("text/event-stream") + || value.eq_ignore_ascii_case("multipart/x-mixed-replace") + ) + }) +} + +fn strip_stale_integrity(headers: &mut Vec) { + headers.retain(|header| { + !matches!( + header.name.to_ascii_lowercase().as_str(), + "accept-ranges" + | "etag" + | "content-md5" + | "digest" + | "content-digest" + | "repr-digest" + | "signature" + | "signature-input" + ) + }); +} + +async fn exchange( + stage: &mut HttpResponseStage, + event: HttpResponseEvent, + chain_deadline: Instant, +) -> Result { + let remaining = chain_deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + return Err("middleware_chain_timeout".into()); + } + let timeout = stage.entry.timeout.min(remaining); + let Some(transport) = stage.transport.as_mut() else { + return Err("middleware_stream_closed".into()); + }; + match tokio::time::timeout(timeout, async { + transport + .sender + .send(event) + .await + .map_err(|_| tonic::Status::unavailable("middleware request stream closed"))?; + transport + .responses + .next() + .await + .ok_or_else(|| tonic::Status::unavailable("middleware result stream closed"))? + }) + .await + { + Ok(Ok(result)) => Ok(result), + Ok(Err(error)) => { + let policy = stage + .entry + .service + .as_ref() + .map_or(MiddlewareDiagnosticPolicy::Preserve, |service| { + service.diagnostic_policy + }); + Err(policy.error_reason(&error)) + } + Err(_) => Err("middleware_timeout".into()), + } +} + +fn body_event(sequence: u64, data: Vec, end_of_stream: bool) -> HttpResponseEvent { + HttpResponseEvent { + event: Some(http_response_event::Event::Body(HttpResponseBodyUnit { + sequence, + payload: Some(http_response_body_unit::Payload::Data(data)), + end_of_stream, + })), + } +} + +fn session_end_event(reason: MiddlewareSessionEndReason) -> HttpResponseEvent { + HttpResponseEvent { + event: Some(http_response_event::Event::SessionEnd( + MiddlewareSessionEnd { + reason: reason as i32, + protocol_error: None, + }, + )), + } +} + +fn body_invocation_with_reason( + stage: &HttpResponseStage, + outcome: HttpResponseInvocationOutcome, + sequence: u64, + input_size: usize, + output_size: usize, + reason_code: Option, +) -> HttpResponseInvocation { + HttpResponseInvocation { + config_name: stage.entry.entry.name.clone(), + implementation: stage.entry.entry.implementation.clone(), + outcome, + sequence: Some(sequence), + input_size, + output_size: Some(output_size), + failed: false, + stage_disabled: false, + reason_code, + failure_category: None, + } +} + +fn collect_diagnostics( + stage: &HttpResponseStage, + mut findings: Vec, + mut metadata: std::collections::HashMap, + all_findings: &mut Vec, + all_metadata: &mut BTreeMap>, +) { + if stage + .entry + .service + .as_ref() + .is_some_and(|service| service.diagnostic_policy == MiddlewareDiagnosticPolicy::Normalize) + { + metadata.clear(); + for finding in &mut findings { + finding.r#type = format!("{}.finding", stage.entry.entry.implementation); + finding.label = super::EXTERNAL_FINDING_LABEL.to_string(); + finding.confidence.clear(); + finding.severity = "medium".into(); + } + } + all_findings.extend(findings.into_iter().map(|finding| NamespacedFinding { + middleware: stage.entry.entry.name.clone(), + finding, + })); + if !metadata.is_empty() { + all_metadata.insert( + stage.entry.entry.name.clone(), + metadata.into_iter().collect(), + ); + } +} + +fn collect_preflight_diagnostics( + entry: &DescribedChainEntry, + findings: Vec, + metadata: std::collections::HashMap, + all_findings: &mut Vec, + all_metadata: &mut BTreeMap>, +) { + let stage = HttpResponseStage { + entry: entry.clone(), + transport: None, + mode: StageMode::HeadersOnly, + next_sequence: 1, + whole_body: Vec::new(), + }; + collect_diagnostics(&stage, findings, metadata, all_findings, all_metadata); +} + +fn collect_preflight_failure( + entry: &DescribedChainEntry, + reason: &str, + invocations: &mut Vec, +) -> Option { + let fail_closed = entry.on_error() == OnError::FailClosed; + invocations.push(HttpResponseInvocation { + config_name: entry.entry.name.clone(), + implementation: entry.entry.implementation.clone(), + outcome: if fail_closed { + HttpResponseInvocationOutcome::FailClosed + } else { + HttpResponseInvocationOutcome::FailOpen + }, + sequence: None, + input_size: 0, + output_size: None, + failed: true, + stage_disabled: true, + reason_code: None, + failure_category: Some(response_failure_category(reason).into()), + }); + fail_closed.then(|| format!("middleware_failed: {reason}")) +} + +fn response_failure_category(reason: &str) -> &'static str { + if reason == "middleware_session_capacity_exhausted" { + "session_capacity" + } else if reason.contains("over_capacity") { + "payload_capacity" + } else if reason.contains("timeout") { + "timeout" + } else if reason.contains("stream_closed") + || reason.contains("stream closed") + || reason.contains("transport") + || reason.contains("unavailable") + { + "transport" + } else if matches!( + reason, + "bodyless_response" + | "partial_response" + | "content_coding_not_identity" + | "cache_control_no_transform" + ) { + "response_not_inspectable" + } else { + "invalid_result" + } +} + +fn empty_preflight_outcome(headers: Vec) -> HttpResponsePreflightOutcome { + HttpResponsePreflightOutcome { + allowed: true, + reason: String::new(), + denial: None, + headers, + session: None, + findings: Vec::new(), + metadata: BTreeMap::new(), + invocations: Vec::new(), + session_capacity_exhausted: false, + } +} + +fn failed_preflight_outcome( + headers: Vec, + reason: String, + findings: Vec, + metadata: BTreeMap>, + invocations: Vec, +) -> HttpResponsePreflightOutcome { + HttpResponsePreflightOutcome { + allowed: false, + reason, + denial: None, + headers, + session: None, + findings, + metadata, + invocations, + session_capacity_exhausted: false, + } +} + +fn blocked_preflight_outcome( + headers: Vec, + denial: super::MiddlewareDenial, + findings: Vec, + metadata: BTreeMap>, + invocations: Vec, +) -> HttpResponsePreflightOutcome { + HttpResponsePreflightOutcome { + allowed: false, + reason: middleware_denial_reason(&denial.config_name, denial.reason_code.as_deref()), + denial: Some(denial), + headers, + session: None, + findings, + metadata, + invocations, + session_capacity_exhausted: false, + } +} + +fn response_session_capacity_exhausted( + entries: Vec, + headers: Vec, +) -> HttpResponsePreflightOutcome { + let mut invocations = Vec::new(); + let fail_closed = entries.iter().any(|entry| { + collect_preflight_failure( + entry, + "middleware_session_capacity_exhausted", + &mut invocations, + ) + .is_some() + }); + HttpResponsePreflightOutcome { + allowed: !fail_closed, + reason: if fail_closed { + "middleware_failed: middleware_session_capacity_exhausted".into() + } else { + String::new() + }, + denial: None, + headers, + session: None, + findings: Vec::new(), + metadata: BTreeMap::new(), + invocations, + session_capacity_exhausted: true, + } +} + +async fn end_stages(stages: &mut [HttpResponseStage], reason: MiddlewareSessionEndReason) { + for stage in stages { + stage.end(reason).await; + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use openshell_core::middleware::{HttpRequestView, InProcessMiddleware}; + use openshell_core::proto::{ + Decision, ExistingHeaderAction, HeaderMutation, HttpRequestResult, HttpResponseBodyResult, + HttpResponseBodyTransform, HttpResponsePreflightInspect, HttpResponsePreflightResult, + HttpResponsePreflightSkip, HttpResponseTrailersResult, MiddlewareBinding, + MiddlewareManifest, WriteHeader, header_mutation, http_response_preflight_result, + }; + use tokio_stream::wrappers::ReceiverStream; + use tokio_stream::wrappers::TcpListenerStream; + + use super::*; + + #[derive(Clone, Copy)] + enum Script { + HeadersOnly, + Stream, + WholeBody, + InvalidSequence, + Configured, + HangBody, + LargeStream, + Skip, + InvalidSkipReason, + TrailerMutation, + InvalidTrailerMutation, + } + + struct ResponseService { + script: Script, + } + + #[derive(Clone)] + struct RemoteResponseService; + + #[tonic::async_trait] + impl openshell_core::proto::middleware::v1::supervisor_middleware_server::SupervisorMiddleware + for RemoteResponseService + { + type EvaluateWebSocketSessionStream = super::super::WebSocketResponseStream; + + async fn describe( + &self, + _request: tonic::Request<()>, + ) -> Result, tonic::Status> { + Ok(tonic::Response::new(response_manifest( + "test/remote-response", + ))) + } + + async fn validate_config( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> + { + Ok(tonic::Response::new( + openshell_core::proto::ValidateConfigResponse { + valid: true, + reason: String::new(), + }, + )) + } + + async fn evaluate_http_request( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Ok(tonic::Response::new(HttpRequestResult { + decision: Decision::Allow as i32, + ..Default::default() + })) + } + + async fn evaluate_web_socket_session( + &self, + _request: tonic::Request< + tonic::Streaming, + >, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented("HTTP response-only service")) + } + } + + #[tonic::async_trait] + impl openshell_core::proto::middleware::v1::http_response_pre_return_server::HttpResponsePreReturn + for RemoteResponseService + { + type EvaluateStream = super::super::HttpResponseResultStream; + + async fn evaluate( + &self, + request: tonic::Request>, + ) -> Result, tonic::Status> { + let mut requests = request.into_inner(); + let (sender, receiver) = mpsc::channel(4); + tokio::spawn(async move { + while let Some(Ok(event)) = requests.next().await { + match event.event { + Some(http_response_event::Event::Preflight(_)) => { + let result = HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some( + http_response_preflight_result::Action::Inspect( + HttpResponsePreflightInspect { + body_mode: + HttpResponseBodyMode::HeadersOnly as i32, + header_mutations: vec![write_header( + "cache-control", + "remote", + )], + }, + ), + ), + ..Default::default() + }, + ), + ), + }; + if sender.send(Ok(result)).await.is_err() { + break; + } + } + Some(http_response_event::Event::SessionEnd(_)) | None => break, + _ => {} + } + } + }); + Ok(tonic::Response::new(Box::pin(ReceiverStream::new(receiver)))) + } + } + + #[tonic::async_trait] + impl InProcessMiddleware for ResponseService { + async fn describe(&self) -> MiddlewareManifest { + MiddlewareManifest { + name: "test/response".into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: openshell_core::proto::SupervisorMiddlewareOperation::HttpResponse + as i32, + phase: openshell_core::proto::SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: if matches!(self.script, Script::LargeStream) { + 128 * 1024 + } else { + 4096 + }, + timeout: if matches!(self.script, Script::HangBody) { + "10ms".into() + } else { + String::new() + }, + }], + expected_audience: String::new(), + } + } + + async fn validate_config( + &self, + _middleware_name: &str, + _config: &prost_types::Struct, + ) -> miette::Result<()> { + Ok(()) + } + + async fn evaluate_http_request( + &self, + _request: HttpRequestView<'_>, + ) -> miette::Result { + Ok(HttpRequestResult { + decision: Decision::Allow as i32, + ..Default::default() + }) + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> Result { + let (sender, receiver) = mpsc::channel(4); + let script = self.script; + tokio::spawn(async move { + let mut selected_script = script; + while let Some(event) = requests.recv().await { + let Some(event) = event.event else { + break; + }; + let result = match event { + http_response_event::Event::Preflight(preflight) => { + if matches!(script, Script::Configured) { + selected_script = match preflight + .config + .as_ref() + .and_then(|config| config.fields.get("mode")) + .and_then(|value| value.kind.as_ref()) + { + Some(prost_types::value::Kind::StringValue(mode)) + if mode == "whole" => + { + Script::WholeBody + } + Some(prost_types::value::Kind::StringValue(mode)) + if mode == "stream" => + { + Script::Stream + } + _ => Script::HeadersOnly, + }; + } + if matches!(selected_script, Script::Skip | Script::InvalidSkipReason) { + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some( + http_response_preflight_result::Action::Skip( + HttpResponsePreflightSkip {}, + ), + ), + reason: if matches!( + selected_script, + Script::InvalidSkipReason + ) { + "x".repeat(MAX_MIDDLEWARE_REASON_BYTES + 1) + } else { + "not selected".into() + }, + reason_code: "path_not_selected".into(), + ..Default::default() + }, + ), + ), + } + } else { + let (body_mode, header_mutations) = match selected_script { + Script::HeadersOnly => ( + HttpResponseBodyMode::HeadersOnly, + vec![write_header("cache-control", "private")], + ), + Script::Stream + | Script::InvalidSequence + | Script::HangBody + | Script::LargeStream + | Script::TrailerMutation + | Script::InvalidTrailerMutation => { + (HttpResponseBodyMode::StreamBytes, Vec::new()) + } + Script::WholeBody => { + (HttpResponseBodyMode::WholeBodyBytes, Vec::new()) + } + Script::Configured + | Script::Skip + | Script::InvalidSkipReason => unreachable!(), + }; + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some( + http_response_preflight_result::Action::Inspect( + HttpResponsePreflightInspect { + body_mode: body_mode as i32, + header_mutations, + }, + ), + ), + ..Default::default() + }, + ), + ), + } + } + } + http_response_event::Event::Body(body) => { + if matches!(selected_script, Script::HangBody) { + continue; + } + let Some(http_response_body_unit::Payload::Data(data)) = body.payload + else { + break; + }; + let replacement = match selected_script { + Script::Stream + | Script::InvalidSequence + | Script::LargeStream + | Script::TrailerMutation + | Script::InvalidTrailerMutation => data.to_ascii_uppercase(), + Script::WholeBody => [b"whole:".as_slice(), &data].concat(), + Script::HeadersOnly + | Script::Configured + | Script::HangBody + | Script::Skip + | Script::InvalidSkipReason => break, + }; + HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: if matches!( + selected_script, + Script::InvalidSequence + ) { + body.sequence + 1 + } else { + body.sequence + }, + action: Some(http_response_body_result::Action::Transform( + HttpResponseBodyTransform { + replacement: Some( + http_response_body_transform::Replacement::Data( + replacement, + ), + ), + }, + )), + ..Default::default() + }, + )), + } + } + http_response_event::Event::Trailers(_) => HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult { + trailer_mutations: match selected_script { + Script::TrailerMutation => { + vec![write_header("x-upstream", "changed")] + } + Script::InvalidTrailerMutation => vec![ + write_header("x-upstream", "changed"), + write_header("x-new", "not-allowed"), + ], + _ => Vec::new(), + }, + ..Default::default() + }, + )), + }, + http_response_event::Event::SessionEnd(_) => break, + }; + if sender.send(Ok(result)).await.is_err() { + break; + } + } + }); + Ok(Box::pin(ReceiverStream::new(receiver))) + } + } + + fn write_header(name: &str, value: &str) -> HeaderMutation { + HeaderMutation { + operation: Some(header_mutation::Operation::Write(WriteHeader { + name: name.into(), + value: value.into(), + on_existing: ExistingHeaderAction::Overwrite as i32, + })), + } + } + + fn response_manifest(name: &str) -> MiddlewareManifest { + MiddlewareManifest { + name: name.into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: openshell_core::proto::SupervisorMiddlewareOperation::HttpResponse + as i32, + phase: openshell_core::proto::SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: 4096, + timeout: String::new(), + }], + expected_audience: String::new(), + } + } + + fn entry(on_error: OnError) -> ChainEntry { + ChainEntry { + name: "response".into(), + implementation: "test/response".into(), + order: 0, + config: prost_types::Struct::default(), + on_error, + } + } + + fn configured_entry(name: &str, order: i32, mode: &str) -> ChainEntry { + ChainEntry { + name: name.into(), + implementation: "test/response".into(), + order, + config: prost_types::Struct { + fields: [( + "mode".into(), + prost_types::Value { + kind: Some(prost_types::value::Kind::StringValue(mode.into())), + }, + )] + .into(), + }, + on_error: OnError::FailClosed, + } + } + + fn input(status_code: u16) -> HttpResponsePreflightInput { + HttpResponsePreflightInput { + context: RequestContext { + request_id: "req-1".into(), + sandbox_id: "sandbox-1".into(), + ..Default::default() + }, + target: HttpRequestTarget { + scheme: "https".into(), + host: "example.com".into(), + port: 443, + method: "GET".into(), + path: "/data".into(), + query: String::new(), + }, + status_code, + declared_body_length: None, + headers: vec![HttpHeader { + name: "content-type".into(), + value: "text/plain".into(), + }], + connection_nominated_headers: Vec::new(), + } + } + + #[tokio::test] + async fn headers_only_preflight_applies_end_to_end_mutation() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::HeadersOnly, + })); + let outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("response preflight"); + + assert!(outcome.allowed); + assert_eq!( + outcome + .headers + .iter() + .find(|header| header.name == "cache-control") + .map(|header| header.value.as_str()), + Some("private") + ); + assert!(outcome.session.is_none()); + } + + #[tokio::test] + async fn stream_mode_transforms_lockstep_units_and_preserves_trailers() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Stream, + })); + let mut outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("response preflight"); + let mut session = outcome.session.take().expect("streaming session"); + + assert_eq!( + session + .push_body(b"hello".to_vec()) + .await + .expect("transform stream unit"), + vec![b"HELLO".to_vec()] + ); + let original_trailers = vec![HttpHeader { + name: "x-upstream".into(), + value: "retained".into(), + }]; + let finish = session + .finish(original_trailers.clone()) + .await + .expect("finish stream"); + assert!(finish.body_units.is_empty()); + assert_eq!(finish.trailers, original_trailers); + } + + #[tokio::test] + async fn whole_body_mode_releases_replacement_only_at_finish() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::WholeBody, + })); + let mut outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("response preflight"); + let mut session = outcome.session.take().expect("whole-body session"); + assert!(session.requires_whole_body()); + assert!( + session + .push_body(b"one".to_vec()) + .await + .expect("buffer first unit") + .is_empty() + ); + assert!( + session + .push_body(b"two".to_vec()) + .await + .expect("buffer second unit") + .is_empty() + ); + + let finish = session.finish(Vec::new()).await.expect("finish whole body"); + assert_eq!(finish.body_units, vec![b"whole:onetwo".to_vec()]); + } + + #[tokio::test] + async fn mixed_profile_chain_respects_policy_order_and_whole_body_barrier() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Configured, + })); + let entries = vec![ + configured_entry("stream", 20, "stream"), + configured_entry("whole", 10, "whole"), + ]; + let mut outcome = runner + .preflight_http_response(&entries, input(200)) + .await + .expect("mixed response preflight"); + let mut session = outcome.session.take().expect("mixed response session"); + assert!(session.requires_whole_body()); + assert!( + session + .push_body(b"hello".to_vec()) + .await + .expect("buffer mixed response") + .is_empty() + ); + let finish = session + .finish(Vec::new()) + .await + .expect("finish mixed chain"); + assert_eq!(finish.body_units, vec![b"WHOLE:HELLO".to_vec()]); + } + + #[tokio::test] + async fn whole_body_overflow_obeys_fail_open_and_fail_closed() { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::WholeBody, + })); + let mut outcome = runner + .preflight_http_response(&[entry(on_error)], input(200)) + .await + .expect("whole-body response preflight"); + let mut session = outcome.session.take().expect("whole-body session"); + let original = vec![b'a'; 4097]; + let pushed = session.push_body(original.clone()).await; + assert_eq!(pushed.is_ok(), allowed); + if allowed { + assert_eq!(pushed.unwrap(), vec![original]); + assert!(!session.requires_whole_body()); + for fill in [b'b', b'c'] { + let unit = vec![fill; MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES]; + assert_eq!( + session + .push_body(unit.clone()) + .await + .expect("fail-open stage must release later units"), + vec![unit] + ); + } + let finish = session.finish(Vec::new()).await.expect("fail-open finish"); + assert!(finish.body_units.is_empty()); + } + } + } + + #[tokio::test] + async fn response_body_timeout_obeys_fail_open_and_fail_closed() { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::HangBody, + })); + let mut outcome = runner + .preflight_http_response(&[entry(on_error)], input(200)) + .await + .expect("timed response preflight"); + let mut session = outcome.session.take().expect("timed response session"); + let result = session.push_body(b"unchanged".to_vec()).await; + assert_eq!(result.is_ok(), allowed); + if let Ok(units) = result { + assert_eq!(units, vec![b"unchanged".to_vec()]); + } + } + } + + #[tokio::test] + async fn stream_unit_limit_never_exceeds_platform_cap() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::LargeStream, + })); + let mut outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("large stream preflight"); + let mut session = outcome.session.take().expect("large stream session"); + assert_eq!( + session.stream_unit_limit(), + MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES + ); + let maximum_unit = vec![b'A'; MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES]; + assert_eq!( + session + .push_body(maximum_unit.clone()) + .await + .expect("maximum stream unit"), + vec![maximum_unit] + ); + assert_eq!( + session + .push_body(vec![b'b'; MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES + 1]) + .await + .expect_err("oversized stream unit") + .reason, + "response_stream_unit_over_capacity" + ); + session + .finish(Vec::new()) + .await + .expect("finish large stream"); + } + + #[tokio::test] + async fn skip_reason_code_is_retained_and_oversized_reason_obeys_on_error() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Skip, + })); + let outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("skip response preflight"); + assert!(outcome.allowed); + assert!(outcome.session.is_none()); + assert_eq!( + outcome.invocations[0].reason_code.as_deref(), + Some("path_not_selected") + ); + + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::InvalidSkipReason, + })); + let outcome = runner + .preflight_http_response(&[entry(on_error)], input(200)) + .await + .expect("invalid skip response preflight"); + assert_eq!(outcome.allowed, allowed); + assert!(outcome.session.is_none()); + } + } + + #[tokio::test] + async fn response_trailers_are_mutated_by_body_stage() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::TrailerMutation, + })); + let mut outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("trailer response preflight"); + let mut session = outcome.session.take().expect("trailer response session"); + session + .push_body(b"body".to_vec()) + .await + .expect("transform response body"); + let trailers = vec![HttpHeader { + name: "x-upstream".into(), + value: "retained".into(), + }]; + let finish = session + .finish(trailers.clone()) + .await + .expect("finish response"); + assert_eq!( + finish.trailers, + vec![HttpHeader { + name: "x-upstream".into(), + value: "changed".into(), + }] + ); + } + + #[tokio::test] + async fn invalid_trailer_mutations_are_atomic_and_keep_failure_diagnostics() { + let trailers = vec![HttpHeader { + name: "x-upstream".into(), + value: "retained".into(), + }]; + for on_error in [OnError::FailOpen, OnError::FailClosed] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::InvalidTrailerMutation, + })); + let mut outcome = runner + .preflight_http_response(&[entry(on_error)], input(200)) + .await + .expect("invalid trailer response preflight"); + let mut session = outcome.session.take().expect("trailer response session"); + session + .push_body(b"body".to_vec()) + .await + .expect("response body exchange"); + session.take_diagnostics(); + + match session.finish(trailers.clone()).await { + Ok(finish) => { + assert_eq!(on_error, OnError::FailOpen); + assert_eq!(finish.trailers, trailers); + assert_eq!( + finish.invocations.last().map(|entry| entry.outcome), + Some(HttpResponseInvocationOutcome::FailOpen) + ); + } + Err(failure) => { + assert_eq!(on_error, OnError::FailClosed); + assert_eq!( + failure + .diagnostics + .invocations + .last() + .map(|entry| entry.outcome), + Some(HttpResponseInvocationOutcome::FailClosed) + ); + } + } + } + } + + #[tokio::test] + async fn invalid_sequence_obeys_fail_open_and_fail_closed() { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::InvalidSequence, + })); + let mut outcome = runner + .preflight_http_response(&[entry(on_error)], input(200)) + .await + .expect("response preflight"); + let mut session = outcome.session.take().expect("stream session"); + let result = session.push_body(b"unchanged".to_vec()).await; + assert_eq!(result.is_ok(), allowed); + if let Ok(units) = result { + assert_eq!(units, vec![b"unchanged".to_vec()]); + } + } + } + + #[tokio::test] + async fn body_inspection_restrictions_obey_fail_open_and_fail_closed() { + let mut cases = Vec::new(); + cases.push(input(206)); + for (name, value) in [ + ("content-range", "bytes 0-3/10"), + ("content-type", "multipart/byteranges; boundary=test"), + ("cache-control", "private, no-transform"), + ("content-encoding", "gzip"), + ] { + let mut candidate = input(200); + candidate.headers.push(HttpHeader { + name: name.into(), + value: value.into(), + }); + cases.push(candidate); + } + for status in [204, 304] { + cases.push(input(status)); + } + let mut head = input(200); + head.target.method = "HEAD".into(); + cases.push(head); + + for candidate in cases { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Stream, + })); + let outcome = runner + .preflight_http_response(&[entry(on_error)], candidate.clone()) + .await + .expect("restricted response preflight"); + assert_eq!(outcome.allowed, allowed); + assert!(outcome.session.is_none()); + } + } + } + + #[tokio::test] + async fn remote_service_executes_through_http_response_pre_return_rpc() { + use openshell_core::proto::middleware::v1::http_response_pre_return_server::HttpResponsePreReturnServer; + use openshell_core::proto::middleware::v1::supervisor_middleware_server::SupervisorMiddlewareServer; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind response middleware"); + let address = listener.local_addr().expect("response middleware address"); + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); + let server = tonic::transport::Server::builder() + .add_service(SupervisorMiddlewareServer::new(RemoteResponseService)) + .add_service(HttpResponsePreReturnServer::new(RemoteResponseService)) + .serve_with_incoming_shutdown(TcpListenerStream::new(listener), async { + let _ = shutdown_rx.await; + }); + let server_task = tokio::spawn(server); + let registry = super::super::MiddlewareRegistry::connect_services( + Vec::new(), + vec![openshell_core::proto::SupervisorMiddlewareService { + name: "remote-response".into(), + grpc_endpoint: format!("http://{address}"), + max_payload_bytes: 4096, + allow_insecure_transport: true, + ..Default::default() + }], + ) + .await + .expect("connect remote response middleware"); + let runner = ChainRunner::from_registry(registry); + let outcome = runner + .preflight_http_response( + &[ChainEntry { + name: "response".into(), + implementation: "remote-response".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: OnError::FailClosed, + }], + input(200), + ) + .await + .expect("remote response preflight"); + + assert!(outcome.allowed); + assert_eq!( + outcome + .headers + .iter() + .find(|header| header.name == "cache-control") + .map(|header| header.value.as_str()), + Some("remote") + ); + assert!(outcome.session.is_none()); + + let _ = shutdown_tx.send(()); + server_task + .await + .expect("join response middleware server") + .expect("serve response middleware"); + } +} diff --git a/crates/openshell-supervisor-network/src/l7/middleware.rs b/crates/openshell-supervisor-network/src/l7/middleware.rs index 6305653f6a..f8a22a9bb9 100644 --- a/crates/openshell-supervisor-network/src/l7/middleware.rs +++ b/crates/openshell-supervisor-network/src/l7/middleware.rs @@ -197,7 +197,7 @@ pub(super) fn websocket_message_finding_events( middleware_finding_events(&outcome.findings) } -fn middleware_finding_events( +pub(super) fn middleware_finding_events( findings: &[openshell_supervisor_middleware::NamespacedFinding], ) -> Vec { findings @@ -423,7 +423,8 @@ pub(super) fn middleware_chain_body_limit( .max() } -pub async fn apply_middleware_chain( +#[allow(clippy::too_many_arguments)] +pub async fn apply_middleware_chain_with_request_id( req: crate::l7::provider::L7Request, client: &mut C, ctx: &L7EvalContext, @@ -431,8 +432,9 @@ pub async fn apply_middleware_chain( runner: &openshell_supervisor_middleware::ChainRunner, generation_guard: &PolicyGenerationGuard, transformed_body_policy: openshell_supervisor_middleware::TransformedBodyPolicy<'_>, + request_id: &str, ) -> Result { - apply_middleware_chain_for_scheme( + apply_middleware_chain_for_scheme_with_request_id( req, client, ctx, @@ -441,12 +443,15 @@ pub async fn apply_middleware_chain( runner, generation_guard, transformed_body_policy, + request_id, ) .await } #[allow(clippy::too_many_arguments)] -pub async fn apply_middleware_chain_for_scheme( +pub async fn apply_middleware_chain_for_scheme_with_request_id< + C: AsyncRead + AsyncWrite + Unpin + Send, +>( req: crate::l7::provider::L7Request, client: &mut C, ctx: &L7EvalContext, @@ -455,6 +460,7 @@ pub async fn apply_middleware_chain_for_scheme, + request_id: &str, ) -> Result { if chain.is_empty() { return Ok(MiddlewareApplyResult::Allowed(req)); @@ -479,7 +485,7 @@ pub async fn apply_middleware_chain_for_scheme, query: String, body: Vec, + request_id: &str, ) -> openshell_supervisor_middleware::HttpRequestInput { openshell_supervisor_middleware::HttpRequestInput { - request_id: uuid::Uuid::new_v4().to_string(), + request_id: request_id.to_string(), sandbox_id: sandbox.sandbox_id.clone(), sandbox_name: sandbox.sandbox_name.clone(), workspace: ctx.workspace.clone(), @@ -637,6 +646,32 @@ pub(super) fn middleware_request_input( } } +#[cfg(test)] +#[allow(clippy::too_many_arguments)] +pub(super) fn middleware_request_input( + sandbox: &openshell_ocsf::SandboxContext, + scheme: &str, + req: &crate::l7::provider::L7Request, + ctx: &L7EvalContext, + headers: Vec<(String, String)>, + connection_nominated_headers: Vec, + query: String, + body: Vec, +) -> openshell_supervisor_middleware::HttpRequestInput { + let request_id = uuid::Uuid::new_v4().to_string(); + middleware_request_input_with_id( + sandbox, + scheme, + req, + ctx, + headers, + connection_nominated_headers, + query, + body, + &request_id, + ) +} + pub(super) fn raw_query_from_request_headers(headers: &[u8]) -> Result { let header_str = std::str::from_utf8(headers).map_err(|_| miette!("HTTP headers contain invalid UTF-8"))?; @@ -1116,7 +1151,7 @@ mod tests { body_length: crate::l7::provider::BodyLength::None, }; - let input = super::middleware_request_input( + let input = super::middleware_request_input_with_id( &sandbox, "https", &req, @@ -1125,11 +1160,13 @@ mod tests { Vec::new(), String::new(), Vec::new(), + "exchange-123", ); assert_eq!(input.sandbox_name, "nightly-build"); assert_eq!(input.sandbox_id, "sbx-123"); assert_eq!(input.workspace, "wrks-default"); + assert_eq!(input.request_id, "exchange-123"); } #[tokio::test] diff --git a/crates/openshell-supervisor-network/src/l7/relay.rs b/crates/openshell-supervisor-network/src/l7/relay.rs index 9e4949b191..2a5de4d6a8 100644 --- a/crates/openshell-supervisor-network/src/l7/relay.rs +++ b/crates/openshell-supervisor-network/src/l7/relay.rs @@ -8,7 +8,7 @@ //! and either forwards or denies the request. use crate::l7::middleware::{ - MiddlewareApplyResult, UninspectableTrafficGate, apply_middleware_chain, + MiddlewareApplyResult, UninspectableTrafficGate, apply_middleware_chain_with_request_id, emit_middleware_uninspectable, middleware_network_input, uninspectable_traffic_gate, }; #[cfg(test)] @@ -288,13 +288,20 @@ async fn relay_http_request_with_credential_rejection( upstream: &mut U, options: crate::l7::rest::RelayRequestOptions<'_>, ctx: &L7EvalContext, + response_middleware: Option>, ) -> Result> where C: AsyncRead + AsyncWrite + Unpin, U: AsyncRead + AsyncWrite + Unpin, { - match crate::l7::rest::relay_http_request_with_options_guarded( - request, client, upstream, options, + match Box::pin( + crate::l7::rest::relay_http_request_with_response_middleware_guarded( + request, + client, + upstream, + options, + response_middleware, + ), ) .await { @@ -310,6 +317,90 @@ where } } +pub(crate) fn http_response_middleware_relay<'a>( + request: &crate::l7::provider::L7Request, + ctx: &'a L7EvalContext, + scheme: &str, + request_id: &str, + chain: &'a [openshell_supervisor_middleware::ChainEntry], + runner: &'a openshell_supervisor_middleware::ChainRunner, + generation_guard: Option<&'a PolicyGenerationGuard>, +) -> crate::l7::rest::HttpResponseMiddlewareRelay<'a> { + let sandbox = openshell_ocsf::ctx::ctx(); + crate::l7::rest::HttpResponseMiddlewareRelay { + chain, + runner, + request_context: openshell_core::proto::RequestContext { + request_id: request_id.to_string(), + sandbox_id: sandbox.sandbox_id.clone(), + sandbox_name: sandbox.sandbox_name.clone(), + workspace: ctx.workspace.clone(), + originating_process: None, + }, + target: openshell_core::proto::HttpRequestTarget { + scheme: scheme.to_string(), + host: ctx.host.clone(), + port: u32::from(ctx.port), + method: request.action.clone(), + path: request.target.clone(), + query: policy_safe_response_query(&request.query_params), + }, + policy_name: &ctx.policy_name, + generation_guard, + whole_body_timeout: crate::http_response_whole_body_timeout(), + } +} + +fn policy_safe_response_query( + query_params: &std::collections::HashMap>, +) -> String { + let mut parameters: Vec<_> = query_params.iter().collect(); + parameters.sort_by_key(|(name, _)| *name); + let mut output = String::new(); + for (name, values) in parameters { + let empty_value = String::new(); + let values = if values.is_empty() { + std::slice::from_ref(&empty_value) + } else { + values.as_slice() + }; + for value in values { + if !output.is_empty() { + output.push('&'); + } + let name = if secrets::contains_reserved_credential_marker(name) { + "[REDACTED]" + } else { + name + }; + let value = if secrets::contains_reserved_credential_marker(value) { + "[REDACTED]" + } else { + value + }; + push_form_component(&mut output, name); + output.push('='); + push_form_component(&mut output, value); + } + } + output +} + +fn push_form_component(output: &mut String, value: &str) { + const HEX: &[u8; 16] = b"0123456789ABCDEF"; + for byte in value.bytes() { + if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') { + output.push(char::from(byte)); + } else if byte == b' ' { + output.push('+'); + } else { + output.push('%'); + output.push(char::from(HEX[usize::from(byte >> 4)])); + output.push(char::from(HEX[usize::from(byte & 0x0f)])); + } + } +} + #[derive(Default)] pub(crate) struct UpgradeRelayOptions<'a> { pub(crate) websocket_request: bool, @@ -736,13 +827,15 @@ where if allowed || (config.enforcement == EnforcementMode::Audit && !force_deny) { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); let websocket_chain = websocket_request.then(|| chain.clone()); // Route selection resolved `config` per request, so re-check the // body against that protocol's policy after every transforming // stage (a no-op for REST and websocket, whose policy inputs the // chain cannot mutate). let validate = transformed_body_validator(config, &engine, ctx, &request_info); - let middleware_result = apply_middleware_chain( + let middleware_result = apply_middleware_chain_with_request_id( req, client, ctx, @@ -750,6 +843,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await; let req = match middleware_result? { @@ -848,6 +942,15 @@ where port: ctx.port, }, ctx, + Some(http_response_middleware_relay( + &req, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await; let outcome_result = match outcome_result { @@ -1461,11 +1564,13 @@ where if allowed || config.enforcement == EnforcementMode::Audit { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); let websocket_chain = websocket_request.then(|| chain.clone()); // REST and websocket-upgrade policy evaluates only the method, // path, and query, which a middleware result cannot mutate, so no // per-stage body re-check is needed. - let middleware_result = apply_middleware_chain( + let middleware_result = apply_middleware_chain_with_request_id( req, client, ctx, @@ -1473,6 +1578,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + &request_id, ) .await; let req = match middleware_result? { @@ -1587,6 +1693,15 @@ where port: ctx.port, }, ctx, + Some(http_response_middleware_relay( + &req_with_auth, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await; let outcome_result = match outcome_result { @@ -1875,12 +1990,14 @@ where if allowed || (config.enforcement == EnforcementMode::Audit && !force_deny) { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); // Policy admitted the original body above; re-check the body // against the same body-aware policy after every transforming // stage so a middleware cannot smuggle a denied operation to the // upstream or the next stage. let validate = transformed_body_validator(config, engine, ctx, &request_info); - let req = match apply_middleware_chain( + let req = match apply_middleware_chain_with_request_id( req, client, ctx, @@ -1888,6 +2005,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await? { @@ -1946,6 +2064,15 @@ where ..Default::default() }, ctx, + Some(http_response_middleware_relay( + &req, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await? else { @@ -2115,12 +2242,14 @@ where if allowed || (config.enforcement == EnforcementMode::Audit && !force_deny) { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); // Policy admitted the original body above; re-check the body // against the same body-aware policy after every transforming // stage so a middleware cannot smuggle a denied operation to the // upstream or the next stage. let validate = transformed_body_validator(config, engine, ctx, &request_info); - let req = match apply_middleware_chain( + let req = match apply_middleware_chain_with_request_id( req, client, ctx, @@ -2128,6 +2257,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await? { @@ -2181,6 +2311,15 @@ where ..Default::default() }, ctx, + Some(http_response_middleware_relay( + &req, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await? else { @@ -2733,6 +2872,8 @@ where ocsf_emit!(event); } + let request_id = uuid::Uuid::new_v4().to_string(); + let mut response_selection = None; let req = if let Some(engine) = middleware_engine { let input = middleware_network_input(ctx); let (chain, generation) = engine.query_middleware_chain_with_generation(&input)?; @@ -2740,9 +2881,10 @@ where return Ok(()); } let runner = engine.middleware_runner()?; + response_selection = Some((chain.clone(), runner.clone())); // The passthrough path enforces no L7 policy, so there is no // body-aware decision to re-check after a transformation. - match apply_middleware_chain( + match apply_middleware_chain_with_request_id( req, client, ctx, @@ -2750,6 +2892,7 @@ where &runner, generation_guard, openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + &request_id, ) .await? { @@ -2811,6 +2954,17 @@ where let scoped_ctx = scoped_context_for_request(ctx, &req_with_auth); let ctx = scoped_ctx.as_ref().unwrap_or(ctx); let resolver = ctx.secret_resolver.as_deref(); + let response_middleware = response_selection.as_ref().map(|(chain, runner)| { + http_response_middleware_relay( + &req_with_auth, + ctx, + "http", + &request_id, + chain, + runner, + Some(generation_guard), + ) + }); // Forward request with credential rewriting and relay the response. // relay_http_request_with_resolver handles both directions: it sends @@ -2826,6 +2980,7 @@ where ..Default::default() }, ctx, + response_middleware, ) .await? else { @@ -3132,6 +3287,7 @@ mod tests { ..options }, &ctx, + None, ) .await .expect("typed credential denial"); @@ -6055,7 +6211,7 @@ network_policies: let (mut app, mut relay_client) = tokio::io::duplex(8192); app.write_all(&body).await.unwrap(); - let result = crate::l7::middleware::apply_middleware_chain_for_scheme( + let result = crate::l7::middleware::apply_middleware_chain_for_scheme_with_request_id( req, &mut relay_client, &ctx, @@ -6064,6 +6220,7 @@ network_policies: &runner, tunnel_engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + "test-request-id", ) .await .expect("apply middleware chain"); @@ -6191,6 +6348,50 @@ network_policies: assert_eq!(input.scheme, "http"); } + #[test] + fn response_middleware_context_reuses_exchange_request_id() { + let req = crate::l7::provider::L7Request { + action: "GET".into(), + target: "/v1/data".into(), + query_params: std::collections::HashMap::from([ + ("cursor".into(), vec!["next page".into()]), + ( + "token".into(), + vec!["openshell:resolve:env:API_TOKEN".into()], + ), + ]), + raw_header: b"GET /v1/data?cursor=next+page&token=sk-live-secret HTTP/1.1\r\nHost: api.example.test\r\n\r\n".to_vec(), + body_length: crate::l7::provider::BodyLength::None, + }; + let ctx = L7EvalContext { + host: "api.example.test".into(), + port: 443, + workspace: "workspace-1".into(), + policy_name: "api".into(), + ..Default::default() + }; + let runner = openshell_supervisor_middleware::ChainRunner::default(); + let chain = Vec::new(); + let response = http_response_middleware_relay( + &req, + &ctx, + "https", + "exchange-123", + &chain, + &runner, + None, + ); + + assert_eq!(response.request_context.request_id, "exchange-123"); + assert_eq!( + response.target.query, + "cursor=next+page&token=%5BREDACTED%5D" + ); + assert!(!response.target.query.contains("sk-live-secret")); + assert!(!response.target.query.contains("API_TOKEN")); + assert_eq!(response.target.scheme, "https"); + } + #[test] fn middleware_ocsf_events_are_audit_safe() { use openshell_supervisor_middleware::{ diff --git a/crates/openshell-supervisor-network/src/l7/rest.rs b/crates/openshell-supervisor-network/src/l7/rest.rs index 93315a671a..ec0da7d135 100644 --- a/crates/openshell-supervisor-network/src/l7/rest.rs +++ b/crates/openshell-supervisor-network/src/l7/rest.rs @@ -12,7 +12,10 @@ use crate::opa::PolicyGenerationGuard; use aws_sigv4::http_request::SignableBody; use base64::Engine as _; use miette::{IntoDiagnostic, Result, miette}; -use openshell_core::proto::{ExistingHeaderAction, HeaderMutation, header_mutation}; +use openshell_core::proto::{ + ExistingHeaderAction, HeaderMutation, HttpHeader, HttpRequestTarget, RequestContext, + header_mutation, +}; use openshell_core::secrets::{ CREDENTIAL_MARKER_SCAN_TAIL_BYTES, SecretResolver, contains_reserved_credential_marker, contains_reserved_credential_marker_bytes, rewrite_http_header_block, @@ -20,7 +23,7 @@ use openshell_core::secrets::{ use openshell_ocsf::ctx::ctx as ocsf_ctx; use sha1::{Digest, Sha1}; use std::collections::{HashMap, HashSet}; -use std::fmt; +use std::fmt::{self, Write as _}; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use tracing::debug; @@ -49,6 +52,7 @@ async fn max_middleware_body_bytes() -> usize { chain[0].max_payload_bytes() } const RELAY_BUF_SIZE: usize = 8192; +const RESPONSE_UNIT_COALESCE_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(2); const HTTP_METHOD_PREFIXES: &[&[u8]] = &[ b"GET ", b"HEAD ", @@ -798,6 +802,31 @@ pub(crate) async fn relay_http_request_with_options_guarded( upstream: &mut U, options: RelayRequestOptions<'_>, ) -> Result +where + C: AsyncRead + AsyncWrite + Unpin, + U: AsyncRead + AsyncWrite + Unpin, +{ + relay_http_request_with_response_middleware_guarded(req, client, upstream, options, None).await +} + +/// Context retained from request evaluation for the matching response hook. +pub(crate) struct HttpResponseMiddlewareRelay<'a> { + pub(crate) chain: &'a [openshell_supervisor_middleware::ChainEntry], + pub(crate) runner: &'a openshell_supervisor_middleware::ChainRunner, + pub(crate) request_context: RequestContext, + pub(crate) target: HttpRequestTarget, + pub(crate) policy_name: &'a str, + pub(crate) generation_guard: Option<&'a PolicyGenerationGuard>, + pub(crate) whole_body_timeout: std::time::Duration, +} + +pub(crate) async fn relay_http_request_with_response_middleware_guarded( + req: &L7Request, + client: &mut C, + upstream: &mut U, + options: RelayRequestOptions<'_>, + response_middleware: Option>, +) -> Result where C: AsyncRead + AsyncWrite + Unpin, U: AsyncRead + AsyncWrite + Unpin, @@ -1152,6 +1181,7 @@ where websocket: websocket_response, client_requested_upgrade, }, + response_middleware, ) .await?; @@ -3124,6 +3154,7 @@ async fn relay_response( upstream: &mut U, client: &mut C, options: RelayResponseOptions, + response_middleware: Option>, ) -> Result where U: AsyncRead + Unpin, @@ -3133,7 +3164,8 @@ where let mut buf = Vec::with_capacity(4096); let mut tmp = [0u8; 1024]; - // Read response headers + // Read response headers. Forward interim responses unchanged, but retain + // the final response head until response middleware preflight completes. loop { if buf.len() > MAX_HEADER_BYTES { return Err(miette!("HTTP response headers exceed limit")); @@ -3149,6 +3181,21 @@ where } buf.extend_from_slice(&tmp[..n]); + while let Some(position) = buf.windows(4).position(|w| w == b"\r\n\r\n") { + let header_end = position + 4; + let header_str = String::from_utf8_lossy(&buf[..header_end]); + let status_code = parse_status_code(&header_str).unwrap_or(200); + if (100..200).contains(&status_code) && status_code != 101 { + client + .write_all(&buf[..header_end]) + .await + .into_diagnostic()?; + client.flush().await.into_diagnostic()?; + buf.drain(..header_end); + continue; + } + break; + } if buf.windows(4).any(|w| w == b"\r\n\r\n") { break; } @@ -3204,6 +3251,24 @@ where }); } + if let Some(response_middleware) = response_middleware + && let Some(outcome) = Box::pin(relay_response_through_middleware( + request_method, + upstream, + client, + response_middleware, + &buf, + header_end, + status_code, + body_length, + server_wants_close, + event_stream, + )) + .await? + { + return Ok(outcome); + } + // Bodiless responses (HEAD, 1xx, 204, 304): forward headers only, skip body if is_bodiless_response(request_method, status_code) { client @@ -3292,584 +3357,2342 @@ where Ok(RelayOutcome::Reusable) } -/// Parse the HTTP status code from a response status line. -/// -/// Expects the first line to look like `HTTP/1.1 200 OK`. -fn parse_status_code(headers: &str) -> Option { - let status_line = headers.lines().next()?; - let code_str = status_line.split_whitespace().nth(1)?; - code_str.parse().ok() -} - -/// Check if the response headers contain `Connection: close`. -fn parse_connection_close(headers: &str) -> bool { - for line in headers.lines().skip(1) { - let lower = line.to_ascii_lowercase(); - if lower.starts_with("connection:") { - let val = lower.split_once(':').map_or("", |(_, v)| v.trim()); - return val.contains("close"); +#[allow(clippy::too_many_arguments)] +async fn relay_response_through_middleware( + request_method: &str, + upstream: &mut U, + client: &mut C, + middleware: HttpResponseMiddlewareRelay<'_>, + buffered: &[u8], + header_end: usize, + status_code: u16, + body_length: BodyLength, + server_wants_close: bool, + event_stream: bool, +) -> Result> +where + U: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + if let Some(guard) = middleware.generation_guard { + guard.ensure_current()?; + } + let header_bytes = &buffered[..header_end]; + let parsed = match parse_response_head_for_middleware(header_bytes) { + Ok(parsed) => parsed, + Err(error) => { + debug!(error = %error, "HTTP response head normalization failed"); + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); + } + }; + let original_headers = parsed.headers.clone(); + let upstream_declared_trailers = parsed.declared_trailers.clone(); + let connection_nominated_headers = parsed.connection_nominated.clone(); + let input = openshell_supervisor_middleware::HttpResponsePreflightInput { + context: middleware.request_context, + target: middleware.target.clone(), + status_code, + declared_body_length: match body_length { + BodyLength::ContentLength(length) => Some(length), + BodyLength::Chunked | BodyLength::None => None, + }, + headers: parsed.headers, + connection_nominated_headers: parsed.connection_nominated, + }; + let preflight = match middleware + .runner + .preflight_http_response(middleware.chain, input) + .await + { + Ok(preflight) => preflight, + Err(error) => { + debug!(error = %error, "HTTP response middleware preflight failed"); + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); + } + }; + debug!( + configured_stage_count = middleware.chain.len(), + active_session = preflight.session.is_some(), + allowed = preflight.allowed, + "HTTP response middleware preflight completed" + ); + for event in crate::l7::middleware::middleware_finding_events(&preflight.findings) { + openshell_ocsf::ocsf_emit!(event); + } + if !preflight.allowed { + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &preflight.invocations, + ); + if let Some(denial) = preflight.denial.as_ref() { + send_response_middleware_denial( + client, + request_method, + middleware.policy_name, + &middleware.target, + denial, + ) + .await?; + } else { + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; } + return Ok(Some(RelayOutcome::Consumed)); } - false -} - -fn response_is_event_stream(headers: &str) -> bool { - headers.lines().skip(1).any(|line| { - let lower = line.to_ascii_lowercase(); - let Some(value) = lower.strip_prefix("content-type:") else { - return false; - }; - value - .split(';') - .next() - .is_some_and(|mime| mime.trim() == "text/event-stream") - }) -} + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &preflight.invocations, + ); -fn validate_websocket_response( - headers: &str, - mode: WebSocketExtensionMode, - websocket: Option<&WebSocketResponseValidation>, -) -> Result<(bool, Option)> { - let Some(validation) = websocket else { - return validate_websocket_response_extensions_preserved(headers, mode) - .map(|compressed| (compressed, None)); + let status_line = response_status_line(header_bytes)?; + let Some(mut session) = preflight.session else { + if preflight.headers == original_headers { + return Ok(None); + } + let outcome = relay_headers_only_response( + request_method, + upstream, + client, + &status_line, + &preflight.headers, + &upstream_declared_trailers, + &buffered[header_end..], + status_code, + body_length, + server_wants_close, + event_stream, + ) + .await?; + return Ok(Some(outcome)); }; - let mut upgrade_websocket = false; - let mut connection_upgrade = false; - let mut accept_count = 0usize; - let mut accept_matches = false; - let mut subprotocol_count = 0usize; - let mut selected_subprotocol = None; - - for line in headers.lines().skip(1) { - let Some((name, value)) = line.split_once(':') else { - continue; - }; - let name = name.trim().to_ascii_lowercase(); - let value = value.trim(); - match name.as_str() { - "upgrade" if header_value_contains_token(value, "websocket") => { - upgrade_websocket = true; - } - "connection" if header_value_contains_token(value, "upgrade") => { - connection_upgrade = true; - } - "sec-websocket-accept" => { - accept_count += 1; - accept_matches = value == validation.expected_accept; - } - "sec-websocket-protocol" => { - subprotocol_count += 1; - if !is_http_token(value) { - return Err(miette!( - "websocket upgrade response has invalid Sec-WebSocket-Protocol" - )); - } - selected_subprotocol = Some(value.to_string()); + let supports_chunked_response = !status_line.starts_with("HTTP/1.0 "); + let bodiless = is_bodiless_response(request_method, status_code); + if bodiless { + let finish = match session.finish(Vec::new()).await { + Ok(finish) => finish, + Err(error) => { + debug!(error = %error, "HTTP response middleware finalization failed"); + emit_http_response_diagnostics( + middleware.policy_name, + &middleware.target, + status_code, + &error.diagnostics, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); } - _ => {} - } + }; + let mut headers = preflight.headers; + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &finish.invocations, + ); + for event in crate::l7::middleware::middleware_finding_events(&finish.findings) { + openshell_ocsf::ocsf_emit!(event); + } + if finish.strip_stale_integrity_headers { + strip_response_integrity_headers(&mut headers); + } + let head = serialize_response_head( + &status_line, + &headers, + ResponseFraming::Preserve(body_length), + server_wants_close, + &[], + ); + client.write_all(&head).await.into_diagnostic()?; + client.flush().await.into_diagnostic()?; + return Ok(Some(if server_wants_close { + RelayOutcome::Consumed + } else { + RelayOutcome::Reusable + })); } - if !upgrade_websocket { - return Err(miette!( - "websocket upgrade response missing Upgrade: websocket" - )); - } - if !connection_upgrade { - return Err(miette!( - "websocket upgrade response missing Connection: Upgrade" - )); - } - if accept_count != 1 || !accept_matches { - return Err(miette!( - "websocket upgrade response has invalid Sec-WebSocket-Accept" - )); + let whole_body = session.requires_whole_body(); + if whole_body { + session.start_whole_body_deadline(middleware.whole_body_timeout); } - if subprotocol_count > 1 { - return Err(miette!( - "websocket upgrade response has multiple Sec-WebSocket-Protocol headers" - )); - } - if let Some(ref protocol) = selected_subprotocol - && !validation - .offered_subprotocols - .iter() - .any(|offered| offered == protocol) - { - return Err(miette!( - "upstream selected WebSocket subprotocol that was not offered" - )); + let unit_limit = session.stream_unit_limit().max(1); + let chunked_output = supports_chunked_response; + let close_delimited_output = !supports_chunked_response; + let declared_trailers = upstream_declared_trailers; + let downstream_trailers = if chunked_output { + declared_trailers.as_slice() + } else { + &[] + }; + let streaming_head = serialize_response_head( + &status_line, + &preflight.headers, + if chunked_output { + ResponseFraming::Chunked + } else { + ResponseFraming::Preserve(BodyLength::None) + }, + server_wants_close || close_delimited_output, + downstream_trailers, + ); + let mut committed = !whole_body; + if committed { + if let Err(error) = client.write_all(&streaming_head).await { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect) + .await; + return Err(error).into_diagnostic(); + } + if let Err(error) = client.flush().await { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect) + .await; + return Err(error).into_diagnostic(); + } } - let actual_extension = normalized_websocket_extension(headers)?; - match (&validation.expected_extension, actual_extension.as_deref()) { - (None, Some(_)) => Err(miette!( - "upstream negotiated WebSocket extension that was not offered" - )), - (None | Some(_), None) => Ok((false, selected_subprotocol)), - (Some(expected), Some(actual)) if expected.eq_ignore_ascii_case(actual) => { - Ok((true, selected_subprotocol)) + let mut reader = BufferedResponseReader::new(upstream, &buffered[header_end..]); + let body_result = relay_normalized_response_body( + &mut reader, + &mut session, + client, + body_length, + server_wants_close, + event_stream, + &mut committed, + chunked_output, + &streaming_head, + unit_limit, + middleware.generation_guard, + middleware.policy_name, + &middleware.target, + status_code, + &connection_nominated_headers, + ) + .await; + let trailers = match body_result { + Ok(trailers) => trailers, + Err(error) => { + let middleware_stop = error.downcast_ref::(); + let end_reason = if middleware + .generation_guard + .is_some_and(PolicyGenerationGuard::is_stale) + { + openshell_core::proto::MiddlewareSessionEndReason::PolicyReload + } else if error + .to_string() + .starts_with("HTTP response client write failed:") + { + openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect + } else if middleware_stop.is_some_and(|stop| stop.failure.denial.is_some()) { + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareDenial + } else if middleware_stop.is_some() { + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareFailure + } else { + openshell_core::proto::MiddlewareSessionEndReason::UpstreamDisconnect + }; + session.end(end_reason).await; + if committed { + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + true, + ); + return Err(error); + } + debug!(error = %error, "HTTP response processing failed before commitment"); + if let Some(denial) = middleware_stop.and_then(|stop| stop.failure.denial.as_ref()) { + send_response_middleware_denial( + client, + request_method, + middleware.policy_name, + &middleware.target, + denial, + ) + .await?; + } else { + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + } + return Ok(Some(RelayOutcome::Consumed)); } - (Some(_), Some(_)) => Err(miette!( - "upstream negotiated WebSocket extension that does not match the safe offer" - )), + }; + + if let Some(guard) = middleware.generation_guard + && let Err(error) = guard.ensure_current() + { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) + .await; + return Err(error); } -} -fn validate_websocket_response_extensions_preserved( - headers: &str, - mode: WebSocketExtensionMode, -) -> Result { - match mode { - WebSocketExtensionMode::Preserve => Ok(false), - WebSocketExtensionMode::PermessageDeflate => { - let offers = websocket_extension_offers(headers)?; - if offers.is_empty() { - Ok(false) + let finish = match session.finish(trailers).await { + Ok(finish) => finish, + Err(error) => { + emit_http_response_diagnostics( + middleware.policy_name, + &middleware.target, + status_code, + &error.diagnostics, + ); + if committed { + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + true, + ); + return Err(miette!( + "HTTP response middleware failed after commitment: {error}" + )); + } + debug!(error = %error, "HTTP response middleware failed before commitment"); + if let Some(denial) = error.denial.as_ref() { + send_response_middleware_denial( + client, + request_method, + middleware.policy_name, + &middleware.target, + denial, + ) + .await?; } else { - Err(miette!( - "upstream negotiated WebSocket extension that was not offered" - )) + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; } + return Ok(Some(RelayOutcome::Consumed)); } + }; + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &finish.invocations, + ); + for event in crate::l7::middleware::middleware_finding_events(&finish.findings) { + openshell_ocsf::ocsf_emit!(event); } -} -fn normalized_websocket_extension(headers: &str) -> Result> { - let offers = websocket_extension_offers(headers)?; - if offers.is_empty() { - return Ok(None); - } - if offers.len() != 1 { - return Err(miette!("upstream negotiated multiple WebSocket extensions")); - } - let offer = &offers[0]; - if !offer.name.eq_ignore_ascii_case("permessage-deflate") { - return Err(miette!( - "upstream negotiated unsupported WebSocket extension" - )); - } - let mut client_no_context_takeover = false; - let mut server_no_context_takeover = false; - let mut seen = HashSet::new(); - for param in &offer.params { - let name = param.name.to_ascii_lowercase(); - if param.value.is_some() || !seen.insert(name.clone()) { - return Err(miette!( - "upstream negotiated unsupported permessage-deflate parameter" - )); + if whole_body && !committed { + let mut headers = preflight.headers; + if finish.strip_stale_integrity_headers { + strip_response_integrity_headers(&mut headers); } - if name == "client_no_context_takeover" { - client_no_context_takeover = true; - } else if name == "server_no_context_takeover" { - server_no_context_takeover = true; + let output_length = finish + .body_units + .iter() + .try_fold(0usize, |total, unit| total.checked_add(unit.len())) + .ok_or_else(|| miette!("HTTP response middleware output length overflow"))?; + let framing = if finish.trailers.is_empty() || !supports_chunked_response { + ResponseFraming::ContentLength(output_length as u64) } else { - return Err(miette!( - "upstream negotiated unsupported permessage-deflate parameter" - )); - } - } - let mut normalized = String::from("permessage-deflate"); - if client_no_context_takeover { - normalized.push_str("; client_no_context_takeover"); - } - if server_no_context_takeover { - normalized.push_str("; server_no_context_takeover"); - } - Ok(Some(normalized)) -} - -/// Check if the client request headers contain both `Upgrade` and -/// `Connection: Upgrade` headers, indicating the client requested a -/// protocol upgrade (e.g. WebSocket). -/// -/// Per RFC 9110 Section 7.8, a server MUST NOT send 101 Switching Protocols -/// unless the client sent these headers. -fn client_requested_upgrade(headers: &str) -> bool { - let mut has_upgrade_header = false; - let mut connection_contains_upgrade = false; - - for line in headers.lines().skip(1) { - let lower = line.to_ascii_lowercase(); - if lower.starts_with("upgrade:") { - has_upgrade_header = true; + ResponseFraming::Chunked + }; + let trailer_names: Vec = if supports_chunked_response { + finish + .trailers + .iter() + .map(|header| header.name.clone()) + .collect() + } else { + Vec::new() + }; + let head = serialize_response_head( + &status_line, + &headers, + framing, + server_wants_close, + &trailer_names, + ); + client.write_all(&head).await.into_diagnostic()?; + if matches!(framing, ResponseFraming::Chunked) { + for unit in &finish.body_units { + write_chunk(client, unit).await?; + } + write_response_trailers(client, &finish.trailers).await?; + } else { + for unit in &finish.body_units { + client.write_all(unit).await.into_diagnostic()?; + } } - if lower.starts_with("connection:") { - let val = lower.split_once(':').map_or("", |(_, v)| v.trim()); - // Connection header can have comma-separated values - if val.split(',').any(|tok| tok.trim() == "upgrade") { - connection_contains_upgrade = true; + } else { + for unit in &finish.body_units { + if chunked_output { + write_chunk(client, unit).await?; + } else { + client.write_all(unit).await.into_diagnostic()?; } } + if chunked_output { + write_response_trailers(client, &finish.trailers).await?; + } } - - has_upgrade_header && connection_contains_upgrade -} - -/// Returns true for responses that MUST NOT contain a message body per RFC 7230 §3.3.3: -/// HEAD responses, 1xx informational, 204 No Content, 304 Not Modified. -fn is_bodiless_response(request_method: &str, status_code: u16) -> bool { - request_method.eq_ignore_ascii_case("HEAD") - || (100..200).contains(&status_code) - || status_code == 204 - || status_code == 304 + client.flush().await.into_diagnostic()?; + Ok(Some( + if (committed && close_delimited_output) + || (matches!(body_length, BodyLength::None) && (server_wants_close || event_stream)) + { + RelayOutcome::Consumed + } else { + RelayOutcome::Reusable + }, + )) } -/// Relay all bytes from reader to writer until EOF or idle timeout. -/// -/// Used for HTTP responses with no explicit framing (no Content-Length, -/// no Transfer-Encoding) where the body is delimited by connection close. -/// An idle timeout prevents blocking when servers keep the TCP connection -/// alive longer than expected (e.g. CDN keep-alive timers). -async fn relay_until_eof(reader: &mut R, writer: &mut W) -> Result<()> +#[allow(clippy::too_many_arguments)] +async fn relay_headers_only_response( + request_method: &str, + upstream: &mut U, + client: &mut C, + status_line: &str, + headers: &[HttpHeader], + declared_trailers: &[String], + overflow: &[u8], + status_code: u16, + body_length: BodyLength, + server_wants_close: bool, + event_stream: bool, +) -> Result where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, + U: AsyncRead + Unpin, + C: AsyncWrite + Unpin, { - let mut buf = [0u8; RELAY_BUF_SIZE]; - loop { - match tokio::time::timeout(RELAY_EOF_IDLE_TIMEOUT, reader.read(&mut buf)).await { - Ok(Ok(0)) => return Ok(()), - Ok(Ok(n)) => { - writer.write_all(&buf[..n]).await.into_diagnostic()?; - writer.flush().await.into_diagnostic()?; + let head = serialize_response_head( + status_line, + headers, + ResponseFraming::Preserve(body_length), + server_wants_close, + declared_trailers, + ); + client.write_all(&head).await.into_diagnostic()?; + + if is_bodiless_response(request_method, status_code) { + client.flush().await.into_diagnostic()?; + return Ok(if server_wants_close { + RelayOutcome::Consumed + } else { + RelayOutcome::Reusable + }); + } + + client.write_all(overflow).await.into_diagnostic()?; + match body_length { + BodyLength::ContentLength(length) => { + let remaining = length.saturating_sub(overflow.len() as u64); + if remaining > 0 { + relay_fixed(upstream, client, remaining, None).await?; } - Ok(Err(e)) => return Err(miette::miette!("{e}")), - Err(_) => { - debug!( - "relay_until_eof idle timeout after {:?}", - RELAY_EOF_IDLE_TIMEOUT - ); - return Ok(()); + } + BodyLength::Chunked => relay_chunked(upstream, client, overflow, None).await?, + BodyLength::None if server_wants_close || event_stream => { + if event_stream { + relay_until_eof_without_idle_timeout(upstream, client).await?; + } else { + relay_until_eof(upstream, client).await?; } + client.flush().await.into_diagnostic()?; + return Ok(RelayOutcome::Consumed); } + BodyLength::None => {} } + client.flush().await.into_diagnostic()?; + Ok(RelayOutcome::Reusable) } -/// Relay all bytes from reader to writer until EOF without an idle timeout. -/// -/// Used for server-sent events, where long idle gaps are part of the protocol -/// and do not mean the response body is complete. -async fn relay_until_eof_without_idle_timeout(reader: &mut R, writer: &mut W) -> Result<()> -where - R: AsyncRead + Unpin, - W: AsyncWrite + Unpin, -{ - let mut buf = [0u8; RELAY_BUF_SIZE]; - loop { - let n = reader.read(&mut buf).await.into_diagnostic()?; - if n == 0 { - return Ok(()); +fn emit_http_response_middleware_invocations( + policy_name: &str, + target: &HttpRequestTarget, + status_code: u16, + invocations: &[openshell_supervisor_middleware::HttpResponseInvocation], +) { + for event in + http_response_middleware_invocation_events(policy_name, target, status_code, invocations) + { + openshell_ocsf::ocsf_emit!(event); + } + for invocation in invocations { + if let Some(event) = + http_response_middleware_fail_open_finding_event(policy_name, target, invocation) + { + openshell_ocsf::ocsf_emit!(event); + } + if let Some(event) = + http_response_middleware_block_finding_event(policy_name, target, invocation) + { + openshell_ocsf::ocsf_emit!(event); } - writer.write_all(&buf[..n]).await.into_diagnostic()?; - writer.flush().await.into_diagnostic()?; } } -/// Detect if the first bytes look like an HTTP request. -/// -/// Checks for common HTTP methods at the start of the stream. -pub fn looks_like_http(peek: &[u8]) -> bool { - HTTP_METHOD_PREFIXES +fn emit_http_response_diagnostics( + policy_name: &str, + target: &HttpRequestTarget, + status_code: u16, + diagnostics: &openshell_supervisor_middleware::HttpResponseDiagnostics, +) { + emit_http_response_middleware_invocations( + policy_name, + target, + status_code, + &diagnostics.invocations, + ); + for event in crate::l7::middleware::middleware_finding_events(&diagnostics.findings) { + openshell_ocsf::ocsf_emit!(event); + } +} + +fn http_response_middleware_invocation_events( + policy_name: &str, + target: &HttpRequestTarget, + status_code: u16, + invocations: &[openshell_supervisor_middleware::HttpResponseInvocation], +) -> Vec { + invocations .iter() - .any(|method| peek.starts_with(method)) + .map(|invocation| { + let outcome = format!("{:?}", invocation.outcome).to_ascii_lowercase(); + let failed = invocation.failed; + let blocked = invocation.outcome + == openshell_supervisor_middleware::HttpResponseInvocationOutcome::BlockDelivery; + openshell_ocsf::HttpActivityBuilder::new(ocsf_ctx()) + .activity(openshell_ocsf::ActivityId::Other) + .action(if blocked { + openshell_ocsf::ActionId::Denied + } else if failed { + openshell_ocsf::ActionId::Other + } else { + openshell_ocsf::ActionId::Allowed + }) + .disposition(if blocked { + openshell_ocsf::DispositionId::Blocked + } else if failed { + openshell_ocsf::DispositionId::Error + } else { + openshell_ocsf::DispositionId::Allowed + }) + .severity(if failed || blocked { + openshell_ocsf::SeverityId::Medium + } else { + openshell_ocsf::SeverityId::Informational + }) + .status(if failed || blocked { + openshell_ocsf::StatusId::Failure + } else { + openshell_ocsf::StatusId::Success + }) + .http_request(openshell_ocsf::HttpRequest::new( + &target.method, + openshell_ocsf::Url::new( + &target.scheme, + &target.host, + &target.path, + u16::try_from(target.port).unwrap_or_default(), + ), + )) + .http_response(openshell_ocsf::HttpResponse { code: status_code }) + .dst_endpoint(openshell_ocsf::Endpoint::from_domain( + &target.host, + u16::try_from(target.port).unwrap_or_default(), + )) + .firewall_rule(policy_name, "supervisor-middleware") + .unmapped("middleware_config", invocation.config_name.as_str()) + .unmapped( + "middleware_implementation", + invocation.implementation.as_str(), + ) + .unmapped("response_middleware_outcome", outcome.as_str()) + .unmapped("sequence", invocation.sequence.unwrap_or_default()) + .unmapped("input_bytes", invocation.input_size) + .unmapped("failed", failed) + .message(format!( + "HTTP_RESPONSE_MIDDLEWARE config={} implementation={} outcome={} sequence={} input_bytes={} failed={failed}", + invocation.config_name, + invocation.implementation, + outcome, + invocation.sequence.unwrap_or_default(), + invocation.input_size, + )) + .build() + }) + .collect() } -pub(crate) fn could_be_http_request_prefix(peek: &[u8]) -> bool { - !peek.is_empty() - && HTTP_METHOD_PREFIXES - .iter() - .any(|method| peek.len() < method.len() && method.starts_with(peek)) +fn http_response_middleware_block_finding_event( + policy_name: &str, + target: &HttpRequestTarget, + invocation: &openshell_supervisor_middleware::HttpResponseInvocation, +) -> Option { + if invocation.outcome + != openshell_supervisor_middleware::HttpResponseInvocationOutcome::BlockDelivery + { + return None; + } + Some( + openshell_ocsf::DetectionFindingBuilder::new(ocsf_ctx()) + .severity(openshell_ocsf::SeverityId::Medium) + .finding_info(openshell_ocsf::FindingInfo::new( + "openshell.middleware.http_response_blocked", + "HTTP response blocked by middleware", + )) + .evidence_pairs(&[ + ("policy", policy_name), + ("middleware_config", invocation.config_name.as_str()), + ( + "middleware_implementation", + invocation.implementation.as_str(), + ), + ("host", target.host.as_str()), + ("phase", "pre_return"), + ]) + .unmapped("middleware_config", invocation.config_name.as_str()) + .unmapped( + "middleware_implementation", + invocation.implementation.as_str(), + ) + .unmapped("phase", "pre_return") + .message("HTTP response delivery blocked by middleware") + .build(), + ) } -pub fn looks_like_http2_prior_knowledge(peek: &[u8]) -> bool { - peek.len() >= MIN_HTTP2_PREFACE_DETECTION_BYTES - && HTTP2_PRIOR_KNOWLEDGE_PREFACE.starts_with(peek) +fn http_response_middleware_fail_open_finding_event( + policy_name: &str, + target: &HttpRequestTarget, + invocation: &openshell_supervisor_middleware::HttpResponseInvocation, +) -> Option { + if !invocation.failed + || invocation.outcome + != openshell_supervisor_middleware::HttpResponseInvocationOutcome::FailOpen + { + return None; + } + let failure_category = invocation + .failure_category + .as_deref() + .unwrap_or("middleware_failure"); + Some( + openshell_ocsf::DetectionFindingBuilder::new(ocsf_ctx()) + .severity(openshell_ocsf::SeverityId::Medium) + .finding_info(openshell_ocsf::FindingInfo::new( + "openshell.middleware.http_response_fail_open", + "HTTP response middleware failed open", + )) + .evidence_pairs(&[ + ("policy", policy_name), + ("middleware_config", invocation.config_name.as_str()), + ( + "middleware_implementation", + invocation.implementation.as_str(), + ), + ("host", target.host.as_str()), + ("phase", "pre_return"), + ("failure_category", failure_category), + ]) + .unmapped("middleware_config", invocation.config_name.as_str()) + .unmapped( + "middleware_implementation", + invocation.implementation.as_str(), + ) + .unmapped("phase", "pre_return") + .unmapped("failure_category", failure_category) + .message("HTTP response middleware failed and response inspection was bypassed") + .build(), + ) } -pub(crate) fn could_be_http2_prior_knowledge_prefix(peek: &[u8]) -> bool { - !peek.is_empty() - && peek.len() < MIN_HTTP2_PREFACE_DETECTION_BYTES - && HTTP2_PRIOR_KNOWLEDGE_PREFACE.starts_with(peek) +fn emit_http_response_middleware_failure( + policy_name: &str, + target: &HttpRequestTarget, + status_code: u16, + committed: bool, +) { + let status_code = status_code.to_string(); + let event = openshell_ocsf::DetectionFindingBuilder::new(ocsf_ctx()) + .severity(openshell_ocsf::SeverityId::High) + .finding_info(openshell_ocsf::FindingInfo::new( + "openshell.middleware.http_response_failure", + "HTTP response middleware delivery failure", + )) + .evidence_pairs(&[ + ("policy", policy_name), + ("host", target.host.as_str()), + ( + "commitment", + if committed { + "after_commit" + } else { + "before_commit" + }, + ), + ("upstream_status", status_code.as_str()), + ]) + .message(if committed { + "HTTP response middleware failed after response commitment" + } else { + "HTTP response middleware failed before response commitment" + }) + .build(); + openshell_ocsf::ocsf_emit!(event); } -/// Check if an IO error represents a benign connection close. -/// -/// TLS peers commonly close the socket without sending a `close_notify` alert. -/// Rustls reports this as `UnexpectedEof`, but it's functionally equivalent -/// to a clean close when no request data has been received yet. -fn is_benign_close(err: &std::io::Error) -> bool { - matches!( - err.kind(), - std::io::ErrorKind::UnexpectedEof - | std::io::ErrorKind::ConnectionReset - | std::io::ErrorKind::BrokenPipe - ) +#[derive(Debug)] +struct ParsedResponseHead { + headers: Vec, + connection_nominated: Vec, + declared_trailers: Vec, } -#[cfg(test)] -#[allow( - clippy::iter_on_single_items, - clippy::manual_string_new, - clippy::collapsible_if, - clippy::cast_possible_truncation, - reason = "Test code: test fixtures and explicit value-shape assertions are idiomatic in tests." -)] -mod tests { - use super::*; - use crate::opa::OpaEngine; - use flate2::{Compress, Compression, Decompress, FlushCompress, FlushDecompress, Status}; - use openshell_core::proposals::AgentProposals; - use openshell_core::secrets::SecretResolver; - use std::pin::Pin; - use std::sync::Arc; - use std::task::{Context, Poll}; - use tokio::io::ReadBuf; - - const TEST_POLICY: &str = include_str!("../../data/sandbox-policy.rego"); - const VALID_WS_KEY: &str = "dGhlIHNhbXBsZSBub25jZQ=="; - const VALID_WS_ACCEPT: &str = "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="; - const TEXT_OPCODE: u8 = 0x1; - - struct CountingReader { - bytes: Vec, - position: usize, - reads: usize, +fn parse_response_head_for_middleware(header_bytes: &[u8]) -> Result { + let header = std::str::from_utf8(header_bytes) + .map_err(|_| miette!("HTTP response headers contain invalid UTF-8"))?; + if parse_status_code(header).is_none() { + return Err(miette!("HTTP response status line is malformed")); } - - impl CountingReader { - fn new(bytes: Vec) -> Self { - Self { - bytes, - position: 0, - reads: 0, + let mut nominated = HashSet::new(); + let mut declared_trailers = Vec::new(); + for line in header.split("\r\n").skip(1) { + let Some((name, value)) = line.split_once(':') else { + continue; + }; + validate_http_field_name(name)?; + validate_http_field_value(value.trim())?; + if name.eq_ignore_ascii_case("connection") { + for token in value + .split(',') + .map(str::trim) + .filter(|token| !token.is_empty()) + { + nominated.insert(token.to_ascii_lowercase()); + } + } else if name.eq_ignore_ascii_case("trailer") { + for token in parse_http_token_list(value)? { + let token = token.to_ascii_lowercase(); + if !declared_trailers.contains(&token) { + declared_trailers.push(token); + } } } } - - impl AsyncRead for CountingReader { - fn poll_read( - mut self: Pin<&mut Self>, - _context: &mut Context<'_>, - buffer: &mut ReadBuf<'_>, - ) -> Poll> { - self.reads += 1; - let available = self.bytes.len().saturating_sub(self.position); - let amount = available.min(buffer.remaining()); - let end = self.position + amount; - buffer.put_slice(&self.bytes[self.position..end]); - self.position = end; - Poll::Ready(Ok(())) - } - } - - fn write_header(name: &str, value: &str, on_existing: ExistingHeaderAction) -> HeaderMutation { - HeaderMutation { - operation: Some(header_mutation::Operation::Write( - openshell_core::proto::WriteHeader { - name: name.into(), - value: value.into(), - on_existing: on_existing as i32, - }, - )), + for trailer in &declared_trailers { + if is_protected_response_field(trailer) || nominated.contains(trailer) { + return Err(miette!("HTTP response declares a protected trailer field")); } } - - fn remove_header(name: &str) -> HeaderMutation { - HeaderMutation { - operation: Some(header_mutation::Operation::Remove( - openshell_core::proto::RemoveHeader { name: name.into() }, - )), + let mut headers = Vec::new(); + for line in header.split("\r\n").skip(1).filter(|line| !line.is_empty()) { + let (name, value) = line + .split_once(':') + .ok_or_else(|| miette!("Malformed HTTP response header field"))?; + validate_http_field_name(name)?; + validate_http_field_value(value.trim())?; + let name = name.to_ascii_lowercase(); + if nominated.contains(&name) || is_protected_response_field(&name) { + continue; } + headers.push(HttpHeader { + name, + value: value.trim().to_string(), + }); } + let mut connection_nominated: Vec<_> = nominated.into_iter().collect(); + connection_nominated.sort(); + Ok(ParsedResponseHead { + headers, + connection_nominated, + declared_trailers, + }) +} - #[test] - fn ordered_header_mutations_replay_against_raw_request() { - let raw = b"GET / HTTP/1.1\r\nHost: example.test\r\nX-OpenShell-Middleware-Chain: first\r\nX-Drop: one\r\nX-Drop: two\r\n\r\n"; - let mutations = [ - write_header( - "x-openshell-middleware-chain", - "second", - ExistingHeaderAction::Append, - ), - write_header( - "x-openshell-middleware-chain", - "ignored", - ExistingHeaderAction::Skip, - ), - write_header( - "x-openshell-middleware-chain", - "replacement", - ExistingHeaderAction::Overwrite, - ), - write_header( - "x-openshell-middleware-chain", - "tail", - ExistingHeaderAction::Append, - ), - remove_header("x-drop"), - ]; - - let updated = String::from_utf8( - apply_header_mutations(raw, &mutations).expect("apply ordered header mutations"), - ) - .expect("UTF-8 request"); - let values: Vec<&str> = updated - .lines() - .filter_map(|line| { - line.split_once(':').and_then(|(name, value)| { - name.eq_ignore_ascii_case("x-openshell-middleware-chain") - .then_some(value.trim()) - }) - }) - .collect(); - assert_eq!(values, vec!["replacement", "tail"]); - assert!(!updated.to_ascii_lowercase().contains("x-drop:")); - assert!(updated.contains("Host: example.test")); +fn validate_http_field_name(name: &str) -> Result<()> { + if name.is_empty() + || !name.bytes().all(|byte| { + byte.is_ascii_alphanumeric() + || matches!( + byte, + b'!' | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + ) + }) + { + return Err(miette!("HTTP response field name is malformed")); } + Ok(()) +} - #[derive(Debug)] - struct CapturedFrame { - fin_opcode: u8, - masked: bool, - payload: Vec, +fn validate_http_field_value(value: &str) -> Result<()> { + if value + .bytes() + .any(|byte| (byte < 0x20 && byte != b'\t') || byte == 0x7f) + { + return Err(miette!("HTTP response field value contains a control byte")); } + Ok(()) +} - async fn read_http_header_block(reader: &mut R) -> Vec { - tokio::time::timeout(std::time::Duration::from_secs(2), async { - let mut header = Vec::new(); - let mut byte = [0u8; 1]; - loop { - reader.read_exact(&mut byte).await.unwrap(); - header.push(byte[0]); - if header.ends_with(b"\r\n\r\n") { - break; - } - } - header - }) - .await - .expect("HTTP header block should arrive") - } +fn is_protected_response_field(name: &str) -> bool { + matches!( + name.to_ascii_lowercase().as_str(), + "connection" + | "content-length" + | "keep-alive" + | "proxy-authenticate" + | "proxy-authorization" + | "proxy-connection" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) +} - async fn read_websocket_frame(reader: &mut R) -> CapturedFrame { - tokio::time::timeout(std::time::Duration::from_secs(2), async { - let mut prefix = [0u8; 2]; - reader.read_exact(&mut prefix).await.unwrap(); - let masked = prefix[1] & 0x80 != 0; - let mut payload_len = u64::from(prefix[1] & 0x7f); - if payload_len == 126 { - let mut extended = [0u8; 2]; - reader.read_exact(&mut extended).await.unwrap(); - payload_len = u64::from(u16::from_be_bytes(extended)); - } else if payload_len == 127 { - let mut extended = [0u8; 8]; - reader.read_exact(&mut extended).await.unwrap(); - payload_len = u64::from_be_bytes(extended); - } - let mut mask_key = [0u8; 4]; - if masked { - reader.read_exact(&mut mask_key).await.unwrap(); - } - let payload_len = usize::try_from(payload_len).unwrap(); - let mut payload = vec![0u8; payload_len]; - reader.read_exact(&mut payload).await.unwrap(); - if masked { - apply_test_mask(&mut payload, mask_key); - } - CapturedFrame { - fin_opcode: prefix[0], - masked, - payload, - } - }) - .await - .expect("WebSocket frame should arrive") - } +fn response_status_line(header_bytes: &[u8]) -> Result { + let line_end = header_bytes + .windows(2) + .position(|window| window == b"\r\n") + .ok_or_else(|| miette!("HTTP response status line is incomplete"))?; + std::str::from_utf8(&header_bytes[..line_end]) + .map(str::to_string) + .map_err(|_| miette!("HTTP response status line contains invalid UTF-8")) +} - async fn policy_local_json_response( - ctx: Arc, - ) -> serde_json::Value { - let (mut client, mut server) = tokio::io::duplex(4096); - let task = tokio::spawn(async move { - crate::policy_local::handle_forward_request( - ctx.as_ref(), - "GET", - "/v1/policy/current", - b"GET http://policy.local/v1/policy/current HTTP/1.1\r\nHost: policy.local\r\n\r\n", - &mut server, - ) - .await - .unwrap(); - }); +#[derive(Clone, Copy)] +enum ResponseFraming { + Preserve(BodyLength), + ContentLength(u64), + Chunked, +} - let mut received = Vec::new(); - client.read_to_end(&mut received).await.unwrap(); - task.await.unwrap(); +fn serialize_response_head( + status_line: &str, + headers: &[HttpHeader], + framing: ResponseFraming, + connection_close: bool, + trailer_names: &[String], +) -> Vec { + let mut output = format!("{status_line}\r\n"); + for header in headers { + output.push_str(&header.name); + output.push_str(": "); + output.push_str(&header.value); + output.push_str("\r\n"); + } + match framing { + ResponseFraming::Preserve(BodyLength::ContentLength(length)) + | ResponseFraming::ContentLength(length) => { + write!(&mut output, "Content-Length: {length}\r\n") + .expect("writing to a String cannot fail"); + } + ResponseFraming::Preserve(BodyLength::Chunked) | ResponseFraming::Chunked => { + output.push_str("Transfer-Encoding: chunked\r\n"); + } + ResponseFraming::Preserve(BodyLength::None) => {} + } + if !trailer_names.is_empty() { + output.push_str("Trailer: "); + output.push_str(&trailer_names.join(", ")); + output.push_str("\r\n"); + } + if connection_close { + output.push_str("Connection: close\r\n"); + } + output.push_str("\r\n"); + output.into_bytes() +} - let response = String::from_utf8(received).unwrap(); - let (_, body) = response.split_once("\r\n\r\n").unwrap(); - serde_json::from_str(body).unwrap() - } +fn strip_response_integrity_headers(headers: &mut Vec) { + headers.retain(|header| { + !matches!( + header.name.to_ascii_lowercase().as_str(), + "accept-ranges" + | "etag" + | "content-md5" + | "digest" + | "content-digest" + | "repr-digest" + | "signature" + | "signature-input" + ) + }); +} - fn masked_frame_with_rsv(opcode: u8, rsv: u8, payload: &[u8]) -> Vec { - let mask_key = [0x37, 0xfa, 0x21, 0x3d]; - let mut frame = Vec::new(); - frame.push(0x80 | rsv | opcode); - write_test_payload_len(&mut frame, 0x80, payload.len()); - frame.extend_from_slice(&mask_key); - let mut masked = payload.to_vec(); - apply_test_mask(&mut masked, mask_key); - frame.extend_from_slice(&masked); - frame - } +struct BufferedResponseReader<'a, R> { + upstream: &'a mut R, + buffered: &'a [u8], + position: usize, + exact_buffer: Vec, + exact_target: Option, + line_buffer: Vec, +} - fn unmasked_frame(opcode: u8, payload: &[u8]) -> Vec { - let mut frame = Vec::new(); - frame.push(0x80 | opcode); - write_test_payload_len(&mut frame, 0, payload.len()); - frame.extend_from_slice(payload); - frame +impl<'a, R: AsyncRead + Unpin> BufferedResponseReader<'a, R> { + fn new(upstream: &'a mut R, buffered: &'a [u8]) -> Self { + Self { + upstream, + buffered, + position: 0, + exact_buffer: Vec::new(), + exact_target: None, + line_buffer: Vec::new(), + } } - fn write_test_payload_len(frame: &mut Vec, mask_bit: u8, payload_len: usize) { - if payload_len < 126 { - frame.push(mask_bit | payload_len as u8); - } else if u16::try_from(payload_len).is_ok() { - frame.push(mask_bit | 0x7e); - frame.extend_from_slice(&(payload_len as u16).to_be_bytes()); - } else { - frame.push(mask_bit | 0x7f); - frame.extend_from_slice(&(payload_len as u64).to_be_bytes()); + async fn read_some(&mut self, limit: usize) -> Result>> { + if self.position < self.buffered.len() { + let end = self.position.saturating_add(limit).min(self.buffered.len()); + let data = self.buffered[self.position..end].to_vec(); + self.position = end; + return Ok(Some(data)); } + let mut data = vec![0u8; limit.max(1)]; + let count = self.upstream.read(&mut data).await.into_diagnostic()?; + if count == 0 { + return Ok(None); + } + data.truncate(count); + Ok(Some(data)) } - fn apply_test_mask(payload: &mut [u8], mask_key: [u8; 4]) { - for (index, byte) in payload.iter_mut().enumerate() { - *byte ^= mask_key[index % 4]; + async fn read_exact_vec(&mut self, length: usize) -> Result> { + match self.exact_target { + Some(target) if target != length => { + return Err(miette!("HTTP response reader exact-read state mismatch")); + } + None => { + self.exact_target = Some(length); + self.exact_buffer.reserve(length); + } + Some(_) => {} + } + while self.exact_buffer.len() < length { + let remaining = length - self.exact_buffer.len(); + let Some(data) = self.read_some(remaining).await? else { + return Err(miette!("HTTP response body ended unexpectedly")); + }; + self.exact_buffer.extend_from_slice(&data); } + self.exact_target = None; + Ok(std::mem::take(&mut self.exact_buffer)) } - fn compress_test_permessage_deflate(payload: &[u8]) -> Vec { - let mut compressor = Compress::new(Compression::fast(), false); - let mut out = Vec::with_capacity(payload.len().saturating_add(128)); + async fn read_line(&mut self) -> Result> { loop { - let consumed = usize::try_from(compressor.total_in()).unwrap(); - if consumed >= payload.len() { - break; + let Some(byte) = self.read_some(1).await? else { + return Err(miette!("HTTP response ended before line terminator")); + }; + self.line_buffer.push(byte[0]); + if self.line_buffer.len() > MAX_HEADER_BYTES { + return Err(miette!("HTTP response line exceeds limit")); } - let before_in = compressor.total_in(); - let before_out = compressor.total_out(); - let status = compressor - .compress_vec(&payload[consumed..], &mut out, FlushCompress::None) - .unwrap(); - if matches!(status, Status::BufError) - || (compressor.total_in() == before_in && compressor.total_out() == before_out) - { - out.reserve(out.capacity().max(1024)); + if self.line_buffer.ends_with(b"\r\n") { + self.line_buffer.truncate(self.line_buffer.len() - 2); + return Ok(std::mem::take(&mut self.line_buffer)); + } + } + } +} + +#[allow(clippy::too_many_arguments)] +async fn relay_normalized_response_body( + reader: &mut BufferedResponseReader<'_, R>, + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + body_length: BodyLength, + server_wants_close: bool, + event_stream: bool, + committed: &mut bool, + chunked_output: bool, + commit_head: &[u8], + unit_limit: usize, + generation_guard: Option<&PolicyGenerationGuard>, + policy_name: &str, + target: &HttpRequestTarget, + status_code: u16, + connection_nominated_headers: &[String], +) -> Result> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + let mut pending = Vec::with_capacity(unit_limit); + let mut framing = ResponseOutputState { + committed, + chunked: chunked_output, + commit_head, + policy_name, + target, + status_code, + }; + match body_length { + BodyLength::ContentLength(mut remaining) => { + while remaining > 0 { + let length = usize::try_from(remaining) + .unwrap_or(unit_limit) + .min(unit_limit); + let unit = read_exact_response_with_deadline( + reader, + length, + session, + client, + &mut framing, + ) + .await?; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + remaining -= unit.len() as u64; + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + &mut framing, + unit_limit, + ) + .await?; + } + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; + Ok(Vec::new()) + } + BodyLength::Chunked => { + let mut size_line = + read_response_line_with_deadline(reader, session, client, &mut framing).await?; + loop { + let size_line_text = std::str::from_utf8(&size_line) + .map_err(|_| miette!("Invalid UTF-8 in response chunk-size line"))?; + let size_token = size_line_text + .split(';') + .next() + .map(str::trim) + .unwrap_or_default(); + let chunk_size = usize::from_str_radix(size_token, 16) + .map_err(|_| miette!("Invalid HTTP response chunk size"))?; + if chunk_size == 0 { + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; + return read_response_trailers( + reader, + session, + client, + &mut framing, + connection_nominated_headers, + ) + .await; + } + let mut remaining = chunk_size; + while remaining > 0 { + let length = remaining.min(unit_limit); + let unit = read_exact_response_with_deadline( + reader, + length, + session, + client, + &mut framing, + ) + .await?; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + remaining -= unit.len(); + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + &mut framing, + unit_limit, + ) + .await?; + } + if read_exact_response_with_deadline(reader, 2, session, client, &mut framing) + .await? + .as_slice() + != b"\r\n" + { + return Err(miette!("HTTP response chunk is missing its terminator")); + } + size_line = if let Ok(line) = tokio::time::timeout( + RESPONSE_UNIT_COALESCE_TIMEOUT, + read_response_line_with_deadline(reader, session, client, &mut framing), + ) + .await + { + line? + } else { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + &mut framing, + ) + .await?; + read_response_line_with_deadline(reader, session, client, &mut framing).await? + }; + } + } + BodyLength::None if server_wants_close || event_stream => loop { + let read = + read_response_with_deadline(reader, unit_limit, session, client, &mut framing); + let next = if pending.is_empty() && event_stream { + read.await? + } else if pending.is_empty() { + match tokio::time::timeout(RELAY_EOF_IDLE_TIMEOUT, read).await { + Ok(result) => result?, + Err(_) => None, + } + } else if let Ok(result) = + tokio::time::timeout(RESPONSE_UNIT_COALESCE_TIMEOUT, read).await + { + result? + } else { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + &mut framing, + ) + .await?; + continue; + }; + let Some(unit) = next else { + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; + return Ok(Vec::new()); + }; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + &mut framing, + unit_limit, + ) + .await?; + }, + BodyLength::None => { + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; + Ok(Vec::new()) + } + } +} + +struct ResponseOutputState<'a> { + committed: &'a mut bool, + chunked: bool, + commit_head: &'a [u8], + policy_name: &'a str, + target: &'a HttpRequestTarget, + status_code: u16, +} + +#[derive(Debug)] +struct ResponseMiddlewareStop { + failure: openshell_supervisor_middleware::HttpResponseMiddlewareFailure, +} + +impl ResponseMiddlewareStop { + fn new(failure: openshell_supervisor_middleware::HttpResponseMiddlewareFailure) -> Self { + Self { failure } + } +} + +impl fmt::Display for ResponseMiddlewareStop { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + formatter, + "HTTP response middleware stopped delivery: {}", + self.failure + ) + } +} + +impl std::error::Error for ResponseMiddlewareStop {} + +impl miette::Diagnostic for ResponseMiddlewareStop {} + +async fn expire_whole_body_deadline( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + framing: &mut ResponseOutputState<'_>, +) -> Result<()> { + let output = session.expire_whole_body_deadline().await; + let diagnostics = session.take_diagnostics(); + emit_http_response_diagnostics( + framing.policy_name, + framing.target, + framing.status_code, + &diagnostics, + ); + let output = output.map_err(ResponseMiddlewareStop::new)?; + deliver_response_units(client, output, framing, session.requires_whole_body()).await +} + +async fn read_response_with_deadline( + reader: &mut BufferedResponseReader<'_, R>, + limit: usize, + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + framing: &mut ResponseOutputState<'_>, +) -> Result>> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + loop { + let Some(deadline) = session.whole_body_deadline() else { + return reader.read_some(limit).await; + }; + match tokio::time::timeout_at(deadline, reader.read_some(limit)).await { + Ok(result) => return result, + Err(_) => expire_whole_body_deadline(session, client, framing).await?, + } + } +} + +async fn read_exact_response_with_deadline( + reader: &mut BufferedResponseReader<'_, R>, + length: usize, + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + framing: &mut ResponseOutputState<'_>, +) -> Result> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + loop { + let Some(deadline) = session.whole_body_deadline() else { + return reader.read_exact_vec(length).await; + }; + match tokio::time::timeout_at(deadline, reader.read_exact_vec(length)).await { + Ok(result) => return result, + Err(_) => expire_whole_body_deadline(session, client, framing).await?, + } + } +} + +async fn read_response_line_with_deadline( + reader: &mut BufferedResponseReader<'_, R>, + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + framing: &mut ResponseOutputState<'_>, +) -> Result> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + loop { + let Some(deadline) = session.whole_body_deadline() else { + return reader.read_line().await; + }; + match tokio::time::timeout_at(deadline, reader.read_line()).await { + Ok(result) => return result, + Err(_) => expire_whole_body_deadline(session, client, framing).await?, + } + } +} + +async fn buffer_normalized_response_bytes( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + pending: &mut Vec, + data: Vec, + framing: &mut ResponseOutputState<'_>, + unit_limit: usize, +) -> Result<()> { + pending.extend_from_slice(&data); + while pending.len() >= unit_limit { + let remainder = pending.split_off(unit_limit); + let unit = std::mem::replace(pending, remainder); + process_response_unit(session, client, unit, framing).await?; + } + Ok(()) +} + +async fn flush_normalized_response_bytes( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + pending: Vec, + framing: &mut ResponseOutputState<'_>, +) -> Result<()> { + if pending.is_empty() { + return Ok(()); + } + process_response_unit(session, client, pending, framing).await +} + +async fn process_response_unit( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + unit: Vec, + framing: &mut ResponseOutputState<'_>, +) -> Result<()> { + let output = session.push_body(unit).await; + let diagnostics = session.take_diagnostics(); + emit_http_response_diagnostics( + framing.policy_name, + framing.target, + framing.status_code, + &diagnostics, + ); + let output = output.map_err(ResponseMiddlewareStop::new)?; + deliver_response_units(client, output, framing, session.requires_whole_body()).await +} + +async fn deliver_response_units( + client: &mut C, + output: Vec>, + framing: &mut ResponseOutputState<'_>, + whole_body_pending: bool, +) -> Result<()> { + if !*framing.committed && !output.is_empty() { + if whole_body_pending { + return Err(miette!( + "whole-body response middleware released output before finalization" + )); + } + // `write_all` may return an error after a partial write. Treat the + // response as committed before the attempt so callers never append a + // canonical error response behind a partially delivered upstream head. + *framing.committed = true; + client + .write_all(framing.commit_head) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + client + .flush() + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + } + if *framing.committed { + for unit in output { + if framing.chunked { + write_chunk(client, &unit) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + } else { + client + .write_all(&unit) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + } + } + client + .flush() + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + } + Ok(()) +} + +async fn read_response_trailers( + reader: &mut BufferedResponseReader<'_, R>, + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + framing: &mut ResponseOutputState<'_>, + connection_nominated_headers: &[String], +) -> Result> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + let mut trailers = Vec::new(); + loop { + let line = read_response_line_with_deadline(reader, session, client, framing).await?; + if line.is_empty() { + return Ok(trailers); + } + let line = std::str::from_utf8(&line) + .map_err(|_| miette!("HTTP response trailer contains invalid UTF-8"))?; + let (name, value) = line + .split_once(':') + .ok_or_else(|| miette!("Malformed HTTP response trailer"))?; + validate_http_field_name(name)?; + validate_http_field_value(value.trim())?; + let name = name.to_ascii_lowercase(); + if is_protected_response_field(&name) || connection_nominated_headers.contains(&name) { + return Err(miette!("HTTP response trailer uses a protected field name")); + } + trailers.push(HttpHeader { + name, + value: value.trim().to_string(), + }); + if trailers.len() > openshell_supervisor_middleware::MAX_MIDDLEWARE_HEADERS { + return Err(miette!("HTTP response trailer count exceeds limit")); + } + } +} + +async fn write_response_trailers( + client: &mut C, + trailers: &[HttpHeader], +) -> Result<()> { + client.write_all(b"0\r\n").await.into_diagnostic()?; + for trailer in trailers { + client + .write_all(format!("{}: {}\r\n", trailer.name, trailer.value).as_bytes()) + .await + .into_diagnostic()?; + } + client.write_all(b"\r\n").await.into_diagnostic()?; + Ok(()) +} + +async fn send_response_middleware_denial( + client: &mut C, + request_method: &str, + policy_name: &str, + target: &HttpRequestTarget, + denial: &openshell_supervisor_middleware::MiddlewareDenial, +) -> Result<()> { + let mut body = serde_json::Map::new(); + body.insert("error".into(), serde_json::json!("middleware_denied")); + body.insert( + "detail".into(), + serde_json::json!("Response blocked by configured middleware"), + ); + body.insert("policy".into(), serde_json::json!(policy_name)); + body.insert("middleware".into(), serde_json::json!(denial.config_name)); + if let Some(reason_code) = &denial.reason_code { + body.insert("reason_code".into(), serde_json::json!(reason_code)); + } + body.insert( + "layer".into(), + serde_json::json!("http_response_pre_return"), + ); + body.insert("method".into(), serde_json::json!(target.method)); + body.insert("path".into(), serde_json::json!(target.path)); + body.insert("host".into(), serde_json::json!(target.host)); + body.insert("port".into(), serde_json::json!(target.port)); + let body = serde_json::to_vec(&serde_json::Value::Object(body)).into_diagnostic()?; + let head = format!( + "HTTP/1.1 403 Forbidden\r\nContent-Type: application/json\r\nContent-Length: {}\r\nX-OpenShell-Policy: {policy_name}\r\nConnection: close\r\n\r\n", + body.len() + ); + client.write_all(head.as_bytes()).await.into_diagnostic()?; + if !request_method.eq_ignore_ascii_case("HEAD") { + client.write_all(&body).await.into_diagnostic()?; + } + client.flush().await.into_diagnostic()?; + Ok(()) +} + +async fn send_response_delivery_failure( + client: &mut C, + request_method: &str, + policy_name: &str, + target: &HttpRequestTarget, +) -> Result<()> { + let body = serde_json::to_vec(&serde_json::json!({ + "error": "response_delivery_failed", + "detail": "The upstream request may have completed, but OpenShell could not deliver its response. Retrying may repeat upstream side effects.", + "policy": policy_name, + "layer": "http_response_pre_return", + "method": target.method, + "path": target.path, + "host": target.host, + "port": target.port, + })) + .into_diagnostic()?; + let head = format!( + "HTTP/1.1 502 Bad Gateway\r\nContent-Type: application/json\r\nContent-Length: {}\r\nX-OpenShell-Policy: {policy_name}\r\nConnection: close\r\n\r\n", + body.len() + ); + client.write_all(head.as_bytes()).await.into_diagnostic()?; + if !request_method.eq_ignore_ascii_case("HEAD") { + client.write_all(&body).await.into_diagnostic()?; + } + client.flush().await.into_diagnostic()?; + Ok(()) +} + +/// Parse the HTTP status code from a response status line. +/// +/// Expects the first line to look like `HTTP/1.1 200 OK`. +fn parse_status_code(headers: &str) -> Option { + let status_line = headers.lines().next()?; + let code_str = status_line.split_whitespace().nth(1)?; + code_str.parse().ok() +} + +/// Check if the response headers contain `Connection: close`. +fn parse_connection_close(headers: &str) -> bool { + for line in headers.lines().skip(1) { + let lower = line.to_ascii_lowercase(); + if lower.starts_with("connection:") { + let val = lower.split_once(':').map_or("", |(_, v)| v.trim()); + return val.contains("close"); + } + } + false +} + +fn response_is_event_stream(headers: &str) -> bool { + headers.lines().skip(1).any(|line| { + let lower = line.to_ascii_lowercase(); + let Some(value) = lower.strip_prefix("content-type:") else { + return false; + }; + value + .split(';') + .next() + .is_some_and(|mime| mime.trim() == "text/event-stream") + }) +} + +fn validate_websocket_response( + headers: &str, + mode: WebSocketExtensionMode, + websocket: Option<&WebSocketResponseValidation>, +) -> Result<(bool, Option)> { + let Some(validation) = websocket else { + return validate_websocket_response_extensions_preserved(headers, mode) + .map(|compressed| (compressed, None)); + }; + + let mut upgrade_websocket = false; + let mut connection_upgrade = false; + let mut accept_count = 0usize; + let mut accept_matches = false; + let mut subprotocol_count = 0usize; + let mut selected_subprotocol = None; + + for line in headers.lines().skip(1) { + let Some((name, value)) = line.split_once(':') else { + continue; + }; + let name = name.trim().to_ascii_lowercase(); + let value = value.trim(); + match name.as_str() { + "upgrade" if header_value_contains_token(value, "websocket") => { + upgrade_websocket = true; + } + "connection" if header_value_contains_token(value, "upgrade") => { + connection_upgrade = true; + } + "sec-websocket-accept" => { + accept_count += 1; + accept_matches = value == validation.expected_accept; + } + "sec-websocket-protocol" => { + subprotocol_count += 1; + if !is_http_token(value) { + return Err(miette!( + "websocket upgrade response has invalid Sec-WebSocket-Protocol" + )); + } + selected_subprotocol = Some(value.to_string()); + } + _ => {} + } + } + + if !upgrade_websocket { + return Err(miette!( + "websocket upgrade response missing Upgrade: websocket" + )); + } + if !connection_upgrade { + return Err(miette!( + "websocket upgrade response missing Connection: Upgrade" + )); + } + if accept_count != 1 || !accept_matches { + return Err(miette!( + "websocket upgrade response has invalid Sec-WebSocket-Accept" + )); + } + if subprotocol_count > 1 { + return Err(miette!( + "websocket upgrade response has multiple Sec-WebSocket-Protocol headers" + )); + } + if let Some(ref protocol) = selected_subprotocol + && !validation + .offered_subprotocols + .iter() + .any(|offered| offered == protocol) + { + return Err(miette!( + "upstream selected WebSocket subprotocol that was not offered" + )); + } + + let actual_extension = normalized_websocket_extension(headers)?; + match (&validation.expected_extension, actual_extension.as_deref()) { + (None, Some(_)) => Err(miette!( + "upstream negotiated WebSocket extension that was not offered" + )), + (None | Some(_), None) => Ok((false, selected_subprotocol)), + (Some(expected), Some(actual)) if expected.eq_ignore_ascii_case(actual) => { + Ok((true, selected_subprotocol)) + } + (Some(_), Some(_)) => Err(miette!( + "upstream negotiated WebSocket extension that does not match the safe offer" + )), + } +} + +fn validate_websocket_response_extensions_preserved( + headers: &str, + mode: WebSocketExtensionMode, +) -> Result { + match mode { + WebSocketExtensionMode::Preserve => Ok(false), + WebSocketExtensionMode::PermessageDeflate => { + let offers = websocket_extension_offers(headers)?; + if offers.is_empty() { + Ok(false) + } else { + Err(miette!( + "upstream negotiated WebSocket extension that was not offered" + )) + } + } + } +} + +fn normalized_websocket_extension(headers: &str) -> Result> { + let offers = websocket_extension_offers(headers)?; + if offers.is_empty() { + return Ok(None); + } + if offers.len() != 1 { + return Err(miette!("upstream negotiated multiple WebSocket extensions")); + } + let offer = &offers[0]; + if !offer.name.eq_ignore_ascii_case("permessage-deflate") { + return Err(miette!( + "upstream negotiated unsupported WebSocket extension" + )); + } + let mut client_no_context_takeover = false; + let mut server_no_context_takeover = false; + let mut seen = HashSet::new(); + for param in &offer.params { + let name = param.name.to_ascii_lowercase(); + if param.value.is_some() || !seen.insert(name.clone()) { + return Err(miette!( + "upstream negotiated unsupported permessage-deflate parameter" + )); + } + if name == "client_no_context_takeover" { + client_no_context_takeover = true; + } else if name == "server_no_context_takeover" { + server_no_context_takeover = true; + } else { + return Err(miette!( + "upstream negotiated unsupported permessage-deflate parameter" + )); + } + } + let mut normalized = String::from("permessage-deflate"); + if client_no_context_takeover { + normalized.push_str("; client_no_context_takeover"); + } + if server_no_context_takeover { + normalized.push_str("; server_no_context_takeover"); + } + Ok(Some(normalized)) +} + +/// Check if the client request headers contain both `Upgrade` and +/// `Connection: Upgrade` headers, indicating the client requested a +/// protocol upgrade (e.g. WebSocket). +/// +/// Per RFC 9110 Section 7.8, a server MUST NOT send 101 Switching Protocols +/// unless the client sent these headers. +fn client_requested_upgrade(headers: &str) -> bool { + let mut has_upgrade_header = false; + let mut connection_contains_upgrade = false; + + for line in headers.lines().skip(1) { + let lower = line.to_ascii_lowercase(); + if lower.starts_with("upgrade:") { + has_upgrade_header = true; + } + if lower.starts_with("connection:") { + let val = lower.split_once(':').map_or("", |(_, v)| v.trim()); + // Connection header can have comma-separated values + if val.split(',').any(|tok| tok.trim() == "upgrade") { + connection_contains_upgrade = true; + } + } + } + + has_upgrade_header && connection_contains_upgrade +} + +/// Returns true for responses that MUST NOT contain a message body per RFC 7230 §3.3.3: +/// HEAD responses, 1xx informational, 204 No Content, 304 Not Modified. +fn is_bodiless_response(request_method: &str, status_code: u16) -> bool { + request_method.eq_ignore_ascii_case("HEAD") + || (100..200).contains(&status_code) + || status_code == 204 + || status_code == 304 +} + +/// Relay all bytes from reader to writer until EOF or idle timeout. +/// +/// Used for HTTP responses with no explicit framing (no Content-Length, +/// no Transfer-Encoding) where the body is delimited by connection close. +/// An idle timeout prevents blocking when servers keep the TCP connection +/// alive longer than expected (e.g. CDN keep-alive timers). +async fn relay_until_eof(reader: &mut R, writer: &mut W) -> Result<()> +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + let mut buf = [0u8; RELAY_BUF_SIZE]; + loop { + match tokio::time::timeout(RELAY_EOF_IDLE_TIMEOUT, reader.read(&mut buf)).await { + Ok(Ok(0)) => return Ok(()), + Ok(Ok(n)) => { + writer.write_all(&buf[..n]).await.into_diagnostic()?; + writer.flush().await.into_diagnostic()?; + } + Ok(Err(e)) => return Err(miette::miette!("{e}")), + Err(_) => { + debug!( + "relay_until_eof idle timeout after {:?}", + RELAY_EOF_IDLE_TIMEOUT + ); + return Ok(()); + } + } + } +} + +/// Relay all bytes from reader to writer until EOF without an idle timeout. +/// +/// Used for server-sent events, where long idle gaps are part of the protocol +/// and do not mean the response body is complete. +async fn relay_until_eof_without_idle_timeout(reader: &mut R, writer: &mut W) -> Result<()> +where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, +{ + let mut buf = [0u8; RELAY_BUF_SIZE]; + loop { + let n = reader.read(&mut buf).await.into_diagnostic()?; + if n == 0 { + return Ok(()); + } + writer.write_all(&buf[..n]).await.into_diagnostic()?; + writer.flush().await.into_diagnostic()?; + } +} + +/// Detect if the first bytes look like an HTTP request. +/// +/// Checks for common HTTP methods at the start of the stream. +pub fn looks_like_http(peek: &[u8]) -> bool { + HTTP_METHOD_PREFIXES + .iter() + .any(|method| peek.starts_with(method)) +} + +pub(crate) fn could_be_http_request_prefix(peek: &[u8]) -> bool { + !peek.is_empty() + && HTTP_METHOD_PREFIXES + .iter() + .any(|method| peek.len() < method.len() && method.starts_with(peek)) +} + +pub fn looks_like_http2_prior_knowledge(peek: &[u8]) -> bool { + peek.len() >= MIN_HTTP2_PREFACE_DETECTION_BYTES + && HTTP2_PRIOR_KNOWLEDGE_PREFACE.starts_with(peek) +} + +pub(crate) fn could_be_http2_prior_knowledge_prefix(peek: &[u8]) -> bool { + !peek.is_empty() + && peek.len() < MIN_HTTP2_PREFACE_DETECTION_BYTES + && HTTP2_PRIOR_KNOWLEDGE_PREFACE.starts_with(peek) +} + +/// Check if an IO error represents a benign connection close. +/// +/// TLS peers commonly close the socket without sending a `close_notify` alert. +/// Rustls reports this as `UnexpectedEof`, but it's functionally equivalent +/// to a clean close when no request data has been received yet. +fn is_benign_close(err: &std::io::Error) -> bool { + matches!( + err.kind(), + std::io::ErrorKind::UnexpectedEof + | std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::BrokenPipe + ) +} + +#[cfg(test)] +#[allow( + clippy::iter_on_single_items, + clippy::manual_string_new, + clippy::collapsible_if, + clippy::cast_possible_truncation, + reason = "Test code: test fixtures and explicit value-shape assertions are idiomatic in tests." +)] +mod tests { + use super::*; + use crate::opa::OpaEngine; + use flate2::{Compress, Compression, Decompress, FlushCompress, FlushDecompress, Status}; + use openshell_core::proposals::AgentProposals; + use openshell_core::proto::{ + Decision, HttpRequestResult, HttpResponseBlockDelivery, HttpResponseBodyMode, + HttpResponseBodyResult, HttpResponseBodyTransform, HttpResponseEvent, + HttpResponseEventResult, HttpResponsePreflightInspect, HttpResponsePreflightResult, + HttpResponseTrailersResult, MiddlewareBinding, MiddlewareManifest, + SupervisorMiddlewareOperation, SupervisorMiddlewarePhase, http_response_body_result, + http_response_body_transform, http_response_body_unit, http_response_event, + http_response_event_result, http_response_preflight_result, + }; + use openshell_core::secrets::SecretResolver; + use std::pin::Pin; + use std::sync::Arc; + use std::task::{Context, Poll}; + use tokio::io::ReadBuf; + use tokio::sync::mpsc; + use tokio_stream::wrappers::ReceiverStream; + + const TEST_POLICY: &str = include_str!("../../data/sandbox-policy.rego"); + const VALID_WS_KEY: &str = "dGhlIHNhbXBsZSBub25jZQ=="; + const VALID_WS_ACCEPT: &str = "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="; + const TEXT_OPCODE: u8 = 0x1; + + #[derive(Clone, Copy)] + enum ResponseRelayScript { + HeadersOnly, + WholeBody, + WholeBodyWithTrailer, + Stream, + BlockPreflight, + BlockWholeBody, + BlockStream, + SlowWholeBody, + InvalidBodySequence, + InvalidWholeBodySequence, + } + + struct ResponseRelayService { + script: ResponseRelayScript, + } + + #[tonic::async_trait] + impl openshell_supervisor_middleware::InProcessMiddleware for ResponseRelayService { + async fn describe(&self) -> MiddlewareManifest { + MiddlewareManifest { + name: "test/response-relay".into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: SupervisorMiddlewareOperation::HttpResponse as i32, + phase: SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: 4096, + timeout: String::new(), + }], + expected_audience: String::new(), + } + } + + async fn validate_config( + &self, + _middleware_name: &str, + _config: &prost_types::Struct, + ) -> Result<()> { + Ok(()) + } + + async fn evaluate_http_request( + &self, + _request: openshell_supervisor_middleware::HttpRequestView<'_>, + ) -> Result { + Ok(HttpRequestResult { + decision: Decision::Allow as i32, + ..Default::default() + }) + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> std::result::Result< + openshell_supervisor_middleware::HttpResponseResultStream, + tonic::Status, + > { + let script = self.script; + let (sender, receiver) = mpsc::channel(4); + tokio::spawn(async move { + while let Some(event) = requests.recv().await { + let Some(event) = event.event else { + break; + }; + let result = match event { + http_response_event::Event::Preflight(_) => { + if matches!(script, ResponseRelayScript::BlockPreflight) { + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some( + http_response_preflight_result::Action::BlockDelivery( + HttpResponseBlockDelivery {}, + ), + ), + reason_code: "content_match".into(), + ..Default::default() + }, + ), + ), + } + } else { + let (body_mode, header_mutations) = match script { + ResponseRelayScript::HeadersOnly => ( + HttpResponseBodyMode::HeadersOnly, + vec![write_header( + "cache-control", + "private", + ExistingHeaderAction::Overwrite, + )], + ), + ResponseRelayScript::WholeBody + | ResponseRelayScript::BlockWholeBody + | ResponseRelayScript::SlowWholeBody + | ResponseRelayScript::InvalidWholeBodySequence + | ResponseRelayScript::WholeBodyWithTrailer => { + (HttpResponseBodyMode::WholeBodyBytes, Vec::new()) + } + ResponseRelayScript::Stream + | ResponseRelayScript::BlockStream + | ResponseRelayScript::InvalidBodySequence => { + (HttpResponseBodyMode::StreamBytes, Vec::new()) + } + ResponseRelayScript::BlockPreflight => unreachable!(), + }; + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some( + http_response_preflight_result::Action::Inspect( + HttpResponsePreflightInspect { + body_mode: body_mode as i32, + header_mutations, + }, + ), + ), + ..Default::default() + }, + ), + ), + } + } + } + http_response_event::Event::Body(body) => { + let Some(http_response_body_unit::Payload::Data(data)) = body.payload + else { + break; + }; + let replacement = match script { + ResponseRelayScript::WholeBody + | ResponseRelayScript::BlockWholeBody + | ResponseRelayScript::SlowWholeBody + | ResponseRelayScript::WholeBodyWithTrailer + | ResponseRelayScript::InvalidWholeBodySequence => { + [b"whole:".as_slice(), &data].concat() + } + ResponseRelayScript::Stream + | ResponseRelayScript::BlockStream + | ResponseRelayScript::InvalidBodySequence => { + data.to_ascii_uppercase() + } + ResponseRelayScript::HeadersOnly + | ResponseRelayScript::BlockPreflight => break, + }; + if matches!(script, ResponseRelayScript::SlowWholeBody) { + tokio::time::sleep(std::time::Duration::from_millis(75)).await; + } + HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: if matches!( + script, + ResponseRelayScript::InvalidBodySequence + | ResponseRelayScript::InvalidWholeBodySequence + ) { + body.sequence + 1 + } else { + body.sequence + }, + action: Some( + if matches!( + script, + ResponseRelayScript::BlockWholeBody + | ResponseRelayScript::BlockStream + ) { + http_response_body_result::Action::BlockDelivery( + HttpResponseBlockDelivery {}, + ) + } else { + http_response_body_result::Action::Transform( + HttpResponseBodyTransform { + replacement: Some( + http_response_body_transform::Replacement::Data( + replacement, + ), + ), + }, + ) + }, + ), + reason_code: if matches!( + script, + ResponseRelayScript::BlockWholeBody + | ResponseRelayScript::BlockStream + ) { + "content_match".into() + } else { + String::new() + }, + ..Default::default() + }, + )), + } + } + http_response_event::Event::Trailers(_) => HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult::default(), + )), + }, + http_response_event::Event::SessionEnd(_) => break, + }; + if sender.send(Ok(result)).await.is_err() { + break; + } + } + }); + Ok(Box::pin(ReceiverStream::new(receiver))) + } + } + + struct CountingReader { + bytes: Vec, + position: usize, + reads: usize, + } + + impl CountingReader { + fn new(bytes: Vec) -> Self { + Self { + bytes, + position: 0, + reads: 0, + } + } + } + + impl AsyncRead for CountingReader { + fn poll_read( + mut self: Pin<&mut Self>, + _context: &mut Context<'_>, + buffer: &mut ReadBuf<'_>, + ) -> Poll> { + self.reads += 1; + let available = self.bytes.len().saturating_sub(self.position); + let amount = available.min(buffer.remaining()); + let end = self.position + amount; + buffer.put_slice(&self.bytes[self.position..end]); + self.position = end; + Poll::Ready(Ok(())) + } + } + + fn write_header(name: &str, value: &str, on_existing: ExistingHeaderAction) -> HeaderMutation { + HeaderMutation { + operation: Some(header_mutation::Operation::Write( + openshell_core::proto::WriteHeader { + name: name.into(), + value: value.into(), + on_existing: on_existing as i32, + }, + )), + } + } + + fn remove_header(name: &str) -> HeaderMutation { + HeaderMutation { + operation: Some(header_mutation::Operation::Remove( + openshell_core::proto::RemoveHeader { name: name.into() }, + )), + } + } + + #[test] + fn ordered_header_mutations_replay_against_raw_request() { + let raw = b"GET / HTTP/1.1\r\nHost: example.test\r\nX-OpenShell-Middleware-Chain: first\r\nX-Drop: one\r\nX-Drop: two\r\n\r\n"; + let mutations = [ + write_header( + "x-openshell-middleware-chain", + "second", + ExistingHeaderAction::Append, + ), + write_header( + "x-openshell-middleware-chain", + "ignored", + ExistingHeaderAction::Skip, + ), + write_header( + "x-openshell-middleware-chain", + "replacement", + ExistingHeaderAction::Overwrite, + ), + write_header( + "x-openshell-middleware-chain", + "tail", + ExistingHeaderAction::Append, + ), + remove_header("x-drop"), + ]; + + let updated = String::from_utf8( + apply_header_mutations(raw, &mutations).expect("apply ordered header mutations"), + ) + .expect("UTF-8 request"); + let values: Vec<&str> = updated + .lines() + .filter_map(|line| { + line.split_once(':').and_then(|(name, value)| { + name.eq_ignore_ascii_case("x-openshell-middleware-chain") + .then_some(value.trim()) + }) + }) + .collect(); + assert_eq!(values, vec!["replacement", "tail"]); + assert!(!updated.to_ascii_lowercase().contains("x-drop:")); + assert!(updated.contains("Host: example.test")); + } + + #[derive(Debug)] + struct CapturedFrame { + fin_opcode: u8, + masked: bool, + payload: Vec, + } + + async fn read_http_header_block(reader: &mut R) -> Vec { + tokio::time::timeout(std::time::Duration::from_secs(2), async { + let mut header = Vec::new(); + let mut byte = [0u8; 1]; + loop { + reader.read_exact(&mut byte).await.unwrap(); + header.push(byte[0]); + if header.ends_with(b"\r\n\r\n") { + break; + } + } + header + }) + .await + .expect("HTTP header block should arrive") + } + + async fn read_websocket_frame(reader: &mut R) -> CapturedFrame { + tokio::time::timeout(std::time::Duration::from_secs(2), async { + let mut prefix = [0u8; 2]; + reader.read_exact(&mut prefix).await.unwrap(); + let masked = prefix[1] & 0x80 != 0; + let mut payload_len = u64::from(prefix[1] & 0x7f); + if payload_len == 126 { + let mut extended = [0u8; 2]; + reader.read_exact(&mut extended).await.unwrap(); + payload_len = u64::from(u16::from_be_bytes(extended)); + } else if payload_len == 127 { + let mut extended = [0u8; 8]; + reader.read_exact(&mut extended).await.unwrap(); + payload_len = u64::from_be_bytes(extended); + } + let mut mask_key = [0u8; 4]; + if masked { + reader.read_exact(&mut mask_key).await.unwrap(); + } + let payload_len = usize::try_from(payload_len).unwrap(); + let mut payload = vec![0u8; payload_len]; + reader.read_exact(&mut payload).await.unwrap(); + if masked { + apply_test_mask(&mut payload, mask_key); + } + CapturedFrame { + fin_opcode: prefix[0], + masked, + payload, + } + }) + .await + .expect("WebSocket frame should arrive") + } + + async fn policy_local_json_response( + ctx: Arc, + ) -> serde_json::Value { + let (mut client, mut server) = tokio::io::duplex(4096); + let task = tokio::spawn(async move { + crate::policy_local::handle_forward_request( + ctx.as_ref(), + "GET", + "/v1/policy/current", + b"GET http://policy.local/v1/policy/current HTTP/1.1\r\nHost: policy.local\r\n\r\n", + &mut server, + ) + .await + .unwrap(); + }); + + let mut received = Vec::new(); + client.read_to_end(&mut received).await.unwrap(); + task.await.unwrap(); + + let response = String::from_utf8(received).unwrap(); + let (_, body) = response.split_once("\r\n\r\n").unwrap(); + serde_json::from_str(body).unwrap() + } + + fn masked_frame_with_rsv(opcode: u8, rsv: u8, payload: &[u8]) -> Vec { + let mask_key = [0x37, 0xfa, 0x21, 0x3d]; + let mut frame = Vec::new(); + frame.push(0x80 | rsv | opcode); + write_test_payload_len(&mut frame, 0x80, payload.len()); + frame.extend_from_slice(&mask_key); + let mut masked = payload.to_vec(); + apply_test_mask(&mut masked, mask_key); + frame.extend_from_slice(&masked); + frame + } + + fn unmasked_frame(opcode: u8, payload: &[u8]) -> Vec { + let mut frame = Vec::new(); + frame.push(0x80 | opcode); + write_test_payload_len(&mut frame, 0, payload.len()); + frame.extend_from_slice(payload); + frame + } + + fn write_test_payload_len(frame: &mut Vec, mask_bit: u8, payload_len: usize) { + if payload_len < 126 { + frame.push(mask_bit | payload_len as u8); + } else if u16::try_from(payload_len).is_ok() { + frame.push(mask_bit | 0x7e); + frame.extend_from_slice(&(payload_len as u16).to_be_bytes()); + } else { + frame.push(mask_bit | 0x7f); + frame.extend_from_slice(&(payload_len as u64).to_be_bytes()); + } + } + + fn apply_test_mask(payload: &mut [u8], mask_key: [u8; 4]) { + for (index, byte) in payload.iter_mut().enumerate() { + *byte ^= mask_key[index % 4]; + } + } + + fn compress_test_permessage_deflate(payload: &[u8]) -> Vec { + let mut compressor = Compress::new(Compression::fast(), false); + let mut out = Vec::with_capacity(payload.len().saturating_add(128)); + loop { + let consumed = usize::try_from(compressor.total_in()).unwrap(); + if consumed >= payload.len() { + break; + } + let before_in = compressor.total_in(); + let before_out = compressor.total_out(); + let status = compressor + .compress_vec(&payload[consumed..], &mut out, FlushCompress::None) + .unwrap(); + if matches!(status, Status::BufError) + || (compressor.total_in() == before_in && compressor.total_out() == before_out) + { + out.reserve(out.capacity().max(1024)); } } loop { @@ -4756,751 +6579,1527 @@ mod tests { panic!("aggregate chunk extensions must be bounded") }; assert!( - error - .to_string() - .contains("chunked body wire representation exceeds configured buffer limit"), - "unexpected error: {error}" + error + .to_string() + .contains("chunked body wire representation exceeds configured buffer limit"), + "unexpected error: {error}" + ); + } + + #[tokio::test] + async fn credential_rewrite_chunked_request_with_trailers_is_rejected() { + let req = L7Request { + action: "POST".into(), + target: "/api".into(), + query_params: HashMap::new(), + raw_header: Vec::new(), + body_length: BodyLength::Chunked, + }; + let headers = + b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\nTrailer: Digest\r\n\r\n"; + let result = collect_and_rewrite_request_body( + &req, + &mut tokio::io::empty(), + headers, + std::str::from_utf8(headers).expect("headers"), + b"1\r\nx\r\n0\r\nDigest: sha-256=:abc123:\r\n\r\n", + None, + None, + ) + .await; + let Err(error) = result else { + panic!("credential rewriting must reject non-empty chunked trailers") + }; + assert!( + error.to_string().contains( + "chunked request trailers are not supported when buffering or transforming request bodies" + ), + "unexpected error: {error}" + ); + } + + #[tokio::test] + async fn collect_chunked_body_reads_payload_in_blocks() { + let payload_len = 64 * 1024; + let mut wire = format!("{payload_len:x}\r\n").into_bytes(); + wire.extend(std::iter::repeat_n(b'x', payload_len)); + wire.extend_from_slice(b"\r\n0\r\n\r\n"); + let mut client = CountingReader::new(wire); + + let body = collect_chunked_body( + &mut client, + &[], + None, + Some(openshell_supervisor_middleware::MAX_MIDDLEWARE_PAYLOAD_BYTES), + ) + .await + .expect("chunked body should decode"); + + assert_eq!(body.len(), payload_len); + assert!( + client.reads <= 32, + "payload should be read in blocks, observed {} reads", + client.reads + ); + } + + #[tokio::test] + async fn extreme_content_length_is_rejected_before_allocation() { + let req = L7Request { + action: "POST".into(), + target: "/upload".into(), + query_params: HashMap::new(), + raw_header: b"POST /upload HTTP/1.1\r\nHost: example.com\r\nContent-Length: 18446744073709551615\r\n\r\n".to_vec(), + body_length: BodyLength::ContentLength(u64::MAX), + }; + let (mut client, _peer) = tokio::io::duplex(1); + + let result = buffer_request_body_for_middleware( + &req, + &mut client, + None, + openshell_supervisor_middleware::MAX_MIDDLEWARE_PAYLOAD_BYTES, + ) + .await + .expect("oversized body should produce a capacity result"); + + assert!(matches!( + result, + BufferResult::OverCapacity { recoverable: true } + )); + } + + #[tokio::test] + async fn middleware_chunked_wire_body_at_cap_is_allowed() { + let max_body_bytes = max_middleware_body_bytes().await; + let payload_len = max_body_bytes - 14; + let mut wire = format!("{payload_len:x}\r\n").into_bytes(); + wire.extend(std::iter::repeat_n(b'x', payload_len)); + wire.extend_from_slice(b"\r\n0\r\n\r\n"); + assert_eq!(wire.len(), max_body_bytes); + + let body = collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) + .await + .expect("wire representation at the cap should be allowed"); + + assert_eq!(body.len(), payload_len); + } + + #[tokio::test] + async fn middleware_chunked_wire_body_over_cap_is_rejected() { + let max_body_bytes = max_middleware_body_bytes().await; + let payload_len = max_body_bytes - 13; + let mut wire = format!("{payload_len:x}\r\n").into_bytes(); + wire.extend(std::iter::repeat_n(b'x', payload_len)); + wire.extend_from_slice(b"\r\n0\r\n\r\n"); + assert_eq!(wire.len(), max_body_bytes + 1); + assert!(payload_len < max_body_bytes); + + let error = + collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) + .await + .expect_err("wire framing over the cap must be rejected"); + + assert!( + matches!(error, CollectChunkedError::OverCapacity), + "over-cap wire body must be OverCapacity, got {error:?}" + ); + } + + #[tokio::test] + async fn middleware_chunked_body_can_exceed_credential_rewrite_limit() { + let max_body_bytes = 1024 * 1024; + let payload_len = 300 * 1024; + assert!(payload_len > MAX_REWRITE_BODY_BYTES); + let mut wire = format!("{payload_len:x}\r\n").into_bytes(); + wire.extend(std::iter::repeat_n(b'x', payload_len)); + wire.extend_from_slice(b"\r\n0\r\n\r\n"); + + let body = collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) + .await + .expect("middleware cap should control chunked body collection"); + + assert_eq!(body.len(), payload_len); + } + + #[tokio::test] + async fn middleware_chunked_invalid_size_is_not_over_capacity() { + let mut raw = + b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n" + .to_vec(); + raw.extend_from_slice(b"xyz\r\n"); + let req = L7Request { + action: "POST".into(), + target: "/api".into(), + query_params: HashMap::new(), + raw_header: raw, + body_length: BodyLength::Chunked, + }; + let err = + buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + .await + .expect_err("invalid chunk framing must surface as an error"); + + assert!( + err.to_string().contains("Invalid chunk size token"), + "unexpected error: {err}" + ); + assert!( + !err.to_string().contains("over_capacity") + && !err.to_string().contains("exceeds configured buffer limit"), + "protocol errors must not be reported as over-capacity: {err}" ); } #[tokio::test] - async fn credential_rewrite_chunked_request_with_trailers_is_rejected() { + async fn middleware_chunked_over_capacity_still_maps_to_buffer_over_capacity() { + let max_body_bytes = 32; + let payload = "hello world that is definitely over the tiny cap"; + let wire = format!("{:x}\r\n{payload}\r\n0\r\n\r\n", payload.len()); + let mut raw = + b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n" + .to_vec(); + raw.extend_from_slice(wire.as_bytes()); let req = L7Request { action: "POST".into(), target: "/api".into(), query_params: HashMap::new(), - raw_header: Vec::new(), + raw_header: raw, body_length: BodyLength::Chunked, }; - let headers = - b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\nTrailer: Digest\r\n\r\n"; - let result = collect_and_rewrite_request_body( - &req, - &mut tokio::io::empty(), - headers, - std::str::from_utf8(headers).expect("headers"), - b"1\r\nx\r\n0\r\nDigest: sha-256=:abc123:\r\n\r\n", - None, - None, - ) - .await; - let Err(error) = result else { - panic!("credential rewriting must reject non-empty chunked trailers") - }; + + let result = + buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, max_body_bytes) + .await + .expect("over-capacity is a BufferResult, not an Err"); + assert!( - error.to_string().contains( - "chunked request trailers are not supported when buffering or transforming request bodies" - ), - "unexpected error: {error}" + matches!(result, BufferResult::OverCapacity { recoverable: false }), + "expected OverCapacity, got {result:?}" ); } #[tokio::test] - async fn collect_chunked_body_reads_payload_in_blocks() { - let payload_len = 64 * 1024; - let mut wire = format!("{payload_len:x}\r\n").into_bytes(); - wire.extend(std::iter::repeat_n(b'x', payload_len)); - wire.extend_from_slice(b"\r\n0\r\n\r\n"); - let mut client = CountingReader::new(wire); + async fn middleware_none_body_with_header_overshoot_is_rejected() { + // Mimic the forward-proxy multi-byte read: headers plus pipelined bytes + // after `\r\n\r\n` on a request with no body framing. + let raw = b"GET /api HTTP/1.1\r\nHost: example.com\r\n\r\nGET /other HTTP/1.1\r\n"; + let req = L7Request { + action: "GET".into(), + target: "/api".into(), + query_params: HashMap::new(), + raw_header: raw.to_vec(), + body_length: BodyLength::None, + }; - let body = collect_chunked_body( - &mut client, - &[], - None, - Some(openshell_supervisor_middleware::MAX_MIDDLEWARE_PAYLOAD_BYTES), - ) - .await - .expect("chunked body should decode"); + let err = + buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + .await + .expect_err("read-ahead leftovers must not become a request body"); - assert_eq!(body.len(), payload_len); assert!( - client.reads <= 32, - "payload should be read in blocks, observed {} reads", - client.reads + err.to_string().contains("no body framing"), + "unexpected error: {err}" ); } #[tokio::test] - async fn extreme_content_length_is_rejected_before_allocation() { + async fn middleware_none_body_without_overshoot_buffers_empty() { + let raw = b"GET /api HTTP/1.1\r\nHost: example.com\r\n\r\n"; let req = L7Request { - action: "POST".into(), - target: "/upload".into(), + action: "GET".into(), + target: "/api".into(), query_params: HashMap::new(), - raw_header: b"POST /upload HTTP/1.1\r\nHost: example.com\r\nContent-Length: 18446744073709551615\r\n\r\n".to_vec(), - body_length: BodyLength::ContentLength(u64::MAX), + raw_header: raw.to_vec(), + body_length: BodyLength::None, }; - let (mut client, _peer) = tokio::io::duplex(1); - let result = buffer_request_body_for_middleware( - &req, + let result = + buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + .await + .expect("empty no-body request should buffer"); + + match result { + BufferResult::Buffered(buffered) => { + assert!(buffered.body.is_empty()); + let rebuilt = rebuild_request_with_buffered_body( + &req, + &buffered.headers, + &buffered.body, + &[], + ) + .expect("rebuild no-body request"); + assert!(matches!(rebuilt.body_length, BodyLength::None)); + let text = String::from_utf8(rebuilt.raw_header).unwrap(); + assert!( + !text.to_ascii_lowercase().contains("content-length"), + "rebuild must preserve no-body framing: {text}" + ); + assert!(!text.contains("GET /other")); + } + other @ BufferResult::OverCapacity { .. } => { + panic!("expected Buffered, got {other:?}") + } + } + } + + /// SEC-009: Bare LF in headers enables header injection. + #[tokio::test] + async fn reject_bare_lf_in_headers() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + // Bare \n between two header values creates a parsing discrepancy + writer + .write_all( + b"GET /api HTTP/1.1\r\nX-Injected: value\nEvil: header\r\nHost: x\r\n\r\n", + ) + .await + .unwrap(); + }); + let result = parse_http_request( &mut client, - None, - openshell_supervisor_middleware::MAX_MIDDLEWARE_PAYLOAD_BYTES, + &crate::l7::path::CanonicalizeOptions::default(), ) - .await - .expect("oversized body should produce a capacity result"); + .await; + assert!(result.is_err(), "Must reject headers with bare LF"); + } - assert!(matches!( - result, - BufferResult::OverCapacity { recoverable: true } - )); + /// SEC-009: Invalid UTF-8 in headers creates interpretation gap. + #[tokio::test] + async fn reject_invalid_utf8_in_headers() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + let mut raw = Vec::new(); + raw.extend_from_slice(b"GET /api HTTP/1.1\r\nHost: x\r\nX-Bad: \xc0\xaf\r\n\r\n"); + writer.write_all(&raw).await.unwrap(); + }); + let result = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await; + assert!(result.is_err(), "Must reject headers with invalid UTF-8"); } #[tokio::test] - async fn middleware_chunked_wire_body_at_cap_is_allowed() { - let max_body_bytes = max_middleware_body_bytes().await; - let payload_len = max_body_bytes - 14; - let mut wire = format!("{payload_len:x}\r\n").into_bytes(); - wire.extend(std::iter::repeat_n(b'x', payload_len)); - wire.extend_from_slice(b"\r\n0\r\n\r\n"); - assert_eq!(wire.len(), max_body_bytes); + async fn reject_malformed_header_fields_before_forwarding() { + let cases = [ + ( + "space continuation", + b"GET /api HTTP/1.1\r\nX-Test: first\r\n continued\r\nHost: x\r\n\r\n".as_slice(), + ), + ( + "tab continuation", + b"GET /api HTTP/1.1\r\nX-Test: first\r\n\tcontinued\r\nHost: x\r\n\r\n".as_slice(), + ), + ( + "missing colon", + b"GET /api HTTP/1.1\r\nX-Test value\r\nHost: x\r\n\r\n".as_slice(), + ), + ( + "whitespace before colon", + b"GET /api HTTP/1.1\r\nX-Test : value\r\nHost: x\r\n\r\n".as_slice(), + ), + ( + "invalid field-name token", + b"GET /api HTTP/1.1\r\nX@Test: value\r\nHost: x\r\n\r\n".as_slice(), + ), + ]; - let body = collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) - .await - .expect("wire representation at the cap should be allowed"); + for (case, raw) in cases { + let (mut client, mut writer) = tokio::io::duplex(4096); + writer.write_all(raw).await.unwrap(); + let result = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await; + assert!(result.is_err(), "{case} must be rejected before forwarding"); + } + } - assert_eq!(body.len(), payload_len); + /// SEC-009: Reject unsupported HTTP versions. + #[tokio::test] + async fn reject_invalid_http_version() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /api JUNK/9.9\r\nHost: x\r\n\r\n") + .await + .unwrap(); + }); + let result = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await; + assert!(result.is_err(), "Must reject unsupported HTTP version"); } #[tokio::test] - async fn middleware_chunked_wire_body_over_cap_is_rejected() { - let max_body_bytes = max_middleware_body_bytes().await; - let payload_len = max_body_bytes - 13; - let mut wire = format!("{payload_len:x}\r\n").into_bytes(); - wire.extend(std::iter::repeat_n(b'x', payload_len)); - wire.extend_from_slice(b"\r\n0\r\n\r\n"); - assert_eq!(wire.len(), max_body_bytes + 1); - assert!(payload_len < max_body_bytes); + async fn parse_http_request_canonicalizes_target_and_rewrites_raw_header() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /public/../secret HTTP/1.1\r\nHost: api.example.com\r\n\r\n") + .await + .unwrap(); + }); + let req = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .expect("request should parse") + .expect("request should exist"); + // Path fed to OPA evaluation is canonical. + assert_eq!(req.target, "/secret"); + // raw_header (forwarded byte-for-byte to upstream) is also canonical + // — this is the invariant the L7 canonicalization PR must uphold. + assert_eq!( + req.raw_header, b"GET /secret HTTP/1.1\r\nHost: api.example.com\r\n\r\n", + "outbound request line must carry the canonical path" + ); + } - let error = - collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) - .await - .expect_err("wire framing over the cap must be rejected"); + #[tokio::test] + async fn parse_http_request_rejects_absolute_authority_mismatched_with_host() { + let (mut client, mut peer) = tokio::io::duplex(1024); + peer.write_all( + b"GET http://attacker.example.test/v1 HTTP/1.1\r\nHost: api.example.test\r\n\r\n", + ) + .await + .unwrap(); + let error = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .expect_err("absolute-form authority mismatch must fail closed"); assert!( - matches!(error, CollectChunkedError::OverCapacity), - "over-cap wire body must be OverCapacity, got {error:?}" + error + .to_string() + .contains("request authority does not match the Host header"), + "{error}" ); } - #[tokio::test] - async fn middleware_chunked_body_can_exceed_credential_rewrite_limit() { - let max_body_bytes = 1024 * 1024; - let payload_len = 300 * 1024; - assert!(payload_len > MAX_REWRITE_BODY_BYTES); - let mut wire = format!("{payload_len:x}\r\n").into_bytes(); - wire.extend(std::iter::repeat_n(b'x', payload_len)); - wire.extend_from_slice(b"\r\n0\r\n\r\n"); + #[test] + fn origin_form_targets_with_embedded_urls_use_host_authority() { + let host: http::uri::Authority = "api.example.test".parse().unwrap(); - let body = collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) - .await - .expect("middleware cap should control chunked body collection"); + for target in ["/fetch/http://example.test", "/?next=http://example.test"] { + assert!( + absolute_form_uri(target).unwrap().is_none(), + "{target} must remain origin-form" + ); + validate_absolute_form_authority(target, Some(&host)) + .expect("embedded URL must not trigger absolute-form validation"); - assert_eq!(body.len(), payload_len); + let raw = format!("GET {target} HTTP/1.1\r\nHost: api.example.test\r\n\r\n"); + let authority = request_authority(raw.as_bytes(), Some(443)) + .unwrap() + .expect("origin-form request with Host must have an authority"); + assert_eq!(authority.authority, host); + assert_eq!(authority.effective_port, 443); + } } #[tokio::test] - async fn middleware_chunked_invalid_size_is_not_over_capacity() { - let mut raw = - b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n" - .to_vec(); - raw.extend_from_slice(b"xyz\r\n"); - let req = L7Request { - action: "POST".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw, - body_length: BodyLength::Chunked, - }; - let err = - buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) - .await - .expect_err("invalid chunk framing must surface as an error"); + async fn parse_http_request_keeps_embedded_url_in_origin_form_path() { + let (mut client, mut peer) = tokio::io::duplex(1024); + peer.write_all( + b"GET /fetch/http://example.test?next=http://other.test HTTP/1.1\r\nHost: api.example.test\r\n\r\n", + ) + .await + .unwrap(); - assert!( - err.to_string().contains("Invalid chunk size token"), - "unexpected error: {err}" - ); - assert!( - !err.to_string().contains("over_capacity") - && !err.to_string().contains("exceeds configured buffer limit"), - "protocol errors must not be reported as over-capacity: {err}" + let request = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .expect("embedded URL origin-form request must parse") + .expect("request must be present"); + assert_eq!(request.target, "/fetch/http:/example.test"); + assert_eq!( + request.raw_header, + b"GET /fetch/http:/example.test?next=http://other.test HTTP/1.1\r\nHost: api.example.test\r\n\r\n", ); } #[tokio::test] - async fn middleware_chunked_over_capacity_still_maps_to_buffer_over_capacity() { - let max_body_bytes = 32; - let payload = "hello world that is definitely over the tiny cap"; - let wire = format!("{:x}\r\n{payload}\r\n0\r\n\r\n", payload.len()); - let mut raw = - b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n" - .to_vec(); - raw.extend_from_slice(wire.as_bytes()); - let req = L7Request { - action: "POST".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw, - body_length: BodyLength::Chunked, - }; - - let result = - buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, max_body_bytes) + async fn parse_http_request_canonicalization_preserves_query_string() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /public/../v1/list?limit=10&sort=asc HTTP/1.1\r\nHost: h\r\n\r\n") .await - .expect("over-capacity is a BufferResult, not an Err"); - - assert!( - matches!(result, BufferResult::OverCapacity { recoverable: false }), - "expected OverCapacity, got {result:?}" + .unwrap(); + }); + let req = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(req.target, "/v1/list"); + assert_eq!( + req.raw_header, b"GET /v1/list?limit=10&sort=asc HTTP/1.1\r\nHost: h\r\n\r\n", + "canonical rewrite must preserve the query string verbatim" ); } #[tokio::test] - async fn middleware_none_body_with_header_overshoot_is_rejected() { - // Mimic the forward-proxy multi-byte read: headers plus pipelined bytes - // after `\r\n\r\n` on a request with no body framing. - let raw = b"GET /api HTTP/1.1\r\nHost: example.com\r\n\r\nGET /other HTTP/1.1\r\n"; - let req = L7Request { - action: "GET".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw.to_vec(), - body_length: BodyLength::None, - }; - - let err = - buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + async fn parse_http_request_leaves_canonical_input_byte_for_byte() { + // When the input is already canonical, the raw_header must pass + // through unchanged — otherwise legitimate traffic pays a rewrite + // cost on every request. + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /api/v1/users HTTP/1.1\r\nHost: api.example.com\r\n\r\n") .await - .expect_err("read-ahead leftovers must not become a request body"); + .unwrap(); + }); + let req = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(req.target, "/api/v1/users"); + assert_eq!( + req.raw_header, + b"GET /api/v1/users HTTP/1.1\r\nHost: api.example.com\r\n\r\n", + ); + } + #[tokio::test] + async fn parse_http_request_rejects_traversal_above_root() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /.. HTTP/1.1\r\nHost: h\r\n\r\n") + .await + .unwrap(); + }); + let result = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await; assert!( - err.to_string().contains("no body framing"), - "unexpected error: {err}" + result.is_err(), + "a target that escapes the path root must be rejected at the parser" ); } #[tokio::test] - async fn middleware_none_body_without_overshoot_buffers_empty() { - let raw = b"GET /api HTTP/1.1\r\nHost: example.com\r\n\r\n"; - let req = L7Request { - action: "GET".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw.to_vec(), - body_length: BodyLength::None, - }; - - let result = - buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + async fn parse_http_request_accepts_encoded_slash_when_endpoint_opts_in() { + // GitLab-style endpoints legitimately embed `%2F` in path segments + // (e.g. `/api/v4/projects/group%2Fproject`). Passing a provider + // constructed with `allow_encoded_slash: true` models the + // endpoint-config wiring that flows from `L7EndpointConfig`. + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /api/v4/projects/group%2Fproject HTTP/1.1\r\nHost: g\r\n\r\n") .await - .expect("empty no-body request should buffer"); + .unwrap(); + }); + let options = crate::l7::path::CanonicalizeOptions { + allow_encoded_slash: true, + ..Default::default() + }; + let req = parse_http_request(&mut client, &options) + .await + .unwrap() + .unwrap(); + assert_eq!(req.target, "/api/v4/projects/group%2Fproject"); + } - match result { - BufferResult::Buffered(buffered) => { - assert!(buffered.body.is_empty()); - let rebuilt = rebuild_request_with_buffered_body( - &req, - &buffered.headers, - &buffered.body, - &[], - ) - .expect("rebuild no-body request"); - assert!(matches!(rebuilt.body_length, BodyLength::None)); - let text = String::from_utf8(rebuilt.raw_header).unwrap(); - assert!( - !text.to_ascii_lowercase().contains("content-length"), - "rebuild must preserve no-body framing: {text}" - ); - assert!(!text.contains("GET /other")); - } - other @ BufferResult::OverCapacity { .. } => { - panic!("expected Buffered, got {other:?}") - } - } + #[tokio::test] + async fn parse_http_request_rejects_encoded_slash_by_default() { + // Default strict options must reject `%2F` — this is the security + // posture for endpoints where an encoded slash would let an + // attacker disagree with the upstream on segment boundaries. + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /api/v4/projects/group%2Fproject HTTP/1.1\r\nHost: g\r\n\r\n") + .await + .unwrap(); + }); + let result = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await; + assert!( + result.is_err(), + "default options must reject encoded slashes in the path" + ); } - /// SEC-009: Bare LF in headers enables header injection. #[tokio::test] - async fn reject_bare_lf_in_headers() { + async fn parse_http_request_preserves_http_10_version_on_rewrite() { + let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { + writer + .write_all(b"GET /a/./b HTTP/1.0\r\nHost: h\r\n\r\n") + .await + .unwrap(); + }); + let req = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(req.target, "/a/b"); + assert!( + req.raw_header.starts_with(b"GET /a/b HTTP/1.0\r\n"), + "rewrite must preserve the original HTTP version, got: {:?}", + String::from_utf8_lossy(&req.raw_header) + ); + } + + #[tokio::test] + async fn parse_http_request_splits_path_and_query_params() { let (mut client, mut writer) = tokio::io::duplex(4096); tokio::spawn(async move { - // Bare \n between two header values creates a parsing discrepancy writer .write_all( - b"GET /api HTTP/1.1\r\nX-Injected: value\nEvil: header\r\nHost: x\r\n\r\n", + b"GET /download?slug=my%2Fskill&tag=foo&tag=bar HTTP/1.1\r\nHost: x\r\n\r\n", ) .await .unwrap(); }); - let result = parse_http_request( + let req = parse_http_request( &mut client, &crate::l7::path::CanonicalizeOptions::default(), ) - .await; - assert!(result.is_err(), "Must reject headers with bare LF"); + .await + .expect("request should parse") + .expect("request should exist"); + assert_eq!(req.target, "/download"); + assert_eq!( + req.query_params.get("slug").cloned(), + Some(vec!["my/skill".to_string()]) + ); + assert_eq!( + req.query_params.get("tag").cloned(), + Some(vec!["foo".to_string(), "bar".to_string()]) + ); } - /// SEC-009: Invalid UTF-8 in headers creates interpretation gap. + /// Regression test: two pipelined requests in a single write must be + /// parsed independently. Before the fix, the 1024-byte `read()` buffer + /// could capture bytes from the second request, which were forwarded + /// upstream as body overflow of the first -- bypassing L7 policy checks. #[tokio::test] - async fn reject_invalid_utf8_in_headers() { + async fn parse_http_request_does_not_overread_next_request() { let (mut client, mut writer) = tokio::io::duplex(4096); + tokio::spawn(async move { - let mut raw = Vec::new(); - raw.extend_from_slice(b"GET /api HTTP/1.1\r\nHost: x\r\nX-Bad: \xc0\xaf\r\n\r\n"); - writer.write_all(&raw).await.unwrap(); + writer + .write_all( + b"GET /allowed HTTP/1.1\r\nHost: example.com\r\n\r\n\ + POST /blocked HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n", + ) + .await + .unwrap(); }); - let result = parse_http_request( + + let first = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .expect("first request should parse") + .expect("expected first request"); + assert_eq!(first.action, "GET"); + assert_eq!(first.target, "/allowed"); + assert!(first.query_params.is_empty()); + assert_eq!( + first.raw_header, b"GET /allowed HTTP/1.1\r\nHost: example.com\r\n\r\n", + "raw_header must contain only the first request's headers" + ); + + let second = parse_http_request( &mut client, &crate::l7::path::CanonicalizeOptions::default(), ) + .await + .expect("second request should parse") + .expect("expected second request"); + assert_eq!(second.action, "POST"); + assert_eq!(second.target, "/blocked"); + assert!(second.query_params.is_empty()); + } + + #[test] + fn http_method_detection() { + assert!(looks_like_http(b"GET / HTTP/1.1\r\n")); + assert!(looks_like_http(b"POST /api HTTP/1.1\r\n")); + assert!(looks_like_http(b"DELETE /foo HTTP/1.1\r\n")); + assert!(could_be_http_request_prefix(b"GE")); + assert!(!could_be_http_request_prefix(b"GET ")); + assert!(!looks_like_http(b"\x00\x00\x00\x08")); // Postgres + assert!(!looks_like_http(HTTP2_PRIOR_KNOWLEDGE_PREFACE)); + assert!(!looks_like_http(b"HELLO")); // Unknown + } + + #[test] + fn http2_prior_knowledge_detection() { + assert!(looks_like_http2_prior_knowledge( + HTTP2_PRIOR_KNOWLEDGE_PREFACE + )); + assert!(looks_like_http2_prior_knowledge( + &HTTP2_PRIOR_KNOWLEDGE_PREFACE[..8] + )); + assert!(could_be_http2_prior_knowledge_prefix(b"PRI * H")); + assert!(!looks_like_http2_prior_knowledge(b"PRI * H")); + assert!(!looks_like_http2_prior_knowledge(b"PRI / HTTP/1.1\r\n")); + } + + #[test] + fn test_parse_status_code() { + assert_eq!( + parse_status_code("HTTP/1.1 200 OK\r\nHost: x\r\n\r\n"), + Some(200) + ); + assert_eq!( + parse_status_code("HTTP/1.1 204 No Content\r\n\r\n"), + Some(204) + ); + assert_eq!( + parse_status_code("HTTP/1.1 304 Not Modified\r\n\r\n"), + Some(304) + ); + assert_eq!( + parse_status_code("HTTP/1.1 100 Continue\r\n\r\n"), + Some(100) + ); + assert_eq!(parse_status_code(""), None); + } + + #[test] + fn test_parse_connection_close() { + assert!(parse_connection_close( + "HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n" + )); + assert!(!parse_connection_close( + "HTTP/1.1 200 OK\r\nConnection: keep-alive\r\n\r\n" + )); + assert!(!parse_connection_close( + "HTTP/1.1 200 OK\r\nHost: x\r\n\r\n" + )); + } + + #[test] + fn test_response_is_event_stream() { + assert!(response_is_event_stream( + "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n" + )); + assert!(response_is_event_stream( + "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream; charset=utf-8\r\n\r\n" + )); + assert!(!response_is_event_stream( + "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n\r\n" + )); + } + + #[test] + fn test_is_bodiless_response() { + assert!(is_bodiless_response("HEAD", 200)); + assert!(is_bodiless_response("GET", 100)); + assert!(is_bodiless_response("GET", 199)); + assert!(is_bodiless_response("GET", 204)); + assert!(is_bodiless_response("GET", 304)); + assert!(!is_bodiless_response("GET", 200)); + assert!(!is_bodiless_response("POST", 201)); + } + + fn response_middleware_fixture( + script: ResponseRelayScript, + ) -> ( + openshell_supervisor_middleware::ChainRunner, + Vec, + ) { + response_middleware_fixture_with_error( + script, + openshell_supervisor_middleware::OnError::FailClosed, + ) + } + + fn response_middleware_fixture_with_error( + script: ResponseRelayScript, + on_error: openshell_supervisor_middleware::OnError, + ) -> ( + openshell_supervisor_middleware::ChainRunner, + Vec, + ) { + let runner = + openshell_supervisor_middleware::ChainRunner::new(Arc::new(ResponseRelayService { + script, + })); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "response".into(), + implementation: "test/response-relay".into(), + order: 0, + config: prost_types::Struct::default(), + on_error, + }]; + (runner, chain) + } + + fn response_middleware_context<'a>( + runner: &'a openshell_supervisor_middleware::ChainRunner, + chain: &'a [openshell_supervisor_middleware::ChainEntry], + method: &str, + ) -> HttpResponseMiddlewareRelay<'a> { + HttpResponseMiddlewareRelay { + chain, + runner, + request_context: RequestContext { + request_id: "request-1".into(), + sandbox_id: "sandbox-1".into(), + ..Default::default() + }, + target: HttpRequestTarget { + scheme: "https".into(), + host: "example.test".into(), + port: 443, + method: method.into(), + path: "/data".into(), + query: String::new(), + }, + policy_name: "test-policy", + generation_guard: None, + whole_body_timeout: std::time::Duration::from_secs(120), + } + } + + async fn run_response_middleware_relay( + response: &'static [u8], + method: &str, + script: ResponseRelayScript, + ) -> (Result, Vec) { + run_response_middleware_relay_with_error( + response, + method, + script, + openshell_supervisor_middleware::OnError::FailClosed, + ) + .await + } + + async fn run_response_middleware_relay_with_error( + response: &'static [u8], + method: &str, + script: ResponseRelayScript, + on_error: openshell_supervisor_middleware::OnError, + ) -> (Result, Vec) { + run_response_middleware_relay_with_timeout( + response, + method, + script, + on_error, + std::time::Duration::from_secs(120), + ) + .await + } + + async fn run_response_middleware_relay_with_timeout( + response: &'static [u8], + method: &str, + script: ResponseRelayScript, + on_error: openshell_supervisor_middleware::OnError, + whole_body_timeout: std::time::Duration, + ) -> (Result, Vec) { + let (runner, chain) = response_middleware_fixture_with_error(script, on_error); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(16 * 1024); + let (mut client_read, mut client_write) = tokio::io::duplex(16 * 1024); + tokio::spawn(async move { + upstream_write.write_all(response).await.unwrap(); + upstream_write.shutdown().await.unwrap(); + }); + let mut middleware = response_middleware_context(&runner, &chain, method); + middleware.whole_body_timeout = whole_body_timeout; + let outcome = relay_response( + method, + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(middleware), + ) .await; - assert!(result.is_err(), "Must reject headers with invalid UTF-8"); + drop(client_write); + let mut delivered = Vec::new(); + client_read.read_to_end(&mut delivered).await.unwrap(); + (outcome, delivered) } #[tokio::test] - async fn reject_malformed_header_fields_before_forwarding() { - let cases = [ - ( - "space continuation", - b"GET /api HTTP/1.1\r\nX-Test: first\r\n continued\r\nHost: x\r\n\r\n".as_slice(), - ), - ( - "tab continuation", - b"GET /api HTTP/1.1\r\nX-Test: first\r\n\tcontinued\r\nHost: x\r\n\r\n".as_slice(), - ), - ( - "missing colon", - b"GET /api HTTP/1.1\r\nX-Test value\r\nHost: x\r\n\r\n".as_slice(), - ), - ( - "whitespace before colon", - b"GET /api HTTP/1.1\r\nX-Test : value\r\nHost: x\r\n\r\n".as_slice(), - ), - ( - "invalid field-name token", - b"GET /api HTTP/1.1\r\nX@Test: value\r\nHost: x\r\n\r\n".as_slice(), - ), - ]; + async fn response_middleware_headers_only_mutates_head_and_preserves_framing() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\nCache-Control: public\r\n\r\nhello", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("cache-control: private\r\n"), + "{delivered}" + ); + assert!(delivered.contains("Content-Length: 5\r\n"), "{delivered}"); + assert!( + !delivered.to_ascii_lowercase().contains("transfer-encoding"), + "{delivered}" + ); + assert!(delivered.ends_with("\r\n\r\nhello"), "{delivered}"); + } - for (case, raw) in cases { - let (mut client, mut writer) = tokio::io::duplex(4096); - writer.write_all(raw).await.unwrap(); - let result = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), - ) - .await; - assert!(result.is_err(), "{case} must be rejected before forwarding"); - } + #[tokio::test] + async fn response_middleware_headers_only_preserves_chunked_and_close_delimited_bodies() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n2;ext=yes\r\nhe\r\n3\r\nllo\r\n0\r\n\r\n", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("Transfer-Encoding: chunked\r\n"), + "{delivered}" + ); + assert!( + delivered.ends_with("2;ext=yes\r\nhe\r\n3\r\nllo\r\n0\r\n\r\n"), + "{delivered}" + ); + + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\nhello", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + assert!( + String::from_utf8(delivered) + .unwrap() + .ends_with("\r\n\r\nhello") + ); } - /// SEC-009: Reject unsupported HTTP versions. #[tokio::test] - async fn reject_invalid_http_version() { - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /api JUNK/9.9\r\nHost: x\r\n\r\n") - .await - .unwrap(); - }); - let result = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_whole_body_delays_commit_and_sets_length() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\nETag: stale\r\n\r\nhello", + "GET", + ResponseRelayScript::WholeBody, ) .await; - assert!(result.is_err(), "Must reject unsupported HTTP version"); + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.contains("Content-Length: 11\r\n"), "{delivered}"); + assert!( + !delivered.to_ascii_lowercase().contains("etag:"), + "{delivered}" + ); + assert!(delivered.ends_with("\r\n\r\nwhole:hello"), "{delivered}"); } #[tokio::test] - async fn parse_http_request_canonicalizes_target_and_rewrites_raw_header() { - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /public/../secret HTTP/1.1\r\nHost: api.example.com\r\n\r\n") - .await - .unwrap(); - }); - let req = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_preflight_block_returns_canonical_403() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::BlockPreflight, ) - .await - .expect("request should parse") - .expect("request should exist"); - // Path fed to OPA evaluation is canonical. - assert_eq!(req.target, "/secret"); - // raw_header (forwarded byte-for-byte to upstream) is also canonical - // — this is the invariant the L7 canonicalization PR must uphold. - assert_eq!( - req.raw_header, b"GET /secret HTTP/1.1\r\nHost: api.example.com\r\n\r\n", - "outbound request line must carry the canonical path" + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.starts_with("HTTP/1.1 403 Forbidden\r\n"), + "{delivered}" ); + assert!( + delivered.contains("\"error\":\"middleware_denied\""), + "{delivered}" + ); + assert!( + delivered.contains("\"reason_code\":\"content_match\""), + "{delivered}" + ); + assert!(delivered.contains("Connection: close\r\n"), "{delivered}"); } #[tokio::test] - async fn parse_http_request_rejects_absolute_authority_mismatched_with_host() { - let (mut client, mut peer) = tokio::io::duplex(1024); - peer.write_all( - b"GET http://attacker.example.test/v1 HTTP/1.1\r\nHost: api.example.test\r\n\r\n", + async fn response_middleware_head_block_reports_length_without_body() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\n", + "HEAD", + ResponseRelayScript::BlockPreflight, ) - .await - .unwrap(); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let header_end = delivered + .windows(4) + .position(|window| window == b"\r\n\r\n") + .unwrap() + + 4; + let head = String::from_utf8(delivered[..header_end].to_vec()).unwrap(); + assert!(head.starts_with("HTTP/1.1 403 Forbidden\r\n"), "{head}"); + assert!(head.contains("Content-Length: "), "{head}"); + assert_eq!(delivered.len(), header_end); + } - let error = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + #[tokio::test] + async fn response_middleware_whole_body_block_returns_403_before_commit() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::BlockWholeBody, ) - .await - .expect_err("absolute-form authority mismatch must fail closed"); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); assert!( - error - .to_string() - .contains("request authority does not match the Host header"), - "{error}" + delivered.starts_with("HTTP/1.1 403 Forbidden\r\n"), + "{delivered}" ); + assert!(!delivered.contains("HTTP/1.1 200 OK"), "{delivered}"); } - #[test] - fn origin_form_targets_with_embedded_urls_use_host_authority() { - let host: http::uri::Authority = "api.example.test".parse().unwrap(); - - for target in ["/fetch/http://example.test", "/?next=http://example.test"] { - assert!( - absolute_form_uri(target).unwrap().is_none(), - "{target} must remain origin-form" - ); - validate_absolute_form_authority(target, Some(&host)) - .expect("embedded URL must not trigger absolute-form validation"); - - let raw = format!("GET {target} HTTP/1.1\r\nHost: api.example.test\r\n\r\n"); - let authority = request_authority(raw.as_bytes(), Some(443)) - .unwrap() - .expect("origin-form request with Host must have an authority"); - assert_eq!(authority.authority, host); - assert_eq!(authority.effective_port, 443); - } + #[tokio::test] + async fn response_middleware_stream_block_aborts_after_commit_without_error_bytes() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::BlockStream, + ) + .await; + assert!(outcome.is_err()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.1 200 OK\r\n"), "{delivered}"); + assert!(!delivered.contains("middleware_denied"), "{delivered}"); + assert!( + !delivered.contains("response_delivery_failed"), + "{delivered}" + ); } #[tokio::test] - async fn parse_http_request_keeps_embedded_url_in_origin_form_path() { - let (mut client, mut peer) = tokio::io::duplex(1024); - peer.write_all( - b"GET /fetch/http://example.test?next=http://other.test HTTP/1.1\r\nHost: api.example.test\r\n\r\n", + async fn response_middleware_whole_body_timeout_obeys_failure_policy() { + let response = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello"; + let (outcome, delivered) = run_response_middleware_relay_with_timeout( + response, + "GET", + ResponseRelayScript::SlowWholeBody, + openshell_supervisor_middleware::OnError::FailOpen, + std::time::Duration::from_millis(15), ) - .await - .unwrap(); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + assert!( + String::from_utf8(delivered) + .unwrap() + .ends_with("\r\n\r\nhello") + ); - let request = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + let (outcome, delivered) = run_response_middleware_relay_with_timeout( + response, + "GET", + ResponseRelayScript::SlowWholeBody, + openshell_supervisor_middleware::OnError::FailClosed, + std::time::Duration::from_millis(15), ) - .await - .expect("embedded URL origin-form request must parse") - .expect("request must be present"); - assert_eq!(request.target, "/fetch/http:/example.test"); - assert_eq!( - request.raw_header, - b"GET /fetch/http:/example.test?next=http://other.test HTTP/1.1\r\nHost: api.example.test\r\n\r\n", + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + assert!( + String::from_utf8(delivered) + .unwrap() + .contains("response_delivery_failed") ); } #[tokio::test] - async fn parse_http_request_canonicalization_preserves_query_string() { - let (mut client, mut writer) = tokio::io::duplex(4096); + async fn response_middleware_whole_body_timeout_does_not_reset_for_trickle_input() { + let (runner, chain) = response_middleware_fixture_with_error( + ResponseRelayScript::WholeBody, + openshell_supervisor_middleware::OnError::FailOpen, + ); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(16 * 1024); + let (mut client_read, mut client_write) = tokio::io::duplex(16 * 1024); tokio::spawn(async move { - writer - .write_all(b"GET /public/../v1/list?limit=10&sort=asc HTTP/1.1\r\nHost: h\r\n\r\n") + upstream_write + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nh") .await .unwrap(); + for byte in b"ello" { + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + upstream_write.write_all(&[*byte]).await.unwrap(); + } + upstream_write.shutdown().await.unwrap(); }); - let req = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + let mut middleware = response_middleware_context(&runner, &chain, "GET"); + middleware.whole_body_timeout = std::time::Duration::from_millis(20); + let outcome = relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(middleware), ) - .await - .unwrap() - .unwrap(); - assert_eq!(req.target, "/v1/list"); - assert_eq!( - req.raw_header, b"GET /v1/list?limit=10&sort=asc HTTP/1.1\r\nHost: h\r\n\r\n", - "canonical rewrite must preserve the query string verbatim" + .await; + drop(client_write); + let mut delivered = Vec::new(); + client_read.read_to_end(&mut delivered).await.unwrap(); + + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.ends_with("5\r\nhello\r\n0\r\n\r\n"), + "{delivered}" ); + assert!(!delivered.contains("whole:hello"), "{delivered}"); } #[tokio::test] - async fn parse_http_request_leaves_canonical_input_byte_for_byte() { - // When the input is already canonical, the raw_header must pass - // through unchanged — otherwise legitimate traffic pays a rewrite - // cost on every request. - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /api/v1/users HTTP/1.1\r\nHost: api.example.com\r\n\r\n") - .await - .unwrap(); - }); - let req = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_streams_normalized_chunks_and_preserves_trailers() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nTrailer: x-upstream\r\n\r\n2;ext=yes\r\nhe\r\n3\r\nllo\r\n0\r\nX-Upstream: kept\r\n\r\n", + "GET", + ResponseRelayScript::Stream, ) - .await - .unwrap() - .unwrap(); - assert_eq!(req.target, "/api/v1/users"); - assert_eq!( - req.raw_header, - b"GET /api/v1/users HTTP/1.1\r\nHost: api.example.com\r\n\r\n", - ); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.contains("Trailer: x-upstream\r\n"), "{delivered}"); + assert!(delivered.contains("5\r\nHELLO\r\n"), "{delivered}"); + assert!(delivered.contains("x-upstream: kept\r\n"), "{delivered}"); + assert!(!delivered.contains("digest:"), "{delivered}"); + assert!(!delivered.contains("ext=yes"), "{delivered}"); } #[tokio::test] - async fn parse_http_request_rejects_traversal_above_root() { - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /.. HTTP/1.1\r\nHost: h\r\n\r\n") - .await - .unwrap(); - }); - let result = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_rejects_malformed_response_fields_before_commit() { + for response in [ + b"HTTP/1.1 200 OK\r\nBad Name: value\r\nContent-Length: 0\r\n\r\n".as_slice(), + b"HTTP/1.1 200 OK\r\nTrailer: content-length\r\nTransfer-Encoding: chunked\r\n\r\n0\r\n\r\n".as_slice(), + b"HTTP/1.1 200 OK\r\nConnection: x-private\r\nTrailer: x-private\r\nTransfer-Encoding: chunked\r\n\r\n0\r\n\r\n".as_slice(), + ] { + let (outcome, delivered) = run_response_middleware_relay( + response, + "GET", + ResponseRelayScript::WholeBody, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.starts_with("HTTP/1.1 502 Bad Gateway\r\n"), + "{delivered}" + ); + } + } + + #[tokio::test] + async fn response_middleware_rejects_malformed_or_protected_upstream_trailers_atomically() { + for response in [ + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\nBad Name: value\r\n\r\n".as_slice(), + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\nContent-Length: 7\r\n\r\n".as_slice(), + ] { + let (outcome, delivered) = run_response_middleware_relay( + response, + "GET", + ResponseRelayScript::WholeBody, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.starts_with("HTTP/1.1 502 Bad Gateway\r\n"), + "{delivered}" + ); + assert!(!delivered.contains("whole:hello"), "{delivered}"); + } + } + + #[tokio::test] + async fn response_middleware_never_uses_chunked_framing_for_http_10() { + for (script, expected_body) in [ + (ResponseRelayScript::HeadersOnly, "hello"), + (ResponseRelayScript::Stream, "HELLO"), + (ResponseRelayScript::WholeBodyWithTrailer, "whole:hello"), + ] { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.0 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + script, + ) + .await; + assert!(outcome.is_ok()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.0 200 OK\r\n"), "{delivered}"); + assert!( + !delivered.to_ascii_lowercase().contains("transfer-encoding"), + "{delivered}" + ); + assert!( + !delivered.to_ascii_lowercase().contains("trailer:"), + "{delivered}" + ); + assert!(delivered.ends_with(expected_body), "{delivered}"); + } + } + + #[tokio::test] + async fn response_middleware_preserves_baseline_connection_outcomes() { + let (outcome, _) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\nConnection: close\r\n\r\nhello", + "GET", + ResponseRelayScript::HeadersOnly, ) .await; - assert!( - result.is_err(), - "a target that escapes the path root must be rejected at the parser" - ); + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\n\r\n", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + assert!(String::from_utf8(delivered).unwrap().ends_with("\r\n\r\n")); } #[tokio::test] - async fn parse_http_request_accepts_encoded_slash_when_endpoint_opts_in() { - // GitLab-style endpoints legitimately embed `%2F` in path segments - // (e.g. `/api/v4/projects/group%2Fproject`). Passing a provider - // constructed with `allow_encoded_slash: true` models the - // endpoint-config wiring that flows from `L7EndpointConfig`. - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /api/v4/projects/group%2Fproject HTTP/1.1\r\nHost: g\r\n\r\n") - .await - .unwrap(); - }); - let options = crate::l7::path::CanonicalizeOptions { - allow_encoded_slash: true, - ..Default::default() - }; - let req = parse_http_request(&mut client, &options) - .await - .unwrap() - .unwrap(); - assert_eq!(req.target, "/api/v4/projects/group%2Fproject"); + async fn response_middleware_forwards_interim_head_before_final_preflight() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 100 Continue\r\nX-Interim: yes\r\n\r\nHTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok", + "GET", + ResponseRelayScript::WholeBody, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.1 100 Continue\r\nX-Interim: yes\r\n\r\n")); + assert!(delivered.ends_with("whole:ok"), "{delivered}"); + } + + #[tokio::test] + async fn response_middleware_handles_bodyless_responses_without_body_events() { + for response in [ + b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\n\r\n".as_slice(), + b"HTTP/1.1 304 Not Modified\r\nContent-Length: 5\r\n\r\n".as_slice(), + ] { + let (outcome, delivered) = + run_response_middleware_relay(response, "GET", ResponseRelayScript::HeadersOnly) + .await; + assert!(outcome.is_ok()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("cache-control: private\r\n"), + "{delivered}" + ); + assert!(delivered.ends_with("\r\n\r\n"), "{delivered}"); + } + + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\n", + "HEAD", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(outcome.is_ok()); + let split = delivered + .windows(4) + .position(|window| window == b"\r\n\r\n") + .unwrap() + + 4; + assert_eq!(&delivered[split..], b""); } #[tokio::test] - async fn parse_http_request_rejects_encoded_slash_by_default() { - // Default strict options must reject `%2F` — this is the security - // posture for endpoints where an encoded slash would let an - // attacker disagree with the upstream on segment boundaries. - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /api/v4/projects/group%2Fproject HTTP/1.1\r\nHost: g\r\n\r\n") - .await - .unwrap(); - }); - let result = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_bypasses_protocol_upgrades() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n\x81\x02ok", + "GET", + ResponseRelayScript::HeadersOnly, ) .await; - assert!( - result.is_err(), - "default options must reject encoded slashes in the path" - ); + assert!(matches!( + outcome.unwrap(), + RelayOutcome::Upgraded { ref overflow, .. } if overflow == b"\x81\x02ok" + )); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(!delivered.contains("cache-control: private"), "{delivered}"); } #[tokio::test] - async fn parse_http_request_preserves_http_10_version_on_rewrite() { - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all(b"GET /a/./b HTTP/1.0\r\nHost: h\r\n\r\n") - .await - .unwrap(); - }); - let req = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_fail_closed_before_commit_returns_canonical_502() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::InvalidWholeBodySequence, ) - .await - .unwrap() - .unwrap(); - assert_eq!(req.target, "/a/b"); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); assert!( - req.raw_header.starts_with(b"GET /a/b HTTP/1.0\r\n"), - "rewrite must preserve the original HTTP version, got: {:?}", - String::from_utf8_lossy(&req.raw_header) + delivered.starts_with("HTTP/1.1 502 Bad Gateway\r\n"), + "{delivered}" ); + assert!(delivered.contains("\"error\":\"response_delivery_failed\"")); + assert!(delivered.contains( + "The upstream request may have completed, but OpenShell could not deliver its response. Retrying may repeat upstream side effects." + )); + assert!(!delivered.contains("invalid_body_sequence"), "{delivered}"); } #[tokio::test] - async fn parse_http_request_splits_path_and_query_params() { - let (mut client, mut writer) = tokio::io::duplex(4096); - tokio::spawn(async move { - writer - .write_all( - b"GET /download?slug=my%2Fskill&tag=foo&tag=bar HTTP/1.1\r\nHost: x\r\n\r\n", - ) - .await - .unwrap(); - }); - let req = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_head_failure_reports_body_length_without_body() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\n", + "HEAD", + ResponseRelayScript::WholeBody, ) - .await - .expect("request should parse") - .expect("request should exist"); - assert_eq!(req.target, "/download"); - assert_eq!( - req.query_params.get("slug").cloned(), - Some(vec!["my/skill".to_string()]) - ); - assert_eq!( - req.query_params.get("tag").cloned(), - Some(vec!["foo".to_string(), "bar".to_string()]) - ); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let split = delivered + .windows(4) + .position(|window| window == b"\r\n\r\n") + .unwrap() + + 4; + let head = String::from_utf8(delivered[..split].to_vec()).unwrap(); + assert!(head.starts_with("HTTP/1.1 502 Bad Gateway\r\n"), "{head}"); + assert!(head.contains("Content-Length: "), "{head}"); + assert_eq!(&delivered[split..], b""); } - /// Regression test: two pipelined requests in a single write must be - /// parsed independently. Before the fix, the 1024-byte `read()` buffer - /// could capture bytes from the second request, which were forwarded - /// upstream as body overflow of the first -- bypassing L7 policy checks. #[tokio::test] - async fn parse_http_request_does_not_overread_next_request() { - let (mut client, mut writer) = tokio::io::duplex(4096); - - tokio::spawn(async move { - writer - .write_all( - b"GET /allowed HTTP/1.1\r\nHost: example.com\r\n\r\n\ - POST /blocked HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n", - ) - .await - .unwrap(); - }); - - let first = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), - ) - .await - .expect("first request should parse") - .expect("expected first request"); - assert_eq!(first.action, "GET"); - assert_eq!(first.target, "/allowed"); - assert!(first.query_params.is_empty()); - assert_eq!( - first.raw_header, b"GET /allowed HTTP/1.1\r\nHost: example.com\r\n\r\n", - "raw_header must contain only the first request's headers" - ); - - let second = parse_http_request( - &mut client, - &crate::l7::path::CanonicalizeOptions::default(), + async fn response_middleware_fail_closed_after_commit_aborts_without_replacement() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::InvalidBodySequence, ) - .await - .expect("second request should parse") - .expect("expected second request"); - assert_eq!(second.action, "POST"); - assert_eq!(second.target, "/blocked"); - assert!(second.query_params.is_empty()); + .await; + assert!(outcome.is_err()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.1 200 OK\r\n"), "{delivered}"); + assert!(!delivered.contains("502 Bad Gateway"), "{delivered}"); } - #[test] - fn http_method_detection() { - assert!(looks_like_http(b"GET / HTTP/1.1\r\n")); - assert!(looks_like_http(b"POST /api HTTP/1.1\r\n")); - assert!(looks_like_http(b"DELETE /foo HTTP/1.1\r\n")); - assert!(could_be_http_request_prefix(b"GE")); - assert!(!could_be_http_request_prefix(b"GET ")); - assert!(!looks_like_http(b"\x00\x00\x00\x08")); // Postgres - assert!(!looks_like_http(HTTP2_PRIOR_KNOWLEDGE_PREFACE)); - assert!(!looks_like_http(b"HELLO")); // Unknown + #[tokio::test] + async fn response_middleware_fail_open_preserves_input_before_and_after_commit() { + for (script, expected_framing) in [ + ( + ResponseRelayScript::InvalidWholeBodySequence, + "Content-Length: 5\r\n", + ), + ( + ResponseRelayScript::InvalidBodySequence, + "Transfer-Encoding: chunked\r\n", + ), + ] { + let (outcome, delivered) = run_response_middleware_relay_with_error( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + script, + openshell_supervisor_middleware::OnError::FailOpen, + ) + .await; + assert!(outcome.is_ok()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.contains(expected_framing), "{delivered}"); + assert!(delivered.contains("hello"), "{delivered}"); + assert!(!delivered.contains("502 Bad Gateway"), "{delivered}"); + } } - #[test] - fn http2_prior_knowledge_detection() { - assert!(looks_like_http2_prior_knowledge( - HTTP2_PRIOR_KNOWLEDGE_PREFACE - )); - assert!(looks_like_http2_prior_knowledge( - &HTTP2_PRIOR_KNOWLEDGE_PREFACE[..8] - )); - assert!(could_be_http2_prior_knowledge_prefix(b"PRI * H")); - assert!(!looks_like_http2_prior_knowledge(b"PRI * H")); - assert!(!looks_like_http2_prior_knowledge(b"PRI / HTTP/1.1\r\n")); + #[tokio::test] + async fn response_middleware_stale_policy_generation_aborts_before_preflight() { + let policy_data = "network_policies: {}\n"; + let engine = OpaEngine::from_strings(TEST_POLICY, policy_data).unwrap(); + let guard = engine + .generation_guard(engine.current_generation()) + .unwrap(); + engine.reload(TEST_POLICY, policy_data).unwrap(); + let (runner, chain) = response_middleware_fixture(ResponseRelayScript::Stream); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); + let (mut client_read, mut client_write) = tokio::io::duplex(4096); + upstream_write + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello") + .await + .unwrap(); + upstream_write.shutdown().await.unwrap(); + let mut context = response_middleware_context(&runner, &chain, "GET"); + context.generation_guard = Some(&guard); + let outcome = relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(context), + ) + .await; + assert!(outcome.is_err()); + drop(client_write); + let mut delivered = Vec::new(); + client_read.read_to_end(&mut delivered).await.unwrap(); + assert!(delivered.is_empty()); } - #[test] - fn test_parse_status_code() { - assert_eq!( - parse_status_code("HTTP/1.1 200 OK\r\nHost: x\r\n\r\n"), - Some(200) - ); - assert_eq!( - parse_status_code("HTTP/1.1 204 No Content\r\n\r\n"), - Some(204) - ); - assert_eq!( - parse_status_code("HTTP/1.1 304 Not Modified\r\n\r\n"), - Some(304) - ); - assert_eq!( - parse_status_code("HTTP/1.1 100 Continue\r\n\r\n"), - Some(100) - ); - assert_eq!(parse_status_code(""), None); + #[tokio::test] + async fn response_middleware_client_disconnect_aborts_stream_delivery() { + let (runner, chain) = response_middleware_fixture(ResponseRelayScript::Stream); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); + let (client_read, mut client_write) = tokio::io::duplex(4096); + drop(client_read); + upstream_write + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello") + .await + .unwrap(); + upstream_write.shutdown().await.unwrap(); + let outcome = relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(response_middleware_context(&runner, &chain, "GET")), + ) + .await; + assert!(outcome.is_err()); } - #[test] - fn test_parse_connection_close() { - assert!(parse_connection_close( - "HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n" - )); - assert!(!parse_connection_close( - "HTTP/1.1 200 OK\r\nConnection: keep-alive\r\n\r\n" - )); - assert!(!parse_connection_close( - "HTTP/1.1 200 OK\r\nHost: x\r\n\r\n" - )); + #[tokio::test] + async fn response_middleware_streams_close_delimited_body_with_owned_framing() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\nhello", + "GET", + ResponseRelayScript::Stream, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("Transfer-Encoding: chunked\r\n"), + "{delivered}" + ); + assert!(delivered.contains("5\r\nHELLO\r\n"), "{delivered}"); } #[test] - fn test_response_is_event_stream() { - assert!(response_is_event_stream( - "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n" - )); - assert!(response_is_event_stream( - "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream; charset=utf-8\r\n\r\n" - )); - assert!(!response_is_event_stream( - "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n\r\n" - )); + fn response_middleware_ocsf_events_omit_content_headers_and_free_form_reasons() { + let target = HttpRequestTarget { + scheme: "https".into(), + host: "example.test".into(), + port: 443, + method: "GET".into(), + path: "/safe".into(), + query: String::new(), + }; + let events = http_response_middleware_invocation_events( + "policy", + &target, + 200, + &[openshell_supervisor_middleware::HttpResponseInvocation { + config_name: "scan".into(), + implementation: "example/scan".into(), + outcome: openshell_supervisor_middleware::HttpResponseInvocationOutcome::FailOpen, + sequence: Some(1), + input_size: 19, + output_size: None, + failed: true, + stage_disabled: true, + reason_code: Some("stable_reason".into()), + failure_category: Some("timeout".into()), + }], + ); + let json = events[0].to_json().unwrap().to_string(); + for forbidden in [ + "secret-response-body", + "authorization", + "content-length", + "middleware said secret", + "stable_reason", + ] { + assert!(!json.contains(forbidden), "{json}"); + } + assert!( + json.to_ascii_lowercase() + .contains("http_response_middleware"), + "{json}" + ); + assert!(json.contains("example/scan"), "{json}"); } #[test] - fn test_is_bodiless_response() { - assert!(is_bodiless_response("HEAD", 200)); - assert!(is_bodiless_response("GET", 100)); - assert!(is_bodiless_response("GET", 199)); - assert!(is_bodiless_response("GET", 204)); - assert!(is_bodiless_response("GET", 304)); - assert!(!is_bodiless_response("GET", 200)); - assert!(!is_bodiless_response("POST", 201)); + fn response_middleware_fail_open_dual_emits_sanitized_findings() { + let target = HttpRequestTarget { + scheme: "https".into(), + host: "example.test".into(), + port: 443, + method: "GET".into(), + path: "/safe".into(), + query: String::new(), + }; + for category in [ + "invalid_result", + "timeout", + "payload_capacity", + "session_capacity", + ] { + let invocation = openshell_supervisor_middleware::HttpResponseInvocation { + config_name: "scan".into(), + implementation: "example/scan".into(), + outcome: openshell_supervisor_middleware::HttpResponseInvocationOutcome::FailOpen, + sequence: Some(1), + input_size: 19, + output_size: None, + failed: true, + stage_disabled: true, + reason_code: None, + failure_category: Some(category.into()), + }; + assert_eq!( + http_response_middleware_invocation_events( + "policy", + &target, + 200, + std::slice::from_ref(&invocation), + ) + .len(), + 1 + ); + let finding = + http_response_middleware_fail_open_finding_event("policy", &target, &invocation) + .expect("fail-open failure must create a detection finding") + .to_json() + .unwrap() + .to_string(); + for expected in [ + "openshell.middleware.http_response_fail_open", + "example.test", + "pre_return", + category, + ] { + assert!(finding.contains(expected), "{finding}"); + } + assert!(!finding.contains("stable_reason"), "{finding}"); + } } #[tokio::test] @@ -5524,6 +8123,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5569,6 +8169,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5619,6 +8220,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5663,6 +8265,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5704,6 +8307,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5742,6 +8346,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5782,6 +8387,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5826,6 +8432,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5869,6 +8476,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5905,6 +8513,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5951,6 +8560,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5998,6 +8608,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -7949,4 +10560,35 @@ mod tests { SigV4PayloadMode::UnsignedPayload ); } + + #[test] + fn response_body_transform_strips_stale_integrity_headers() { + let mut headers = [ + "accept-ranges", + "etag", + "content-md5", + "digest", + "content-digest", + "repr-digest", + "signature", + "signature-input", + "content-type", + ] + .into_iter() + .map(|name| HttpHeader { + name: name.to_string(), + value: "value".to_string(), + }) + .collect(); + + strip_response_integrity_headers(&mut headers); + + assert_eq!( + headers, + vec![HttpHeader { + name: "content-type".to_string(), + value: "value".to_string(), + }] + ); + } } diff --git a/crates/openshell-supervisor-network/src/lib.rs b/crates/openshell-supervisor-network/src/lib.rs index 4fec48b300..592228bc79 100644 --- a/crates/openshell-supervisor-network/src/lib.rs +++ b/crates/openshell-supervisor-network/src/lib.rs @@ -8,6 +8,24 @@ //! owned by the orchestrator; this crate produces denials but does not //! aggregate them. +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::Duration; + +static HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS: AtomicU64 = + AtomicU64::new(openshell_core::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS); + +/// Configure the supervisor-wide wall-clock bound for whole-response buffering. +pub fn set_http_response_whole_body_timeout(timeout: Duration) { + let milliseconds = u64::try_from(timeout.as_millis()) + .unwrap_or(u64::MAX) + .max(1); + HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS.store(milliseconds, Ordering::Relaxed); +} + +pub(crate) fn http_response_whole_body_timeout() -> Duration { + Duration::from_millis(HTTP_RESPONSE_WHOLE_BODY_TIMEOUT_MS.load(Ordering::Relaxed)) +} + pub mod identity; pub mod inference_routes; pub mod l7; diff --git a/crates/openshell-supervisor-network/src/proxy.rs b/crates/openshell-supervisor-network/src/proxy.rs index 177d640fd8..e915ae9bdf 100644 --- a/crates/openshell-supervisor-network/src/proxy.rs +++ b/crates/openshell-supervisor-network/src/proxy.rs @@ -1272,6 +1272,7 @@ impl ForwardMiddlewarePipeline<'_> { request: crate::l7::provider::L7Request, client: &mut C, chain: Vec, + request_id: &str, ) -> Result where C: TokioAsyncRead + TokioAsyncWrite + Unpin + Send, @@ -1290,7 +1291,7 @@ impl ForwardMiddlewarePipeline<'_> { None => openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, }; - crate::l7::middleware::apply_middleware_chain_for_scheme( + crate::l7::middleware::apply_middleware_chain_for_scheme_with_request_id( request, client, self.ctx, @@ -1299,6 +1300,7 @@ impl ForwardMiddlewarePipeline<'_> { self.runner, self.generation_guard, transformed_body_policy, + request_id, ) .await } @@ -1751,7 +1753,7 @@ async fn handle_tcp_connection( let target = parts.next().unwrap_or(""); if method != "CONNECT" { - return handle_forward_proxy( + return Box::pin(handle_forward_proxy( method, target, &buf[..], @@ -1768,7 +1770,7 @@ async fn handle_tcp_connection( dynamic_credentials, denial_tx.as_ref(), activity_tx.as_ref(), - ) + )) .await; } @@ -4727,6 +4729,15 @@ struct ForwardRelayOptions<'a> { signing_region: &'a str, host: &'a str, port: u16, + response_middleware: Option>, +} + +struct ForwardResponseMiddleware<'a> { + ctx: &'a crate::l7::relay::L7EvalContext, + scheme: &'a str, + request_id: &'a str, + chain: &'a [openshell_supervisor_middleware::ChainEntry], + runner: &'a openshell_supervisor_middleware::ChainRunner, } async fn relay_rewritten_forward_request( @@ -4747,16 +4758,28 @@ where .map_or(rewritten.len(), |p| p + 4); let header_str = String::from_utf8_lossy(&rewritten[..header_end]); let body_length = crate::l7::rest::parse_body_length(&header_str)?; - let (_, query_params) = crate::l7::rest::parse_target_query(path)?; + let (request_path, query_params) = crate::l7::rest::parse_target_query(path)?; let req = crate::l7::provider::L7Request { action: method.to_string(), - target: path.to_string(), + target: request_path, query_params, raw_header: rewritten, body_length, }; - crate::l7::rest::relay_http_request_with_options_guarded( + let response_middleware = options.response_middleware.map(|middleware| { + crate::l7::relay::http_response_middleware_relay( + &req, + middleware.ctx, + middleware.scheme, + middleware.request_id, + middleware.chain, + middleware.runner, + Some(options.generation_guard), + ) + }); + + crate::l7::rest::relay_http_request_with_response_middleware_guarded( &req, client, upstream, @@ -4773,6 +4796,7 @@ where host: options.host, port: options.port, }, + response_middleware, ) .await } @@ -5699,7 +5723,9 @@ async fn handle_forward_proxy( .await?; return Ok(()); } + let request_id = uuid::Uuid::new_v4().to_string(); let websocket_chain = forward_websocket_request.then(|| chain.clone()); + let mut response_selection = None; if !chain.is_empty() { let middleware_runner = opa_engine.middleware_runner()?; let request = crate::l7::rest::request_from_buffered_http( @@ -5723,7 +5749,8 @@ async fn handle_forward_proxy( generation_guard: &forward_generation_guard, l7_reevaluation, }; - forward_request_bytes = match pipeline.apply(request, client, chain).await? { + response_selection = Some((chain.clone(), middleware_runner.clone())); + forward_request_bytes = match pipeline.apply(request, client, chain, &request_id).await? { crate::l7::middleware::MiddlewareApplyResult::Allowed(request) => request.raw_header, crate::l7::middleware::MiddlewareApplyResult::Denied { denial, .. } => { emit_activity_simple(activity_tx, true, "middleware"); @@ -6031,6 +6058,15 @@ async fn handle_forward_proxy( signing_region, host: &host_lc, port, + response_middleware: response_selection.as_ref().map(|(chain, runner)| { + ForwardResponseMiddleware { + ctx: &l7_ctx, + scheme: &scheme, + request_id: &request_id, + chain, + runner, + } + }), }, ) .await; @@ -6450,6 +6486,118 @@ mod tests { release: Arc, } + struct ForwardResponseHeadersMiddleware { + expected_path: String, + forbidden_path_fragment: String, + block: bool, + } + + #[tonic::async_trait] + impl openshell_core::middleware::InProcessMiddleware for ForwardResponseHeadersMiddleware { + async fn describe(&self) -> openshell_core::proto::MiddlewareManifest { + openshell_core::proto::MiddlewareManifest { + name: "test/forward-response".into(), + service_version: "test".into(), + bindings: vec![openshell_core::proto::MiddlewareBinding { + operation: openshell_core::proto::SupervisorMiddlewareOperation::HttpResponse + as i32, + phase: openshell_core::proto::SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: 8192, + timeout: String::new(), + }], + expected_audience: String::new(), + } + } + + async fn validate_config( + &self, + _middleware_name: &str, + _config: &prost_types::Struct, + ) -> Result<()> { + Ok(()) + } + + async fn evaluate_http_request( + &self, + _request: openshell_core::middleware::HttpRequestView<'_>, + ) -> Result { + Ok(openshell_core::proto::HttpRequestResult { + decision: openshell_core::proto::Decision::Allow as i32, + ..Default::default() + }) + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> std::result::Result + { + let (sender, receiver) = mpsc::channel(2); + let expected_path = self.expected_path.clone(); + let forbidden_path_fragment = self.forbidden_path_fragment.clone(); + let block = self.block; + tokio::spawn(async move { + while let Some(event) = requests.recv().await { + match event.event { + Some(openshell_core::proto::http_response_event::Event::Preflight( + preflight, + )) => { + let target = preflight.target.expect("response target"); + assert_eq!(target.path, expected_path); + assert!(!target.path.contains(&forbidden_path_fragment)); + let action = if block { + openshell_core::proto::http_response_preflight_result::Action::BlockDelivery( + openshell_core::proto::HttpResponseBlockDelivery {}, + ) + } else { + openshell_core::proto::http_response_preflight_result::Action::Inspect( + openshell_core::proto::HttpResponsePreflightInspect { + body_mode: openshell_core::proto::HttpResponseBodyMode::HeadersOnly as i32, + header_mutations: vec![openshell_core::proto::HeaderMutation { + operation: Some( + openshell_core::proto::header_mutation::Operation::Write( + openshell_core::proto::WriteHeader { + name: "x-forward-response-test".into(), + value: "selected".into(), + on_existing: openshell_core::proto::ExistingHeaderAction::Overwrite as i32, + }, + ), + ), + }], + }, + ) + }; + let result = openshell_core::proto::HttpResponseEventResult { + result: Some( + openshell_core::proto::http_response_event_result::Result::PreflightResult( + openshell_core::proto::HttpResponsePreflightResult { + action: Some(action), + reason_code: if block { + "query_guard".into() + } else { + String::new() + }, + ..Default::default() + }, + ), + ), + }; + if sender.send(Ok(result)).await.is_err() { + break; + } + } + Some(openshell_core::proto::http_response_event::Event::SessionEnd(_)) + | None => break, + Some(_) => panic!("headers-only response received an unexpected event"), + } + } + }); + Ok(Box::pin(tokio_stream::wrappers::ReceiverStream::new( + receiver, + ))) + } + } + #[tonic::async_trait] impl openshell_core::middleware::InProcessMiddleware for BlockingForwardMiddleware { async fn describe(&self) -> openshell_core::proto::MiddlewareManifest { @@ -6630,7 +6778,7 @@ network_policies: tokio::time::timeout( std::time::Duration::from_secs(30), - handle_forward_proxy( + Box::pin(handle_forward_proxy( "GET", &target, request.as_bytes(), @@ -6647,7 +6795,7 @@ network_policies: None, None, None, - ), + )), ) .await .expect("denied preflight must complete without an upstream response") @@ -6763,7 +6911,7 @@ network_policies: let (mut proxy_connection, _) = proxy_listener.accept().await.unwrap(); let handler = tokio::spawn(async move { - handle_forward_proxy( + Box::pin(handle_forward_proxy( "GET", &target, request.as_bytes(), @@ -6780,7 +6928,7 @@ network_policies: None, None, None, - ) + )) .await }); let scenario = tokio::time::timeout(std::time::Duration::from_secs(60), async { @@ -7684,7 +7832,7 @@ network_policies: let (_app, mut client) = tokio::io::duplex(8192); let outcome = pipeline - .apply(request, &mut client, chain) + .apply(request, &mut client, chain, "test-request-id") .await .expect("forward middleware pipeline"); @@ -7781,7 +7929,10 @@ network_policies: state.revoke_static_provider_environment(2); release.notify_one(); }; - let (outcome, ()) = tokio::join!(pipeline.apply(request, &mut client, chain), revoke); + let (outcome, ()) = tokio::join!( + pipeline.apply(request, &mut client, chain, "test-request-id"), + revoke + ); let request = match outcome.expect("middleware pipeline") { crate::l7::middleware::MiddlewareApplyResult::Allowed(request) => request, crate::l7::middleware::MiddlewareApplyResult::Denied { .. } => { @@ -7895,6 +8046,171 @@ network_policies: .unwrap() } + #[tokio::test] + async fn plaintext_forward_relay_applies_response_middleware() { + let guard = forward_test_guard(); + let ctx = crate::l7::relay::L7EvalContext { + host: "api.example.test".into(), + port: 80, + request_default_port: Some(80), + policy_name: "forward".into(), + binary_path: "/usr/bin/curl".into(), + ..Default::default() + }; + let runner = openshell_supervisor_middleware::ChainRunner::new(Arc::new( + ForwardResponseHeadersMiddleware { + expected_path: "/demo".into(), + forbidden_path_fragment: "not-present".into(), + block: false, + }, + )); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "response".into(), + implementation: "test/forward-response".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: openshell_supervisor_middleware::OnError::FailClosed, + }]; + let request = b"GET /demo HTTP/1.1\r\nHost: api.example.test\r\n\r\n".to_vec(); + let (mut proxy_to_upstream, mut upstream) = tokio::io::duplex(8192); + let (mut app, mut proxy_to_client) = tokio::io::duplex(8192); + let upstream_task = tokio::spawn(async move { + let mut request = vec![0; 1024]; + let size = upstream.read(&mut request).await.unwrap(); + assert!(request[..size].ends_with(b"\r\n\r\n")); + upstream + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok") + .await + .unwrap(); + }); + + let outcome = relay_rewritten_forward_request( + "GET", + "/demo", + request, + &mut proxy_to_client, + &mut proxy_to_upstream, + ForwardRelayOptions { + generation_guard: &guard, + credential_generation: None, + websocket_extensions: crate::l7::rest::WebSocketExtensionMode::Preserve, + secret_resolver: None, + request_body_credential_rewrite: false, + deny_uninspected_credentials: false, + credential_signing: crate::l7::CredentialSigning::None, + signing_service: "", + signing_region: "", + host: "api.example.test", + port: 80, + response_middleware: Some(ForwardResponseMiddleware { + ctx: &ctx, + scheme: "http", + request_id: "correlated-request-id", + chain: &chain, + runner: &runner, + }), + }, + ) + .await + .expect("plaintext forward relay"); + assert!(matches!( + outcome, + crate::l7::provider::RelayOutcome::Reusable + )); + upstream_task.await.unwrap(); + drop(proxy_to_client); + let mut response = Vec::new(); + app.read_to_end(&mut response).await.unwrap(); + let response = String::from_utf8(response).unwrap(); + assert!(response.contains("x-forward-response-test: selected\r\n")); + assert!(response.ends_with("\r\n\r\nok")); + } + + #[tokio::test] + async fn plaintext_forward_response_denial_never_echoes_query_secret() { + const SECRET: &str = "sk-forward-query-secret"; + let guard = forward_test_guard(); + let ctx = crate::l7::relay::L7EvalContext { + host: "api.example.test".into(), + port: 80, + request_default_port: Some(80), + policy_name: "forward".into(), + binary_path: "/usr/bin/curl".into(), + ..Default::default() + }; + let runner = openshell_supervisor_middleware::ChainRunner::new(Arc::new( + ForwardResponseHeadersMiddleware { + expected_path: "/demo".into(), + forbidden_path_fragment: SECRET.into(), + block: true, + }, + )); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "response".into(), + implementation: "test/forward-response".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: openshell_supervisor_middleware::OnError::FailClosed, + }]; + let target = format!("/demo?access_token={SECRET}"); + let request = + format!("GET {target} HTTP/1.1\r\nHost: api.example.test\r\n\r\n").into_bytes(); + let (mut proxy_to_upstream, mut upstream) = tokio::io::duplex(8192); + let (mut app, mut proxy_to_client) = tokio::io::duplex(8192); + let upstream_task = tokio::spawn(async move { + let mut request = vec![0; 1024]; + let size = upstream.read(&mut request).await.unwrap(); + assert!(request[..size].ends_with(b"\r\n\r\n")); + upstream + .write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok") + .await + .unwrap(); + }); + + let outcome = relay_rewritten_forward_request( + "GET", + &target, + request, + &mut proxy_to_client, + &mut proxy_to_upstream, + ForwardRelayOptions { + generation_guard: &guard, + credential_generation: None, + websocket_extensions: crate::l7::rest::WebSocketExtensionMode::Preserve, + secret_resolver: None, + request_body_credential_rewrite: false, + deny_uninspected_credentials: false, + credential_signing: crate::l7::CredentialSigning::None, + signing_service: "", + signing_region: "", + host: "api.example.test", + port: 80, + response_middleware: Some(ForwardResponseMiddleware { + ctx: &ctx, + scheme: "http", + request_id: "correlated-request-id", + chain: &chain, + runner: &runner, + }), + }, + ) + .await + .expect("plaintext forward response denial"); + assert!(matches!( + outcome, + crate::l7::provider::RelayOutcome::Consumed + )); + upstream_task.await.unwrap(); + drop(proxy_to_client); + let mut response = Vec::new(); + app.read_to_end(&mut response).await.unwrap(); + let response = String::from_utf8(response).unwrap(); + assert!(response.starts_with("HTTP/1.1 403 Forbidden\r\n")); + assert!(response.contains("\"path\":\"/demo\"")); + assert!(!response.contains(SECRET)); + assert!(!response.contains("access_token")); + } + async fn relay_forward_request_and_capture( method: &str, path: &str, @@ -7979,6 +8295,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await?; @@ -8244,6 +8561,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await?; @@ -11098,7 +11416,7 @@ network_policies: let (_app, mut client) = tokio::io::duplex(8192); let allowed = pipeline - .apply(request, &mut client, chain) + .apply(request, &mut client, chain, "test-request-id") .await .expect("middleware pipeline"); let crate::l7::middleware::MiddlewareApplyResult::Allowed(request) = allowed else { @@ -11400,6 +11718,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await @@ -11481,6 +11800,7 @@ network_policies: signing_region: "us-west-2", host: "api.example.com", port: 80, + response_middleware: None, }, ) .await @@ -11570,6 +11890,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await; @@ -11620,6 +11941,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await; diff --git a/docs/reference/gateway-config.mdx b/docs/reference/gateway-config.mdx index 1c846c29a8..8b1cb5c584 100644 --- a/docs/reference/gateway-config.mdx +++ b/docs/reference/gateway-config.mdx @@ -131,6 +131,9 @@ provider_profile_sources = [ # Operator-run supervisor middleware. The gRPC endpoint must be reachable from # both the gateway and sandbox supervisors. +[openshell.supervisor] +http_response_whole_body_timeout = "120s" + [[openshell.supervisor.middleware]] name = "local-content-guard" grpc_endpoint = "https://host.openshell.internal:50051" @@ -290,6 +293,11 @@ The gateway flushes buffered spans during shutdown, so spans from in-flight requ Register operator-run supervisor middleware services with one or more `[[openshell.supervisor.middleware]]` entries. Registration is static and operator-owned; changing it requires restarting the gateway. ```toml +[openshell.supervisor] +# One non-resetting wall-clock limit for response accumulation and whole-body +# middleware barriers. Accepts a positive integer followed by ms, s, or m. +http_response_whole_body_timeout = "120s" + [[openshell.supervisor.middleware]] name = "local-content-guard" grpc_endpoint = "https://host.openshell.internal:50051" @@ -299,13 +307,15 @@ max_payload_bytes = 262144 timeout = "500ms" ``` -Each service implements the supervisor middleware gRPC contract and exposes bindings through `Describe`. Policies reference the operator-owned registration `name`, attaching the complete middleware and all of its bindings. Bindings are identified by operation and phase. A manifest may expose at most one binding for each operation and phase pair. V1 supports `HttpRequest/pre_credentials` and `WebSocketMessage/pre_credentials`, so a service can inspect HTTP, WebSocket, or both. Registration names must be unique, and operator-run registrations cannot claim the reserved `openshell/` namespace. The service-reported manifest name is diagnostic metadata and does not need to match the registration name. +Each service implements the supervisor middleware gRPC contract and exposes bindings through `Describe`. Policies reference the operator-owned registration `name`, attaching the complete middleware and all of its bindings. Bindings are identified by operation and phase. A manifest may expose at most one binding for each operation and phase pair. V1 supports `HttpRequest/pre_credentials`, `HttpResponse/pre_return`, and `WebSocketMessage/pre_credentials`. Registration names must be unique, and operator-run registrations cannot claim the reserved `openshell/` namespace. The service-reported manifest name is diagnostic metadata and does not need to match the registration name. The gateway connects to every registered service and validates `Describe` before it starts. The service must therefore be running before the gateway. Policy creation and full policy updates call `ValidateConfig`; an unavailable service or invalid middleware configuration rejects the policy before persistence. -`max_payload_bytes` is the shared operator limit for inspectable logical payloads across every binding exposed by the service. It caps HTTP request and replacement bodies as well as complete WebSocket text messages and replacements. The value must be greater than zero, no larger than each binding's advertised `max_payload_bytes` capability, and no larger than the 4 MiB platform maximum. OpenShell rejects oversized values instead of silently clamping them. Binary WebSocket messages are not exposed to V1 middleware, so this field does not limit binary pass-through. Middleware gRPC servers should allow messages of at least 4 MiB plus 293 KiB so a maximum-size payload and its protobuf envelope fit on the transport. +`max_payload_bytes` is the shared operator limit for inspectable logical payloads across every binding exposed by the service. It caps HTTP request and response units, replacement bodies, and complete WebSocket text messages and replacements. Whole-response inspection uses it as the stage's total body limit. Streaming response inspection applies it to each self-contained unit, with a platform maximum of 64 KiB per input unit. The value must be greater than zero, no larger than each binding's advertised `max_payload_bytes` capability, and no larger than the 4 MiB platform maximum. OpenShell rejects oversized values instead of silently clamping them. Binary WebSocket messages are not exposed to V1 middleware, so this field does not limit binary pass-through. Middleware gRPC servers should allow messages of at least 4 MiB plus 293 KiB so a maximum-size payload and its protobuf envelope fit on the transport. + +`timeout` is the operator-configured service-wide RPC timeout. It accepts the same compact duration syntax as gateway interceptors: an integer followed by `ms` or `s`, such as `500ms` or `2s`. Values must be between `10ms` and `30s`, inclusive. Omit the field to use the 500 ms platform default. A binding may advertise a shorter `timeout` in the `Describe` manifest, but it cannot extend the operator-configured deadline; OpenShell uses the smaller value. OpenShell validates both levels before accepting the service. The operator-configured service timeout applies to `Describe` and `ValidateConfig`. The effective binding timeout applies to HTTP request evaluation, HTTP response preflight and unit exchanges, WebSocket preflight, and each WebSocket message. Accepted streaming protocols have no connection-wide RPC deadline. -`timeout` is the operator-configured service-wide RPC timeout. It accepts the same compact duration syntax as gateway interceptors: an integer followed by `ms` or `s`, such as `500ms` or `2s`. Values must be between `10ms` and `30s`, inclusive. Omit the field to use the 500 ms platform default. A binding may advertise a shorter `timeout` in the `Describe` manifest, but it cannot extend the operator-configured deadline; OpenShell uses the smaller value. OpenShell validates both levels before accepting the service. The operator-configured service timeout applies to `Describe` and `ValidateConfig`. The effective binding timeout applies only to `EvaluateHttpRequest`, WebSocket preflight, and each WebSocket message. An accepted WebSocket stream has no connection-wide RPC deadline. +`http_response_whole_body_timeout` is a supervisor-wide safety bound, not a middleware RPC timeout. It defaults to `120s` and accepts a positive integer followed by `ms`, `s`, or `m`. The timer starts after ordered response preflight selects a `WHOLE_BODY_BYTES` stage and before OpenShell reads the first response body byte. It does not reset as chunks arrive. It ends after the normalized body reaches end of stream and every whole-body barrier produces output. On expiry, each active whole-body stage follows its configured `on_error` policy: fail-open releases OpenShell-owned bytes through later stages, while fail-closed returns the canonical pre-commit 502 response. Changing this field requires restarting the gateway. The service `grpc_endpoint` supports plaintext `http://` and TLS `https://`. HTTPS uses the platform trust store unless `tls_ca_cert_path` names a certificate-only PEM bundle. OpenShell rejects bundles containing private keys, loads the certificates at gateway startup, and distributes only public certificates to sandbox supervisors; normal TLS hostname verification still applies. `audience` sets the exact audience for gateway-minted service tokens and defaults to `urn:openshell:extension:middleware:`. After authenticated `Describe` succeeds, OpenShell treats a non-empty manifest `expected_audience` as a consistency assertion and refuses to start when it differs from the configured audience. A strict verifier may reject an incorrect audience before returning the manifest. diff --git a/examples/supervisor-middleware-content-guard/Cargo.lock b/examples/supervisor-middleware-content-guard/Cargo.lock index f31d5be9b5..18bcbcb1b7 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.lock +++ b/examples/supervisor-middleware-content-guard/Cargo.lock @@ -206,6 +206,17 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527" +[[package]] +name = "chacha20" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" +dependencies = [ + "cfg-if", + "cpufeatures", + "rand_core", +] + [[package]] name = "clap" version = "4.6.1" @@ -252,6 +263,16 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" +[[package]] +name = "combine" +version = "4.6.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfc320937d09e6de266b31b9afb480f197d7a861be86be7cb2ea7e5d1bfffc5e" +dependencies = [ + "bytes", + "memchr", +] + [[package]] name = "core-foundation" version = "0.10.1" @@ -268,6 +289,27 @@ version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" +[[package]] +name = "cpufeatures" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ca28b0ae3115b884660db4118d803791fd6756b6e88f39c0f3f7859060d7566" +dependencies = [ + "libc", +] + +[[package]] +name = "critical-section" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "790eea4361631c5e7d22598ecd5723ff611904e3344ce8720784c93e3d83d40b" + +[[package]] +name = "data-encoding" +version = "2.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4583a4551df46e2792f82ceeac45e850d2e2d5debba0b91f102385cda5b11f06" + [[package]] name = "displaydoc" version = "0.2.6" @@ -399,6 +441,7 @@ dependencies = [ "cfg-if", "libc", "r-efi", + "rand_core", ] [[package]] @@ -453,6 +496,25 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hickory-proto" +version = "0.26.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e2da0694c15b44c6f68a6b05e0233617008c54080e31d6eb848d858a9c5b38d" +dependencies = [ + "data-encoding", + "idna", + "ipnet", + "jni", + "once_cell", + "rand", + "ring", + "thiserror", + "tinyvec", + "tracing", + "url", +] + [[package]] name = "http" version = "1.4.2" @@ -699,12 +761,71 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn", +] + [[package]] name = "libc" version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "libyml" +version = "0.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3302702afa434ffa30847a83305f0a69d6abd74293b6554c18ec85c7ef30c980" +dependencies = [ + "anyhow", + "version_check", +] + [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -832,6 +953,10 @@ name = "once_cell" version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +dependencies = [ + "critical-section", + "portable-atomic", +] [[package]] name = "once_cell_polyfill" @@ -854,6 +979,8 @@ dependencies = [ "prost-types", "protoc-bin-vendored", "rustix", + "rustls", + "rustls-pemfile", "serde", "serde_json", "thiserror", @@ -878,12 +1005,26 @@ dependencies = [ "tower", ] +[[package]] +name = "openshell-policy" +version = "0.0.0" +dependencies = [ + "hickory-proto", + "miette", + "openshell-core", + "prost-types", + "serde", + "serde_json", + "serde_yml", +] + [[package]] name = "openshell-supervisor-middleware-content-guard" version = "0.0.0" dependencies = [ "clap", "openshell-core", + "openshell-policy", "prost-types", "tokio", "tokio-stream", @@ -968,6 +1109,12 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "portable-atomic" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05c8b63e8d9609db387f0324918f81d68fe27748f084ef092fb35954d0539a85" + [[package]] name = "potential_utf" version = "0.1.5" @@ -1148,6 +1295,23 @@ version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "rand" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +dependencies = [ + "chacha20", + "getrandom 0.4.3", + "rand_core", +] + +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + [[package]] name = "redox_syscall" version = "0.5.18" @@ -1206,6 +1370,15 @@ version = "0.1.27" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b50b8869d9fc858ce7266cce0194bd74df58b9d0e3f6df3a9fc8eb470d95c09d" +[[package]] +name = "rustc_version" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" +dependencies = [ + "semver", +] + [[package]] name = "rustix" version = "1.1.4" @@ -1246,6 +1419,15 @@ dependencies = [ "security-framework", ] +[[package]] +name = "rustls-pemfile" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "rustls-pki-types" version = "1.14.1" @@ -1266,6 +1448,21 @@ dependencies = [ "untrusted", ] +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + [[package]] name = "schannel" version = "0.1.29" @@ -1304,6 +1501,12 @@ dependencies = [ "libc", ] +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + [[package]] name = "serde" version = "1.0.228" @@ -1347,6 +1550,21 @@ dependencies = [ "zmij", ] +[[package]] +name = "serde_yml" +version = "0.0.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59e2dd588bf1597a252c3b920e0143eb99b0f76e4e082f4c92ce34fbc9e71ddd" +dependencies = [ + "indexmap", + "itoa", + "libyml", + "memchr", + "ryu", + "serde", + "version_check", +] + [[package]] name = "shlex" version = "2.0.1" @@ -1363,6 +1581,22 @@ dependencies = [ "libc", ] +[[package]] +name = "simd_cesu8" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11031e251abf8611c80f460e19dbdeb54a66db918e49c65a7065b46ac7aec520" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "slab" version = "0.4.12" @@ -1515,6 +1749,21 @@ dependencies = [ "zerovec", ] +[[package]] +name = "tinyvec" +version = "1.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cf0ded5c4e56918d8f8a339e1bb67d038d3bc6d144ac407904015ba2e4cde9b" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + [[package]] name = "tokio" version = "1.52.3" @@ -1775,6 +2024,22 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + [[package]] name = "want" version = "0.3.1" @@ -1790,6 +2055,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "windows-link" version = "0.2.1" diff --git a/examples/supervisor-middleware-content-guard/Cargo.toml b/examples/supervisor-middleware-content-guard/Cargo.toml index ef0d1f47c3..9a9f98381a 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.toml +++ b/examples/supervisor-middleware-content-guard/Cargo.toml @@ -20,6 +20,9 @@ tokio = { version = "1.43", features = ["macros", "rt-multi-thread"] } tokio-stream = "0.1" tonic = { version = "0.14", features = ["transport"] } +[dev-dependencies] +openshell-policy = { path = "../../crates/openshell-policy" } + [[bin]] name = "supervisor-middleware-content-guard" path = "src/main.rs" diff --git a/examples/supervisor-middleware-content-guard/README.md b/examples/supervisor-middleware-content-guard/README.md index 53eda94210..21b2ae87f7 100644 --- a/examples/supervisor-middleware-content-guard/README.md +++ b/examples/supervisor-middleware-content-guard/README.md @@ -8,18 +8,18 @@ SPDX-License-Identifier: Apache-2.0 > [!WARNING] > Supervisor middleware is a research preview. Its policy and service contracts may change without compatibility guarantees. Use it only to prototype and evaluate middleware integrations. -This example implements an operator-run supervisor middleware service. It scans UTF-8 HTTP request bodies and complete client-to-upstream WebSocket text messages for configured literal strings, then either replaces every match or denies the request or message. Findings report only aggregate counts and never include configured terms or inspected content. +This example implements request, response, and client WebSocket bindings in one operator-run supervisor middleware service. The response binding demonstrates header-only inspection, whole-body and streaming transforms, trailer mutation, and block delivery. > [!WARNING] > This intentionally simple implementation demonstrates the supervisor middleware service contract. It is not a complete or reliable content guard and must not be used as a security control. It handles only UTF-8 HTTP request bodies and WebSocket text messages with case-sensitive literal terms, merges overlapping literal match ranges before redaction, and does not address encodings, transformations, normalization, binary WebSocket messages, upstream-to-client messages, or adversarial inputs that a production content guard must handle. ## Prerequisites -Install `cargo`, `curl`, `jq`, and `openssl` on the host before running the smoke script. +Install `cargo`, `curl`, `jq`, `openssl`, and Python 3 on the host before running the smoke script. ## Run the smoke example -Run the end-to-end smoke suite to build and start a local gateway, start the content-guard service, create a sandbox, and send the same request body to two destinations: +Run the end-to-end smoke suite to build a local gateway and sandbox supervisor, start the content-guard service, create a sandbox, and send the same request body to two destinations: ```shell ./examples/supervisor-middleware-content-guard/smoke.sh --test-suite @@ -39,6 +39,8 @@ The script creates the sandbox and prints the guarded and unguarded request comm CONTENT_GUARD_SMOKE_HOST=192.168.1.10 ./examples/supervisor-middleware-content-guard/smoke.sh --test-suite ``` +The gateway auto-detects its compute driver. Set `CONTENT_GUARD_SMOKE_DRIVER=docker` or `CONTENT_GUARD_SMOKE_DRIVER=podman` if more than one local runtime is installed and auto-detection selects the wrong one. + ## Run manually Start the service before starting the gateway. Bind to all host interfaces so a local containerized gateway and sandbox supervisor can reach it: @@ -54,6 +56,7 @@ Add the service registration to your local gateway TOML: [[openshell.supervisor.middleware]] name = "content-guard-example" grpc_endpoint = "http://host.openshell.internal:50051" +allow_insecure_transport = true max_payload_bytes = 262144 timeout = "500ms" ``` @@ -84,6 +87,34 @@ curl -sS https://httpbin.org/anything \ The echoed JSON body contains `[FILTERED]` instead of the configured term. +## HTTP response behavior + +Start the included raw HTTP upstream in another terminal: + +```shell +python3 examples/supervisor-middleware-content-guard/upstream.py +``` + +From the sandbox, exercise the response protocol: + +```shell +curl -i http://host.openshell.internal:18081/headers-only +curl -i http://host.openshell.internal:18081/whole-body +curl -i --raw http://host.openshell.internal:18081/stream +curl -i --raw http://host.openshell.internal:18081/stream-close +curl -i http://host.openshell.internal:18081/block +``` + +| Path | Mode | Result | +| --- | --- | --- | +| `/headers-only` | `HEADERS_ONLY` | Adds `x-example-response-mode` without changing content-length framing. | +| `/whole-body` | `WHOLE_BODY_BYTES` | Prefixes the normalized body with `[whole]`. | +| `/stream` | `STREAM_BYTES` | Uppercases normalized units and changes the existing `x-example-body-bytes` trailer to `11`. | +| `/stream-close` | `STREAM_BYTES` | Uppercases a close-delimited `text/event-stream` response. | +| `/block` | `WHOLE_BODY_BYTES` | Returns OpenShell's canonical 403 response with reason code `content_match`. | + +The upstream deliberately uses content-length, chunked, and close-delimited responses. OpenShell normalizes transport framing only for body-processing modes and validates trailer changes against the names supplied by the upstream. + ## WebSocket behavior For a selected WebSocket upgrade, the service accepts preflight, waits for the session-start notification, and evaluates each complete client-to-upstream text message. Redact mode returns a replacement message, while deny mode returns `content_match` and OpenShell closes the session according to middleware policy. Session-start and session-end events are notifications and do not produce results. @@ -107,4 +138,4 @@ config: - prototype-secret ``` -The implementation supports `HTTP_REQUEST/PRE_CREDENTIALS` and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`, advertises a 256 KiB limit for each operation, and inherits the service-wide RPC timeout. The gateway registration's `max_payload_bytes` may set a smaller shared limit. A binding can advertise a shorter timeout, but it cannot extend the operator-configured timeout. +The implementation supports `HTTP_REQUEST/PRE_CREDENTIALS`, `HTTP_RESPONSE/PRE_RETURN`, and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. It advertises a 256 KiB limit for each operation and inherits the service-wide RPC timeout. The gateway registration's `max_payload_bytes` may set a smaller shared limit. A binding can advertise a shorter timeout, but it cannot extend the operator-configured timeout. diff --git a/examples/supervisor-middleware-content-guard/policy.yaml b/examples/supervisor-middleware-content-guard/policy.yaml index ff3d9ef89e..b72af7ac2f 100644 --- a/examples/supervisor-middleware-content-guard/policy.yaml +++ b/examples/supervisor-middleware-content-guard/policy.yaml @@ -18,6 +18,7 @@ network_middlewares: endpoints: include: - httpbin.org + - host.openshell.internal network_policies: httpbin: @@ -44,3 +45,27 @@ network_policies: path: /anything binaries: - path: /usr/bin/curl + response-framing-demo: + name: Response framing demo + endpoints: + - host: host.openshell.internal + port: 18081 + protocol: rest + rules: + - allow: + method: GET + path: /headers-only + - allow: + method: GET + path: /whole-body + - allow: + method: GET + path: /stream + - allow: + method: GET + path: /stream-close + - allow: + method: GET + path: /block + binaries: + - path: /usr/bin/curl diff --git a/examples/supervisor-middleware-content-guard/smoke.sh b/examples/supervisor-middleware-content-guard/smoke.sh index b30475c8ca..106d544891 100755 --- a/examples/supervisor-middleware-content-guard/smoke.sh +++ b/examples/supervisor-middleware-content-guard/smoke.sh @@ -24,6 +24,8 @@ Options: Environment: CONTENT_GUARD_SMOKE_HOST Non-loopback host address reachable from both the gateway and sandbox containers. + CONTENT_GUARD_SMOKE_DRIVER + Optional compute driver name, such as docker or podman. EOF } @@ -103,6 +105,7 @@ detect_service_host() { } SERVICE_HOST="$(detect_service_host)" +COMPUTE_DRIVER="${CONTENT_GUARD_SMOKE_DRIVER:-}" if [[ "$SERVICE_HOST" == "localhost" || "$SERVICE_HOST" == "::1" || "$SERVICE_HOST" == 127.* || "$SERVICE_HOST" == *:* ]]; then echo "CONTENT_GUARD_SMOKE_HOST must be a non-loopback IPv4 address: $SERVICE_HOST" >&2 exit 1 @@ -115,6 +118,8 @@ GATEWAY_CONFIG="$TMPDIR/gateway.toml" SETUP_LOG="$LOG_DIR/setup.log" GATEWAY_LOG="$LOG_DIR/gateway.log" MIDDLEWARE_LOG="$LOG_DIR/middleware.log" +UPSTREAM_LOG="$LOG_DIR/upstream.log" +SANDBOX_LOG="$LOG_DIR/sandbox.log" RUN_ID="content-guard-smoke-$$-$RANDOM" # Sandbox names are capped at 19 characters. Use a short prefix with # the PID for uniqueness; keep the full RUN_ID for gateway identity. @@ -141,6 +146,11 @@ cleanup() { wait "$MIDDLEWARE_PID" 2>/dev/null || true fi + if [[ -n "${UPSTREAM_PID:-}" ]]; then + kill "$UPSTREAM_PID" 2>/dev/null || true + wait "$UPSTREAM_PID" 2>/dev/null || true + fi + if [[ "$status" -eq 0 ]]; then rm -rf "$TMPDIR" else @@ -214,8 +224,12 @@ ttl_secs = 0 [[openshell.supervisor.middleware]] name = "content-guard-example" grpc_endpoint = "http://$SERVICE_HOST:$MIDDLEWARE_PORT" +allow_insecure_transport = true max_payload_bytes = 262144 timeout = "500ms" + +[openshell.drivers.docker] +supervisor_bin = "$ROOT/target/debug/openshell-sandbox" EOF } @@ -239,11 +253,13 @@ generate_gateway_jwt_bundle() { dump_logs() { local label path - for label in setup gateway middleware; do + for label in setup gateway middleware upstream sandbox; do case "$label" in setup) path="$SETUP_LOG" ;; gateway) path="$GATEWAY_LOG" ;; middleware) path="$MIDDLEWARE_LOG" ;; + upstream) path="$UPSTREAM_LOG" ;; + sandbox) path="$SANDBOX_LOG" ;; esac printf '\n--- %s log: %s ---\n' "$label" "$path" >&2 if [[ -f "$path" ]]; then @@ -254,8 +270,23 @@ dump_logs() { done } +capture_sandbox_log() { + local container_id + + if [[ "$SANDBOX_CREATED" -ne 1 || "$COMPUTE_DRIVER" != "docker" ]] || + ! command -v docker >/dev/null 2>&1; then + return + fi + + container_id="$(docker ps -aq --filter "name=$SANDBOX_NAME" | head -n 1)" + if [[ -n "$container_id" ]]; then + docker logs "$container_id" >"$SANDBOX_LOG" 2>&1 || true + fi +} + fail() { printf 'FAIL %s\n' "$1" >&2 + capture_sandbox_log dump_logs exit 1 } @@ -316,15 +347,40 @@ wait_for_middleware() { fail "content guard service is reachable at $SERVICE_HOST:$MIDDLEWARE_PORT" } +start_upstream() { + printf 'INFO starting response framing upstream at %s:18081\n' "$SERVICE_HOST" + python3 "$EXAMPLE_DIR/upstream.py" >"$UPSTREAM_LOG" 2>&1 & + UPSTREAM_PID=$! +} + +wait_for_upstream() { + for _ in {1..30}; do + if ! kill -0 "$UPSTREAM_PID" 2>/dev/null; then + fail "response framing upstream starts" + fi + if curl -fsS --max-time 1 "http://127.0.0.1:18081/headers-only" >/dev/null 2>&1; then + printf 'INFO response framing upstream is ready\n' + return + fi + sleep 1 + done + fail "response framing upstream is reachable" +} + start_gateway() { + local -a driver_args=() + if [[ -n "$COMPUTE_DRIVER" ]]; then + driver_args=(--drivers "$COMPUTE_DRIVER") + fi printf 'INFO starting gateway\n' env -u OPENSHELL_DRIVERS "$GATEWAY_BIN" \ + "${driver_args[@]}" \ --config "$GATEWAY_CONFIG" \ --bind-address 127.0.0.1 \ --port "$GATEWAY_PORT" \ --health-port "$HEALTH_PORT" \ --metrics-port 0 \ - --log-level info \ + --log-level "${CONTENT_GUARD_SMOKE_LOG_LEVEL:-info}" \ --disable-tls \ --db-url "sqlite://$TMPDIR/gateway.db" >"$GATEWAY_LOG" 2>&1 & GATEWAY_PID=$! @@ -354,10 +410,10 @@ create_sandbox() { "$CLI_BIN" --gateway-endpoint "$GATEWAY_ENDPOINT" ) + SANDBOX_CREATED=1 run_setup_step \ "creating content guard sandbox" \ - "${CLI[@]}" sandbox create --name "$SANDBOX_NAME" --policy "$EXAMPLE_DIR/policy.yaml" --keep --no-tty -- /bin/sh -lc true - SANDBOX_CREATED=1 + "${CLI[@]}" sandbox create --name "$SANDBOX_NAME" --policy "$EXAMPLE_DIR/policy.yaml" --no-tty --detach -- sleep infinity } request() { @@ -368,14 +424,48 @@ request() { --data '{"note":"prototype-secret"}' } +response_request() { + local path="$1" + "${CLI[@]}" sandbox exec --name "$SANDBOX_NAME" --no-tty -- \ + curl -sS -i --raw --max-time 20 "http://host.openshell.internal:18081/$path" +} + run_suite() { local guarded_output="$LOG_DIR/guarded.out" local unguarded_output="$LOG_DIR/unguarded.out" + local response_output="$LOG_DIR/response.out" printf 'INFO sending guarded request to httpbin.org\n' if ! request httpbin.org >"$guarded_output" 2>>"$SETUP_LOG"; then fail "guarded request completes" fi + + printf 'INFO exercising HTTP response middleware modes\n' + if ! response_request headers-only >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fiq 'x-example-response-mode: headers-only' "$response_output" || + ! grep -Fq 'headers-only' "$response_output"; then + fail "headers-only response middleware" + fi + if ! response_request whole-body >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq '[whole] whole body' "$response_output"; then + fail "whole-body response middleware" + fi + if ! response_request stream >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'STREAM BODY' "$response_output" || + ! grep -Fiq 'x-example-body-bytes: 11' "$response_output"; then + fail "stream response middleware with trailer mutation" + fi + if ! response_request stream-close >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'DATA: STREAM CLOSE' "$response_output"; then + fail "close-delimited SSE response middleware" + fi + if ! response_request block >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'HTTP/1.1 403 Forbidden' "$response_output" || + ! grep -Fq 'middleware_denied' "$response_output" || + ! grep -Fq 'content_match' "$response_output"; then + fail "response middleware block" + fi + printf 'PASS HTTP response middleware modes\n' if grep -Fq '[FILTERED]' "$guarded_output" && ! grep -Fq 'prototype-secret' "$guarded_output"; then printf 'PASS guarded request is filtered\n' else @@ -439,15 +529,19 @@ require_command cargo require_command curl require_command jq require_command openssl +require_command python3 ROOT_TARGET_DIR="$(cargo_target_dir "$ROOT/Cargo.toml")" EXAMPLE_TARGET_DIR="$(cargo_target_dir "$EXAMPLE_DIR/Cargo.toml")" GATEWAY_BIN="$ROOT_TARGET_DIR/debug/openshell-gateway" CLI_BIN="$ROOT_TARGET_DIR/debug/openshell" MIDDLEWARE_BIN="$EXAMPLE_TARGET_DIR/debug/supervisor-middleware-content-guard" run_setup_step "building gateway" cargo build --quiet -p openshell-gateway --bin openshell-gateway +run_setup_step "building sandbox supervisor" cargo build --quiet -p openshell-sandbox --bin openshell-sandbox run_setup_step "building content guard" cargo build --quiet --manifest-path "$EXAMPLE_DIR/Cargo.toml" run_setup_step "building CLI" cargo build --quiet -p openshell-cli --bin openshell generate_gateway_jwt_bundle +start_upstream +wait_for_upstream start_middleware wait_for_middleware start_gateway diff --git a/examples/supervisor-middleware-content-guard/src/main.rs b/examples/supervisor-middleware-content-guard/src/main.rs index 8d714264e7..effce70a5a 100644 --- a/examples/supervisor-middleware-content-guard/src/main.rs +++ b/examples/supervisor-middleware-content-guard/src/main.rs @@ -6,20 +6,31 @@ use std::net::SocketAddr; use std::ops::Range; use clap::Parser; -use openshell_core::middleware::WebSocketResponseStream; +use openshell_core::middleware::{HttpResponseResultStream, WebSocketResponseStream}; +use openshell_core::proto::middleware::v1::http_response_pre_return_server::{ + HttpResponsePreReturn, HttpResponsePreReturnServer, +}; use openshell_core::proto::middleware::v1::supervisor_middleware_server::{ SupervisorMiddleware, SupervisorMiddlewareServer, }; use openshell_core::proto::{ - Decision, Finding, HttpRequestEvaluation, HttpRequestResult, MiddlewareBinding, - MiddlewareManifest, SupervisorMiddlewareOperation, SupervisorMiddlewarePhase, - ValidateConfigRequest, ValidateConfigResponse, WebSocketMessage, WebSocketMessageResult, - WebSocketPreflightAction, WebSocketPreflightDecision, WebSocketSessionEvent, - WebSocketSessionEventResult, web_socket_message, web_socket_message_result, + Decision, ExistingHeaderAction, Finding, HeaderMutation, HttpRequestEvaluation, + HttpRequestResult, HttpResponseBlockDelivery, HttpResponseBodyMode, HttpResponseBodyResult, + HttpResponseBodyTransform, HttpResponseEvent, HttpResponseEventResult, + HttpResponsePreflightInspect, HttpResponsePreflightResult, HttpResponsePreflightSkip, + HttpResponseTrailersResult, MiddlewareBinding, MiddlewareManifest, + SupervisorMiddlewareOperation, SupervisorMiddlewarePhase, ValidateConfigRequest, + ValidateConfigResponse, WebSocketMessage, WebSocketMessageResult, WebSocketPreflightAction, + WebSocketPreflightDecision, WebSocketSessionEvent, WebSocketSessionEventResult, WriteHeader, + header_mutation, http_response_body_result, http_response_body_transform, + http_response_body_unit, http_response_event, http_response_event_result, + http_response_preflight_result, web_socket_message, web_socket_message_result, web_socket_session_event, web_socket_session_event_result, }; use prost_types::Struct; use prost_types::value::Kind; +use tokio::sync::mpsc; +use tokio_stream::wrappers::ReceiverStream; use tokio_stream::{Stream, StreamExt}; use tonic::transport::Server; use tonic::{Request, Response, Status}; @@ -238,6 +249,12 @@ impl SupervisorMiddleware for ContentGuard { max_payload_bytes: MAX_PAYLOAD_BYTES, timeout: String::new(), }, + MiddlewareBinding { + operation: SupervisorMiddlewareOperation::HttpResponse as i32, + phase: SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: MAX_PAYLOAD_BYTES, + timeout: String::new(), + }, ], expected_audience: String::new(), })) @@ -282,6 +299,222 @@ impl SupervisorMiddleware for ContentGuard { } } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ResponseMode { + HeadersOnly, + WholeBody, + Stream, + StreamClose, + Block, +} + +#[derive(Debug, Default)] +struct ResponseSessionState { + selected: Option, + next_sequence: u64, + body_ended: bool, +} + +impl ResponseSessionState { + fn preflight( + &mut self, + preflight: openshell_core::proto::HttpResponsePreflight, + ) -> Result { + if self.selected.is_some() { + return Err(Status::failed_precondition("duplicate response preflight")); + } + GuardConfig::parse(preflight.config.as_ref()).map_err(Status::invalid_argument)?; + let path = preflight + .target + .as_ref() + .map(|target| target.path.as_str()) + .unwrap_or_default(); + let Some(selected) = response_mode_for_path(path) else { + return Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some(http_response_preflight_result::Action::Skip( + HttpResponsePreflightSkip {}, + )), + reason_code: "path_not_selected".into(), + ..Default::default() + }, + )), + }); + }; + self.selected = Some(selected); + self.next_sequence = 1; + let body_mode = match selected { + ResponseMode::HeadersOnly => HttpResponseBodyMode::HeadersOnly, + ResponseMode::WholeBody | ResponseMode::Block => HttpResponseBodyMode::WholeBodyBytes, + ResponseMode::Stream | ResponseMode::StreamClose => HttpResponseBodyMode::StreamBytes, + }; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some(http_response_preflight_result::Action::Inspect( + HttpResponsePreflightInspect { + body_mode: body_mode as i32, + header_mutations: vec![write_header( + "x-example-response-mode", + match selected { + ResponseMode::HeadersOnly => "headers-only", + ResponseMode::WholeBody => "whole-body", + ResponseMode::Stream => "stream", + ResponseMode::StreamClose => "stream-close", + ResponseMode::Block => "block", + }, + )], + }, + )), + ..Default::default() + }, + )), + }) + } + + fn body( + &mut self, + body: openshell_core::proto::HttpResponseBodyUnit, + ) -> Result { + let selected = self + .selected + .ok_or_else(|| Status::failed_precondition("body arrived before preflight"))?; + if selected == ResponseMode::HeadersOnly || self.body_ended { + return Err(Status::failed_precondition( + "body event is invalid for the response session state", + )); + } + if body.sequence != self.next_sequence { + return Err(Status::invalid_argument( + "unexpected response body sequence", + )); + } + self.next_sequence = self.next_sequence.saturating_add(1); + self.body_ended = body.end_of_stream; + let Some(http_response_body_unit::Payload::Data(data)) = body.payload else { + return Err(Status::invalid_argument("body data is required")); + }; + let (action, reason_code) = match selected { + ResponseMode::WholeBody => ( + http_response_body_result::Action::Transform(HttpResponseBodyTransform { + replacement: Some(http_response_body_transform::Replacement::Data( + [b"[whole] ".as_slice(), &data].concat(), + )), + }), + String::new(), + ), + ResponseMode::Stream | ResponseMode::StreamClose => ( + http_response_body_result::Action::Transform(HttpResponseBodyTransform { + replacement: Some(http_response_body_transform::Replacement::Data( + data.to_ascii_uppercase(), + )), + }), + String::new(), + ), + ResponseMode::Block => ( + http_response_body_result::Action::BlockDelivery(HttpResponseBlockDelivery {}), + "content_match".into(), + ), + ResponseMode::HeadersOnly => unreachable!(), + }; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: body.sequence, + action: Some(action), + reason_code, + ..Default::default() + }, + )), + }) + } + + fn trailers(&self) -> Result { + if !self.body_ended { + return Err(Status::failed_precondition( + "trailers arrived before the final body result", + )); + } + let trailer_mutations = if self.selected == Some(ResponseMode::Stream) { + vec![write_header("x-example-body-bytes", "11")] + } else { + Vec::new() + }; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult { + trailer_mutations, + ..Default::default() + }, + )), + }) + } +} + +fn response_mode_for_path(path: &str) -> Option { + match path { + "/headers-only" => Some(ResponseMode::HeadersOnly), + "/whole-body" => Some(ResponseMode::WholeBody), + "/stream" => Some(ResponseMode::Stream), + "/stream-close" => Some(ResponseMode::StreamClose), + "/block" => Some(ResponseMode::Block), + _ => None, + } +} + +fn write_header(name: &str, value: &str) -> HeaderMutation { + HeaderMutation { + operation: Some(header_mutation::Operation::Write(WriteHeader { + name: name.into(), + value: value.into(), + on_existing: ExistingHeaderAction::Overwrite as i32, + })), + } +} + +#[tonic::async_trait] +impl HttpResponsePreReturn for ContentGuard { + type EvaluateStream = HttpResponseResultStream; + + async fn evaluate( + &self, + request: Request>, + ) -> Result, Status> { + let mut events = request.into_inner(); + let (sender, receiver) = mpsc::channel(4); + tokio::spawn(async move { + let mut state = ResponseSessionState::default(); + while let Some(event) = events.next().await { + let result = match event { + Ok(event) => match event.event { + Some(http_response_event::Event::Preflight(preflight)) => { + state.preflight(preflight) + } + Some(http_response_event::Event::Body(body)) => state.body(body), + Some(http_response_event::Event::Trailers(_)) => state.trailers(), + Some(http_response_event::Event::SessionEnd(_)) => break, + None => Err(Status::invalid_argument("response event is required")), + }, + Err(error) => Err(error), + }; + match result { + Ok(result) => { + if sender.send(Ok(result)).await.is_err() { + break; + } + } + Err(error) => { + let _ = sender.send(Err(error)).await; + break; + } + } + } + }); + Ok(Response::new(Box::pin(ReceiverStream::new(receiver)))) + } +} + fn validate_phase(phase: i32) -> Result<(), String> { if phase != PHASE as i32 { return Err(format!("unsupported phase '{phase}'")); @@ -470,6 +703,7 @@ async fn main() -> Result<(), Box> { println!("serving {MANIFEST_NAME} on http://{}", cli.bind); Server::builder() .add_service(SupervisorMiddlewareServer::new(ContentGuard)) + .add_service(HttpResponsePreReturnServer::new(ContentGuard)) .serve(cli.bind) .await?; Ok(()) @@ -478,7 +712,10 @@ async fn main() -> Result<(), Box> { #[cfg(test)] mod tests { use super::*; - use openshell_core::proto::{MiddlewareSessionEnd, WebSocketPreflight, WebSocketSessionStart}; + use openshell_core::proto::{ + HttpRequestTarget, HttpResponseBodyUnit, HttpResponsePreflight, MiddlewareSessionEnd, + WebSocketPreflight, WebSocketSessionStart, + }; use prost_types::{ListValue, Value}; use std::collections::BTreeMap; @@ -511,13 +748,13 @@ mod tests { } #[tokio::test] - async fn manifest_advertises_http_and_websocket_bindings() { + async fn manifest_advertises_request_response_and_websocket_bindings() { let manifest = SupervisorMiddleware::describe(&ContentGuard, Request::new(())) .await .expect("describe") .into_inner(); - assert_eq!(manifest.bindings.len(), 2); + assert_eq!(manifest.bindings.len(), 3); assert_eq!( manifest.bindings[0].operation, SupervisorMiddlewareOperation::HttpRequest as i32 @@ -528,6 +765,126 @@ mod tests { SupervisorMiddlewareOperation::WebsocketMessage as i32 ); assert_eq!(manifest.bindings[1].max_payload_bytes, MAX_PAYLOAD_BYTES); + assert_eq!( + manifest.bindings[2].operation, + SupervisorMiddlewareOperation::HttpResponse as i32 + ); + assert_eq!( + manifest.bindings[2].phase, + SupervisorMiddlewarePhase::PreReturn as i32 + ); + } + + fn response_preflight(path: &str) -> HttpResponsePreflight { + HttpResponsePreflight { + target: Some(HttpRequestTarget { + path: path.into(), + ..Default::default() + }), + config: Some(config("redact", &["prototype-secret"], Some("[FILTERED]"))), + ..Default::default() + } + } + + #[test] + fn response_paths_select_all_modes() { + for (path, expected) in [ + ("/headers-only", HttpResponseBodyMode::HeadersOnly), + ("/whole-body", HttpResponseBodyMode::WholeBodyBytes), + ("/stream", HttpResponseBodyMode::StreamBytes), + ("/stream-close", HttpResponseBodyMode::StreamBytes), + ("/block", HttpResponseBodyMode::WholeBodyBytes), + ] { + let mut state = ResponseSessionState::default(); + let result = state.preflight(response_preflight(path)).unwrap(); + let Some(http_response_event_result::Result::PreflightResult(result)) = result.result + else { + panic!("expected preflight result"); + }; + let Some(http_response_preflight_result::Action::Inspect(inspect)) = result.action + else { + panic!("expected inspect action"); + }; + assert_eq!(inspect.body_mode, expected as i32); + } + } + + #[test] + fn response_whole_body_transforms_and_block_is_typed() { + let mut whole = ResponseSessionState::default(); + whole.preflight(response_preflight("/whole-body")).unwrap(); + let result = whole + .body(HttpResponseBodyUnit { + sequence: 1, + payload: Some(http_response_body_unit::Payload::Data(b"body".to_vec())), + end_of_stream: true, + }) + .unwrap(); + let Some(http_response_event_result::Result::BodyResult(body)) = result.result else { + panic!("expected body result"); + }; + let Some(http_response_body_result::Action::Transform(transform)) = body.action else { + panic!("expected body transform"); + }; + let Some(http_response_body_transform::Replacement::Data(data)) = transform.replacement + else { + panic!("expected data replacement"); + }; + assert_eq!(data, b"[whole] body"); + + let mut block = ResponseSessionState::default(); + block.preflight(response_preflight("/block")).unwrap(); + let result = block + .body(HttpResponseBodyUnit { + sequence: 1, + payload: Some(http_response_body_unit::Payload::Data( + b"prototype-secret".to_vec(), + )), + end_of_stream: true, + }) + .unwrap(); + let Some(http_response_event_result::Result::BodyResult(body)) = result.result else { + panic!("expected body result"); + }; + assert!(matches!( + body.action, + Some(http_response_body_result::Action::BlockDelivery(_)) + )); + assert_eq!(body.reason_code, "content_match"); + } + + #[test] + fn response_stream_returns_the_required_trailer_exchange() { + let mut state = ResponseSessionState::default(); + state.preflight(response_preflight("/stream")).unwrap(); + state + .body(HttpResponseBodyUnit { + sequence: 1, + payload: Some(http_response_body_unit::Payload::Data(b"body".to_vec())), + end_of_stream: true, + }) + .unwrap(); + let result = state.trailers().unwrap(); + let Some(http_response_event_result::Result::TrailersResult(trailers)) = result.result + else { + panic!("expected trailers result"); + }; + assert_eq!(trailers.trailer_mutations.len(), 1); + } + + #[test] + fn response_paths_outside_the_example_are_skipped() { + let mut state = ResponseSessionState::default(); + let result = state.preflight(response_preflight("/outside")).unwrap(); + let Some(http_response_event_result::Result::PreflightResult(result)) = result.result + else { + panic!("expected preflight result"); + }; + assert!(matches!( + result.action, + Some(http_response_preflight_result::Action::Skip(_)) + )); + assert_eq!(result.reason_code, "path_not_selected"); } #[tokio::test] @@ -737,4 +1094,11 @@ mod tests { assert_eq!(parsed.mode, Mode::Redact); assert_eq!(parsed.replacement, DEFAULT_REPLACEMENT); } + + #[test] + fn example_policy_is_valid() { + let policy = openshell_policy::parse_sandbox_policy(include_str!("../policy.yaml")) + .expect("example policy must parse"); + openshell_policy::validate_sandbox_policy(&policy).expect("example policy must be valid"); + } } diff --git a/examples/supervisor-middleware-content-guard/upstream.py b/examples/supervisor-middleware-content-guard/upstream.py new file mode 100644 index 0000000000..02897ddcfc --- /dev/null +++ b/examples/supervisor-middleware-content-guard/upstream.py @@ -0,0 +1,67 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import socketserver + + +class Handler(socketserver.BaseRequestHandler): + def handle(self): + request = b"" + while b"\r\n\r\n" not in request: + block = self.request.recv(4096) + if not block: + return + request += block + path = request.split(b" ", 2)[1] + if path == b"/headers-only": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Content-Length: 12\r\n\r\n" + b"headers-only" + ) + elif path == b"/whole-body": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Transfer-Encoding: chunked\r\n\r\n" + b"6\r\nwhole \r\n4\r\nbody\r\n0\r\n\r\n" + ) + elif path == b"/stream": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Trailer: x-example-body-bytes\r\n" + b"Transfer-Encoding: chunked\r\n\r\n" + b"6\r\nstream\r\n5\r\n body\r\n" + b"0\r\nX-Example-Body-Bytes: 0\r\n\r\n" + ) + elif path == b"/stream-close": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/event-stream\r\n" + b"Connection: close\r\n\r\n" + b"data: stream close\n\n" + ) + elif path == b"/block": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Content-Length: 16\r\n\r\n" + b"prototype-secret" + ) + else: + response = b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n" + self.request.sendall(response) + + +class DemoServer(socketserver.ThreadingTCPServer): + allow_reuse_address = True + + +with DemoServer(("0.0.0.0", 18081), Handler) as server: + print("response framing demo upstream listening on 0.0.0.0:18081", flush=True) + try: + server.serve_forever() + except KeyboardInterrupt: + pass diff --git a/proto/sandbox.proto b/proto/sandbox.proto index 51139ba461..0fe09899b7 100644 --- a/proto/sandbox.proto +++ b/proto/sandbox.proto @@ -394,6 +394,10 @@ message GetSandboxConfigResponse { // False also covers older gateways that do not advertise this capability; // supervisors preserve their legacy unauthenticated connection behavior. bool extension_authentication_enabled = 12; + // Supervisor-wide wall-clock limit for accumulating and processing a response + // through whole-body middleware. Zero means the supervisor default for + // compatibility with older gateways. + uint64 http_response_whole_body_timeout_ms = 13; } // Connection details for one operator-registered supervisor middleware service. diff --git a/sdk/go/proto/sandboxv1/sandbox.pb.go b/sdk/go/proto/sandboxv1/sandbox.pb.go index 8da143ebaa..9b4fa8a9d7 100644 --- a/sdk/go/proto/sandboxv1/sandbox.pb.go +++ b/sdk/go/proto/sandboxv1/sandbox.pb.go @@ -1826,6 +1826,10 @@ type GetSandboxConfigResponse struct { // False also covers older gateways that do not advertise this capability; // supervisors preserve their legacy unauthenticated connection behavior. ExtensionAuthenticationEnabled bool `protobuf:"varint,12,opt,name=extension_authentication_enabled,json=extensionAuthenticationEnabled,proto3" json:"extension_authentication_enabled,omitempty"` + // Supervisor-wide wall-clock limit for accumulating and processing a response + // through whole-body middleware. Zero means the supervisor default for + // compatibility with older gateways. + HttpResponseWholeBodyTimeoutMs uint64 `protobuf:"varint,13,opt,name=http_response_whole_body_timeout_ms,json=httpResponseWholeBodyTimeoutMs,proto3" json:"http_response_whole_body_timeout_ms,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -1944,6 +1948,13 @@ func (x *GetSandboxConfigResponse) GetExtensionAuthenticationEnabled() bool { return false } +func (x *GetSandboxConfigResponse) GetHttpResponseWholeBodyTimeoutMs() uint64 { + if x != nil { + return x.HttpResponseWholeBodyTimeoutMs + } + return 0 +} + // Connection details for one operator-registered supervisor middleware service. // V1 supports plaintext and server-authenticated TLS gRPC. type SupervisorMiddlewareService struct { @@ -2208,7 +2219,7 @@ const file_sandbox_proto_rawDesc = "" + "\x05value\"\x86\x01\n" + "\x10EffectiveSetting\x128\n" + "\x05value\x18\x01 \x01(\v2\".openshell.sandbox.v1.SettingValueR\x05value\x128\n" + - "\x05scope\x18\x02 \x01(\x0e2\".openshell.sandbox.v1.SettingScopeR\x05scope\"\xd1\x06\n" + + "\x05scope\x18\x02 \x01(\x0e2\".openshell.sandbox.v1.SettingScopeR\x05scope\"\x9e\a\n" + "\x18GetSandboxConfigResponse\x12;\n" + "\x06policy\x18\x01 \x01(\v2#.openshell.sandbox.v1.SandboxPolicyR\x06policy\x12\x18\n" + "\aversion\x18\x02 \x01(\rR\aversion\x12\x1f\n" + @@ -2223,7 +2234,8 @@ const file_sandbox_proto_rawDesc = "" + "\tworkspace\x18\n" + " \x01(\tR\tworkspace\x12C\n" + "\x1epolicy_validation_failure_mode\x18\v \x01(\tR\x1bpolicyValidationFailureMode\x12H\n" + - " extension_authentication_enabled\x18\f \x01(\bR\x1eextensionAuthenticationEnabled\x1ac\n" + + " extension_authentication_enabled\x18\f \x01(\bR\x1eextensionAuthenticationEnabled\x12K\n" + + "#http_response_whole_body_timeout_ms\x18\r \x01(\x04R\x1ehttpResponseWholeBodyTimeoutMs\x1ac\n" + "\rSettingsEntry\x12\x10\n" + "\x03key\x18\x01 \x01(\tR\x03key\x12<\n" + "\x05value\x18\x02 \x01(\v2&.openshell.sandbox.v1.EffectiveSettingR\x05value:\x028\x01\"\x99\x02\n" + diff --git a/skills/debug-openshell-cluster/SKILL.md b/skills/debug-openshell-cluster/SKILL.md index 8e86bc0643..4501934ee4 100644 --- a/skills/debug-openshell-cluster/SKILL.md +++ b/skills/debug-openshell-cluster/SKILL.md @@ -106,13 +106,15 @@ The gateway calls each interceptor's `Describe` RPC and validates its manifest a For operator-run supervisor middleware, inspect `[[openshell.supervisor.middleware]]`, service reachability, and both gateway and supervisor logs: ```bash -rg -n 'supervisor|middleware|grpc_endpoint|tls_ca_cert_path|audience|allow_insecure_transport|max_payload_bytes|timeout|gateway_jwt' /etc/openshell/gateway.toml +rg -n 'supervisor|middleware|grpc_endpoint|tls_ca_cert_path|audience|allow_insecure_transport|max_payload_bytes|timeout|http_response_whole_body_timeout|gateway_jwt' /etc/openshell/gateway.toml journalctl -u --no-pager --lines=200 journalctl -u openshell-gateway --no-pager --lines=200 openshell logs --tail --source sandbox ``` -The middleware service must start before the gateway and be reachable from both the gateway and sandbox supervisors. Gateway startup fails if `Describe` is unavailable, a manifest exposes duplicate operation/phase bindings, the registration claims the reserved `openshell/` namespace, or payload and timeout limits are invalid. Supported V1 bindings are `HTTP_REQUEST/PRE_CREDENTIALS` and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. When gateway JWT signing is disabled, supervisors preserve the legacy unauthenticated connector and do not request extension credentials. When signing is enabled, credential acquisition and verification failures are fail closed: check HTTPS trust and hostname validation, audience and issuer agreement, the token `kid`, gateway `RefreshSandboxToken` errors, and middleware logs. Changing a registration requires a gateway restart. A policy update can also fail before persistence if the selected implementation rejects its `network_middlewares` config. +The middleware service must start before the gateway and be reachable from both the gateway and sandbox supervisors. Gateway startup fails if `Describe` is unavailable, a manifest exposes duplicate operation/phase bindings, the registration claims the reserved `openshell/` namespace, or payload and timeout limits are invalid. Supported V1 bindings are `HTTP_REQUEST/PRE_CREDENTIALS`, `HTTP_RESPONSE/PRE_RETURN`, and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. When gateway JWT signing is disabled, supervisors preserve the legacy unauthenticated connector and do not request extension credentials. When signing is enabled, credential acquisition and verification failures are fail closed: check HTTPS trust and hostname validation, audience and issuer agreement, the token `kid`, gateway `RefreshSandboxToken` errors, and middleware logs. Changing a registration or `http_response_whole_body_timeout` requires a gateway restart. A policy update can also fail before persistence if the selected implementation rejects its `network_middlewares` config. + +For response failures, distinguish a deliberate `middleware_denied` decision from `response_delivery_failed`. Before response commitment, they produce canonical 403 and 502 responses respectively. After commitment, OpenShell aborts without adding an error body, final chunk, or trailer. A whole-body accumulation timeout is one non-resetting wall-clock deadline shared across response reads and whole-body barriers; inspect the active stage's `on_error`, the supervisor timeout setting, and `whole_body_accumulation_timeout` diagnostics. Header-only stages preserve upstream body framing. Body-processing stages normalize framing and send a trailer exchange, including an empty trailer set, after the final body result. At request time, distinguish attachment, binding selection, coverage, denial, and failure. A host-matched HTTP-only attachment can inspect the upgrade GET but does not join the WebSocket chain; the connection proceeds under either `on_error` mode and emits `binding_not_selected` coverage. A selected WebSocket stage receives text messages only. Binary messages pass under both modes, emit `unsupported_message_type` coverage, and consume a session sequence without an RPC. An explicit `middleware_denied` result is always enforced. WebSocket preflight returns `INSPECT`, voluntary `SKIP`, or authoritative `DENY`; `DENY` rejects the upgrade before upstream contact under both `on_error` modes. A selected-stage failure follows the policy-local `on_error`: `fail_closed` blocks the HTTP request or closes the WebSocket, while `fail_open` bypasses only that stage and emits a detection finding. A fail-open per-message capacity failure bypasses that message without disabling the stage. A timeout, transport failure, stream closure, missing or invalid response, duplicate or regressed sequence, or other failure that makes an established WebSocket stream unreliable disables that stage for later messages on the connection and emits `openshell.middleware.websocket_stage_disabled`. Confirm preflight, session-start, and session-end in service logs. OpenShell best-effort sends at most one session-end to each still-writable opened stage, including a preflight that terminates before session start; distinguish `MIDDLEWARE_DENIAL` from `MIDDLEWARE_FAILURE`. WebSocket message sequences are allocated session-wide; each stage receives a strictly increasing subset, so gaps are valid when binary messages or other units are not delivered to that stage. Zero, duplicate, or regressed sequences are protocol errors. If a running supervisor cannot install a new registry, it preserves its last-known-good generation and emits a configuration failure event. @@ -696,6 +698,7 @@ configuration — check that the gateway spawned the driver binary you expect | Policy mutation returns `FAILED_PRECONDITION` for endpoint ambiguity | Equally specific effective endpoint selectors disagree on connection or request-processing metadata | CLI error, base and provider-composed policy, affected profile attachments; confirm no new revision was stored | | Supervisor enters policy quarantine | A runtime candidate failed validation while `policy_validation_failure_mode = "fail_closed"` | Sandbox OCSF config/finding events, validation rationale, active generation, `previous_policy_active` | | HTTP request returns `middleware_failed` or `middleware_denied`, or WebSocket closes with `1008` | Selected stage failed or explicitly denied admitted traffic | Sandbox OCSF logs; policy-local middleware config; service availability; binding operation; `on_error` | +| HTTP response becomes canonical `403 middleware_denied`, `502 response_delivery_failed`, or closes mid-body | Response middleware blocked, failed before commitment, or stopped delivery after commitment | Sandbox OCSF response middleware events; `HTTP_RESPONSE/PRE_RETURN` binding; `on_error`; `http_response_whole_body_timeout`; service stream lifecycle | | WebSocket upgrades but a host-matched middleware receives no preflight or message RPC | The implementation did not advertise `WEBSOCKET_MESSAGE/PRE_CREDENTIALS` | `WEBSOCKET_MIDDLEWARE_COVERAGE state=binding_not_selected`; service `Describe`; the upgrade GET may still have used its HTTP binding | | Binary WebSocket message passes without a middleware RPC | Binary is unsupported by the V1 text-message binding under both `on_error` modes | `WEBSOCKET_MIDDLEWARE_COVERAGE state=unsupported_message_type`; the next text RPC may have a valid sequence gap | | WebSocket messages stop reaching middleware after one failure | A fail-open stage stream was disabled for the rest of the connection | `openshell.middleware.websocket_stage_disabled`; middleware timeout/stream/protocol logs. A per-message capacity bypass alone leaves the stage active. Reconnect to create a fresh stream after a genuine stream failure | diff --git a/skills/generate-sandbox-policy/SKILL.md b/skills/generate-sandbox-policy/SKILL.md index 73c0863df7..2f7436e567 100644 --- a/skills/generate-sandbox-policy/SKILL.md +++ b/skills/generate-sandbox-policy/SKILL.md @@ -83,7 +83,7 @@ Regardless of tier, extract (or infer) these from the user's description: | **Paths** | Specific URL paths or patterns | Only for custom/fine-grained | | **Enforcement** | `enforce` or `audit`? Default to `enforce`. | No — has a default | | **Binary** | Which binary/process should have access | Yes — ask if not stated | -| **Middleware** | Whether admitted HTTP requests or client WebSocket text messages need an ordered built-in or operator-run processing stage | No | +| **Middleware** | Whether admitted HTTP requests, final HTTP responses, or client WebSocket text messages need an ordered built-in or operator-run processing stage | No | If the host and access level are clear but binaries are not specified, ask the user which binary or process will be making the requests. Suggest common defaults like `/usr/bin/curl`, `/usr/local/bin/claude`, etc. @@ -209,11 +209,12 @@ Is L7 inspection needed? ### Middleware Decision -Add `network_middlewares` only when the user asks to inspect, transform, redact, or independently authorize admitted HTTP requests or client WebSocket text messages. Middleware runs after network and L7 policy admission and before provider credential injection. +Add `network_middlewares` only when the user asks to inspect, transform, redact, or independently authorize admitted HTTP requests, final HTTP responses, or client WebSocket text messages. Request middleware runs after network and L7 policy admission and before provider credential injection. Response middleware runs on the matching final response before it returns to the sandbox. - Use `openshell/regex` without gateway registration for fixed-pattern redaction of UTF-8 HTTP request bodies or complete client-to-upstream WebSocket text messages. - Use an operator-owned middleware name only when it is already registered under `[[openshell.supervisor.middleware]]` and reachable from both the gateway and sandbox supervisors. - Confirm that a requested WebSocket implementation exposes a `WEBSOCKET_MESSAGE/PRE_CREDENTIALS` binding. `openshell/regex` exposes this binding. A host-matched HTTP-only implementation may inspect the upgrade GET but does not join the post-upgrade chain; messages pass and OpenShell emits `binding_not_selected` coverage regardless of `on_error`. +- Confirm that requested response processing exposes `HTTP_RESPONSE/PRE_RETURN`. Response stages choose header-only, whole-body, or streaming inspection independently. Whole-body stages delay response commitment and are bounded by the gateway's supervisor-wide accumulation timeout. Intentional blocks produce a canonical 403 before commitment and abort the connection after commitment. - WebSocket middleware runs for both `ws://` and `wss://` and receives complete client text messages only. Binary messages pass under both error modes and emit `unsupported_message_type` coverage for active stages. Upstream-to-client messages remain uninspected. Do not claim that V1 provides all-message WebSocket inspection. - Treat `fail_open` on WebSocket as a session-scoped bypass: if the stage stream fails, OpenShell disables it for later messages on that connection and emits a state-change finding. Prefer `fail_closed` for required redaction or authorization. - `on_error` governs failures after an advertised operation binding is selected. It does not apply to an unadvertised WebSocket binding or binary-message pass-through. An explicit HTTP, WebSocket preflight, or WebSocket message denial is authoritative under both `fail_open` and `fail_closed`. @@ -380,6 +381,7 @@ Before presenting the policy to the user, verify correctness **and** flag breadt - [ ] Middleware `order` values are unique and no selected chain exceeds 10 stages - [ ] No fail-closed middleware selector can cover a `tls: skip` endpoint - [ ] Any required WebSocket control advertises `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`, and the user understands that V1 does not inspect binary messages +- [ ] Any required response control advertises `HTTP_RESPONSE/PRE_RETURN`, and whole-body buffering fits the registered payload limit and supervisor deadline - [ ] Endpoints contributed by a credentialed provider are not L4-only or `tls: skip` unless `allow_uninspected_credentials: true` explicitly records the exception ### Schema Warnings (log-only, but should be fixed) diff --git a/skills/openshell-cli/SKILL.md b/skills/openshell-cli/SKILL.md index c46f569789..dc3466f8bb 100644 --- a/skills/openshell-cli/SKILL.md +++ b/skills/openshell-cli/SKILL.md @@ -498,12 +498,14 @@ Edit `current-policy.yaml` to allow the blocked actions. **For policy content au - TLS termination configuration - Enforcement modes (`audit` vs `enforce`) - Binary matching patterns -- Ordered `network_middlewares`, host selection, HTTP and WebSocket bindings, and `fail_open` or `fail_closed` behavior +- Ordered `network_middlewares`, host selection, HTTP request/response and WebSocket bindings, and `fail_open` or `fail_closed` behavior `network_policies` and `network_middlewares` can be modified at runtime when the selected compute driver supports live policy updates. Use `--wait` to verify that the active runtime loaded the revision; do not infer enforcement from the gateway accepting the update. If `filesystem_policy`, `landlock`, or `process` need changes, the sandbox must be recreated. Built-in middleware such as `openshell/regex` needs no gateway registration. An operator-run middleware must already be registered under `[[openshell.supervisor.middleware]]`; changing that static registration requires a gateway restart. Middleware can inspect parsed HTTP request bodies and complete client-to-upstream WebSocket text messages over both `ws://` and `wss://` when the implementation advertises the matching binding. The built-in `openshell/regex` advertises both bindings and applies its fixed patterns to UTF-8 text. A host-matched HTTP-only attachment can inspect the upgrade GET but does not join the WebSocket chain; look for `binding_not_selected` coverage. Binary messages pass under both `on_error` modes and active stages emit `unsupported_message_type` coverage; upstream-to-client messages remain uninspected. A broken fail-open WebSocket stage is disabled for the rest of that connection; inspect sandbox OCSF logs for `openshell.middleware.websocket_stage_disabled`. +An operator-run implementation can also advertise `HTTP_RESPONSE/PRE_RETURN`. Each selected response stage chooses header-only, whole-body, or streaming inspection. Whole-body inspection delays downstream commitment and uses the supervisor-wide `http_response_whole_body_timeout` from gateway configuration. A valid middleware block returns the canonical 403 before commitment; after commitment OpenShell aborts the response without adding error bytes. Use sandbox OCSF events to distinguish an intentional block from `response_delivery_failed` and fail-open bypass. + ### Step 5: Push the updated policy ```bash