diff --git a/.github/workflows/branch-checks.yml b/.github/workflows/branch-checks.yml index 5d4a98c61a..8427102ce8 100644 --- a/.github/workflows/branch-checks.yml +++ b/.github/workflows/branch-checks.yml @@ -172,6 +172,7 @@ jobs: OPENSHELL_TELEMETRY_ENABLED: "false" run: | cargo nextest run --profile ci --workspace --features openshell-server/test-support + cargo test --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml - name: Verify telemetry can be compiled out run: | diff --git a/architecture/sandbox-limits.md b/architecture/sandbox-limits.md index 9635bc1c30..a14f999996 100644 --- a/architecture/sandbox-limits.md +++ b/architecture/sandbox-limits.md @@ -73,7 +73,7 @@ Middleware also validates every non-body envelope component. Important examples include 64 KiB service config, 4 KiB request context, 32 KiB target data, 128 request headers totaling 64 KiB, 64 header mutations, 32 findings per stage, and 64 metadata entries. The detailed external contract lives in -[Supervisor Middleware](../docs/extensibility/supervisor-middleware.mdx). +[Supervisor Middleware](../docs/extensibility/supervisor-middleware/index.mdx). The work semaphore bounds aggregate buffered middleware input to approximately `32 × 4 MiB`, plus bounded envelope and parser overhead. It is a concurrency diff --git a/architecture/sandbox.md b/architecture/sandbox.md index ba5b3c59b0..017df3b3d5 100644 --- a/architecture/sandbox.md +++ b/architecture/sandbox.md @@ -165,6 +165,11 @@ host selectors choose the chain independently of the network rule that admitted the request. Policy-local map keys identify configs, while built-in names or operator-owned registration names identify implementations. +The configured-literal content-guard example shares matching semantics across +request bodies, complete response bodies, and client WebSocket text messages. +It requires whole-body response inspection and returns a middleware failure +when that mode is unavailable. + Built-ins run in-process against a borrowed view of the chain's current HTTP request state. Operator services retain the bounded protobuf/gRPC contract, and the remote adapter materializes an owned HTTP evaluation only when a request @@ -187,6 +192,28 @@ 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`. +Valid HTTP that cannot fit the response middleware protocol, including non-UTF-8 +header values or an oversized preflight envelope, fails each selected stage +according to its `on_error` policy. An all-fail-open chain relays the original +bytes; a fail-closed stage prevents delivery. The relay validates HTTP syntax +and protected trailer declarations before allowing this bypass. + +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, +120-second accumulation deadline per response, defined in the response relay. +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. Streaming input units flush after bounded coalescing even within a +content-length body or transfer chunk. Coalescing cancels only input acquisition; +deadline transitions and downstream writes finish outside those timeouts. +The response runtime caps aggregate retained body data across stages and pending +output at 8 MiB. A transformation that exceeds the budget follows its stage's +failure policy, preserving its input when failing open. + 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 @@ -213,8 +240,9 @@ against body-aware L7 policy before later stages or the upstream can observe them. Requests, results, chain length, execution time, and diagnostics are bounded; external free-form diagnostic text is not exposed in responses or security logs. See -[Supervisor Middleware](../docs/extensibility/supervisor-middleware.mdx) for -configuration and protocol details. +[Supervisor Middleware](../docs/extensibility/supervisor-middleware/index.mdx) for +an introduction, or the [configuration guide](../docs/extensibility/supervisor-middleware/configure.mdx) +for service registration and policy attachment. `https://inference.local` is special. It bypasses OPA network policy and is handled by the inference interception path: diff --git a/crates/openshell-supervisor-middleware/src/lib.rs b/crates/openshell-supervisor-middleware/src/lib.rs index 1065c18faf..dcf695539a 100644 --- a/crates/openshell-supervisor-middleware/src/lib.rs +++ b/crates/openshell-supervisor-middleware/src/lib.rs @@ -5,8 +5,16 @@ pub mod headers; mod remote; +mod response; mod websocket; +pub use response::{ + HttpResponseDiagnostics, HttpResponseFinish, HttpResponseInvocation, + HttpResponseInvocationOutcome, HttpResponseMiddlewareFailure, HttpResponsePreflightInput, + HttpResponsePreflightOutcome, HttpResponseSession, MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES, + MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES, +}; + pub use websocket::{ WebSocketCoverage, WebSocketCoverageState, WebSocketInvocation, WebSocketInvocationOutcome, WebSocketMessageAdmission, WebSocketMessageOutcome, WebSocketMessageType, @@ -626,6 +634,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 +849,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 +3703,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 +3717,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..3e444a48b5 --- /dev/null +++ b/crates/openshell-supervisor-middleware/src/response.rs @@ -0,0 +1,3076 @@ +// 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; +/// Maximum logical body bytes retained across a session's stage buffers and +/// pending output. Temporary exchange copies have the per-binding payload cap. +pub const MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES: usize = 8 * 1024 * 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, + retained_body_bytes: usize, + 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; + self.release_body_bytes(&released); + 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; + // Between pushes, the first active whole-body barrier owns all input + // not returned to the relay (at most the 4 MiB binding cap). Later + // barriers cannot receive bytes until it finishes or disables itself; + // finish consumes the session and expiry disables all such barriers. + // Replacement admission reserves an additional upstream unit below. + self.retained_body_bytes += data.len(); + debug_assert!(self.retained_body_bytes <= MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES); + let output = self.process_units_from(0, vec![data], deadline).await?; + if !self.defer_output_until_finish { + self.release_body_bytes(&output); + 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); + self.release_body_bytes(&released); + 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; + } + + fn release_body_bytes(&mut self, units: &[Vec]) { + self.retained_body_bytes -= units.iter().map(Vec::len).sum::(); + } + + 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 replacement_size = match &decision.action { + BodyAction::Transform(replacement) + | BodyAction::SkipRemaining(CurrentBodyAction::Transform(replacement)) => { + Some(replacement.len()) + } + _ => None, + }; + if let Some(replacement_size) = replacement_size { + let retained = self.retained_body_bytes - input_size + replacement_size; + // Reserve room for one more normalized upstream unit. Whole-body + // barriers bound the input retained between calls to push_body. + if retained + > MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES - MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES + { + return self + .handle_stage_failure( + index, + "response_body_aggregate_over_capacity", + Some(sequence), + original, + ) + .await; + } + self.retained_body_bytes = retained; + } + 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 { + /// Apply selected stages' failure policies when valid HTTP cannot be encoded + /// in the middleware protocol. The caller must validate HTTP safety first. + pub fn http_response_input_unrepresentable( + &self, + entries: &[DescribedChainEntry], + ) -> HttpResponsePreflightOutcome { + response_preflight_input_failure(entries, Vec::new(), "response_input_unrepresentable") + } + + pub async fn preflight_http_response( + &self, + entries: &[ChainEntry], + input: HttpResponsePreflightInput, + ) -> miette::Result { + let described = self.describe_http_response_chain(entries).await?; + if described.is_empty() { + return Ok(empty_preflight_outcome(input.headers)); + } + if validate_preflight_input(&input).is_err() { + return Ok(response_preflight_input_failure( + &described, + input.headers, + "response_input_over_capacity", + )); + } + 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 { + sender + .send(HttpResponseEvent { + event: Some(http_response_event::Event::Preflight(preflight)), + }) + .await + .map_err(|_| tonic::Status::unavailable("middleware request stream closed"))?; + let mut responses = service + .service + .open_http_response_pre_return(receiver) + .await?; + 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, + retained_body_bytes: 0, + 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 > 0 { + 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" + | "response_input_unrepresentable" + | "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_preflight_input_failure( + entries: &[DescribedChainEntry], + headers: Vec, + reason: &str, +) -> HttpResponsePreflightOutcome { + let mut outcome = empty_preflight_outcome(headers); + for entry in entries { + if let Some(reason) = collect_preflight_failure(entry, reason, &mut outcome.invocations) { + outcome.allowed = false; + outcome.reason = reason; + break; + } + } + outcome +} + +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, + Expansion, + DeleteBody, + SkipBody, + 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(); + // Exercise servers that inspect the initial request before sending + // response headers, rather than returning a stream immediately. + let first = requests.next().await.expect("initial request"); + assert!(matches!(&first, Ok(HttpResponseEvent { + event: Some(http_response_event::Event::Preflight(_)) + }))); + let mut requests = futures::stream::iter([first]).chain(requests); + 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 | Script::Expansion + ) { + 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::Expansion + | Script::DeleteBody + | Script::SkipBody + | 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::Expansion => vec![b'x'; 128 * 1024], + Script::DeleteBody => Vec::new(), + Script::SkipBody => b"replacement".to_vec(), + 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, + }; + let transform = HttpResponseBodyTransform { + replacement: Some(http_response_body_transform::Replacement::Data( + replacement, + )), + }; + let action = if matches!(selected_script, Script::SkipBody) { + http_response_body_result::Action::SkipRemaining( + openshell_core::proto::HttpResponseBodySkipRemaining { + current: Some( + http_response_body_skip_remaining::Current::Transform( + transform, + ), + ), + }, + ) + } else { + http_response_body_result::Action::Transform(transform) + }; + 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(action), + ..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 response_preflight_envelope_limits_obey_selected_stage_policies() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::HeadersOnly, + })); + for limit in 0..4 { + let mut input = input(200); + match limit { + 0 => input.context.request_id = "x".repeat(MAX_MIDDLEWARE_CONTEXT_BYTES + 1), + 1 => input.target.path = "x".repeat(MAX_MIDDLEWARE_TARGET_BYTES + 1), + 2 => input.headers = vec![input.headers[0].clone(); MAX_MIDDLEWARE_HEADERS + 1], + _ => input.headers[0].value = "x".repeat(MAX_MIDDLEWARE_HEADER_BYTES + 1), + } + for last_policy in [OnError::FailOpen, OnError::FailClosed] { + let entries = [entry(OnError::FailOpen), entry(last_policy)]; + let outcome = runner + .preflight_http_response(&entries, input.clone()) + .await + .unwrap(); + assert_eq!(outcome.allowed, last_policy == OnError::FailOpen); + assert_eq!(outcome.headers, input.headers); + assert!(outcome.session.is_none()); + assert_eq!(outcome.invocations.len(), 2); + assert!( + outcome + .invocations + .iter() + .all(|invocation| invocation.failed && invocation.stage_disabled) + ); + assert_eq!( + outcome.invocations[1].failure_category.as_deref(), + Some("payload_capacity") + ); + } + } + let described = runner + .describe_http_response_chain(&[entry(OnError::FailOpen), entry(OnError::FailClosed)]) + .await + .unwrap(); + let outcome = runner.http_response_input_unrepresentable(&described); + assert!(!outcome.allowed); + assert_eq!(outcome.invocations.len(), 2); + assert!(outcome.invocations.iter().all(|invocation| { + invocation.failure_category.as_deref() == Some("response_not_inspectable") + })); + } + + #[test] + fn stream_mode_requires_only_one_byte_of_payload_capacity() { + let mut described = DescribedChainEntry { + entry: entry(OnError::FailClosed), + service: None, + binding: None, + max_payload_bytes: 1, + timeout: Duration::from_millis(500), + }; + + let modes = permitted_body_modes(&input(200), &described, None); + assert!(modes.contains(&(HttpResponseBodyMode::StreamBytes as i32))); + + described.max_payload_bytes = 0; + let modes = permitted_body_modes(&input(200), &described, None); + assert!(!modes.contains(&(HttpResponseBodyMode::StreamBytes as i32))); + } + + struct ReadPreflightBeforeOpening { + failure: Option, + } + + #[tonic::async_trait] + impl InProcessMiddleware for ReadPreflightBeforeOpening { + async fn describe(&self) -> MiddlewareManifest { + let mut manifest = response_manifest("test/response"); + manifest.bindings[0].timeout = "10ms".into(); + manifest + } + + async fn validate_config(&self, _: &str, _: &prost_types::Struct) -> miette::Result<()> { + Ok(()) + } + + async fn evaluate_http_request( + &self, + _: HttpRequestView<'_>, + ) -> miette::Result { + unreachable!() + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> Result { + let first = requests.recv().await.expect("initial preflight"); + assert!(matches!( + first.event, + Some(http_response_event::Event::Preflight(_)) + )); + if let Some(hang) = self.failure { + if hang { + futures::future::pending::<()>().await; + } + return Err(tonic::Status::unavailable("startup failed")); + } + let response = 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::new(), + }, + )), + ..Default::default() + }, + )), + }; + Ok(Box::pin(futures::stream::iter([Ok(response)]))) + } + } + + #[tokio::test] + async fn preflight_can_be_read_before_open_returns() { + let runner = ChainRunner::new(Arc::new(ReadPreflightBeforeOpening { failure: None })); + let outcome = tokio::time::timeout( + Duration::from_secs(1), + runner.preflight_http_response(&[entry(OnError::FailClosed)], input(200)), + ) + .await + .expect("bounded startup") + .expect("preflight"); + assert!(outcome.allowed, "{}", outcome.reason); + } + + #[tokio::test] + async fn preflight_opening_failure_obeys_policy_and_releases_admission() { + for hang in [false, true] { + for on_error in [OnError::FailOpen, OnError::FailClosed] { + let runner = ChainRunner::new(Arc::new(ReadPreflightBeforeOpening { + failure: Some(hang), + })); + let permits = runner.registry.session_admission.available_permits(); + let outcome = tokio::time::timeout( + Duration::from_secs(1), + runner.preflight_http_response(&[entry(on_error)], input(200)), + ) + .await + .expect("bounded opening failure") + .unwrap(); + assert_eq!(outcome.allowed, on_error == OnError::FailOpen); + assert!(outcome.session.is_none()); + assert_eq!( + runner.registry.session_admission.available_permits(), + permits + ); + assert!(outcome.invocations[0].failed); + } + } + } + + #[tokio::test] + async fn multiple_whole_body_barriers_preserve_accounting_on_overflow_and_expiry() { + for expire in [false, true] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Configured, + })); + let mut entries = vec![ + configured_entry("first", 0, "whole"), + configured_entry("second", 1, "whole"), + configured_entry("stream", 2, "stream"), + ]; + for entry in &mut entries { + entry.on_error = OnError::FailOpen; + } + let mut outcome = runner + .preflight_http_response(&entries, input(200)) + .await + .unwrap(); + let mut session = outcome.session.take().unwrap(); + for _ in 0..2 { + assert!( + session + .push_body(vec![b'a'; 2048]) + .await + .unwrap() + .is_empty() + ); + assert!(session.retained_body_bytes <= 4096); + assert_eq!( + session + .stages + .iter() + .filter(|stage| !stage.whole_body.is_empty()) + .count(), + 1 + ); + } + let output = if expire { + session.start_whole_body_deadline(Duration::ZERO); + session.expire_whole_body_deadline().await.unwrap() + } else { + session.push_body(vec![b'a'; 2048]).await.unwrap() + }; + assert_eq!( + output.concat(), + vec![b'A'; if expire { 4096 } else { 6144 }] + ); + assert_eq!(session.retained_body_bytes, 0); + assert_eq!( + session.push_body(b"next".to_vec()).await.unwrap().concat(), + b"NEXT" + ); + assert_eq!(session.retained_body_bytes, 0); + assert!( + session + .finish(Vec::new()) + .await + .unwrap() + .body_units + .is_empty() + ); + } + } + + #[tokio::test] + async fn deleted_and_skip_remaining_units_release_body_accounting() { + for script in [Script::DeleteBody, Script::SkipBody] { + let runner = ChainRunner::new(Arc::new(ResponseService { script })); + let mut outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .unwrap(); + let mut session = outcome.session.take().unwrap(); + for index in 0..3 { + let output = session.push_body(b"original".to_vec()).await.unwrap(); + let expected = match script { + Script::DeleteBody => Vec::new(), + Script::SkipBody if index == 0 => b"replacement".to_vec(), + Script::SkipBody => b"original".to_vec(), + _ => unreachable!(), + }; + assert_eq!(output.concat(), expected); + assert_eq!(session.retained_body_bytes, 0); + } + assert!( + session + .finish(Vec::new()) + .await + .unwrap() + .body_units + .is_empty() + ); + } + } + + #[tokio::test] + async fn expanding_stages_obey_aggregate_budget_and_failure_policy() { + for on_error in [OnError::FailClosed, OnError::FailOpen] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Expansion, + })); + let permits = runner.registry.session_admission.available_permits(); + let entries = (0..9) + .map(|order| { + let mut entry = entry(on_error); + entry.name = format!("expand-{order}"); + entry.order = order; + entry + }) + .collect::>(); + let mut outcome = runner + .preflight_http_response(&entries, input(200)) + .await + .unwrap(); + let mut session = outcome.session.take().unwrap(); + let result = tokio::time::timeout(Duration::from_secs(5), session.push_body(vec![1])) + .await + .expect("bounded expansion"); + match result { + Ok(output) => { + assert_eq!(on_error, OnError::FailOpen); + assert!( + output.iter().map(Vec::len).sum::() + < MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES + ); + assert!(output.iter().flatten().all(|byte| *byte == b'x')); + assert_eq!(session.retained_body_bytes, 0); + assert!( + session + .invocations + .iter() + .any(|invocation| invocation.outcome + == HttpResponseInvocationOutcome::FailOpen) + ); + // Holding returned output applies backpressure: subsequent + // stage work starts only when the relay calls again. + drop(output); + for _ in 0..3 { + let output = session.push_body(vec![1]).await.unwrap(); + assert!( + output.iter().map(Vec::len).sum::() + < MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES + ); + assert_eq!(session.retained_body_bytes, 0); + } + let finish = session.finish(Vec::new()).await.unwrap(); + assert!( + finish.body_units.iter().map(Vec::len).sum::() + < MAX_HTTP_RESPONSE_RETAINED_BODY_BYTES + ); + } + Err(failure) => { + assert_eq!(on_error, OnError::FailClosed); + assert!( + failure + .reason + .contains("response_body_aggregate_over_capacity") + ); + session + .end(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + } + } + assert_eq!( + runner.registry.session_admission.available_permits(), + permits + ); + } + } + + #[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(()); + tokio::time::timeout(Duration::from_secs(2), server_task) + .await + .expect("bounded server shutdown") + .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 aec52cf9cd..9bb0af5f5b 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: super::rest::DEFAULT_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..2797916e2f 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,35 @@ 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 +} + +/// Default wall-clock bound shared by whole-body stages in one response. +pub(crate) const DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT: std::time::Duration = + std::time::Duration::from_mins(2); + +/// 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 +1185,7 @@ where websocket: websocket_response, client_requested_upgrade, }, + response_middleware, ) .await?; @@ -3124,6 +3158,7 @@ async fn relay_response( upstream: &mut U, client: &mut C, options: RelayResponseOptions, + response_middleware: Option>, ) -> Result where U: AsyncRead + Unpin, @@ -3133,7 +3168,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 +3185,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 +3255,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,569 +3361,2456 @@ 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]; + // Ordinary responses retain the HTTP parser's byte-preserving behavior. + // Response-specific normalization and limits apply only to selected hooks. + if middleware.chain.is_empty() { + return Ok(None); + } + let parsed = match middleware + .runner + .describe_http_response_chain(middleware.chain) + .await + { + Ok(described) if described.is_empty() => return Ok(None), + Ok(described) => { + parse_response_head_for_middleware(header_bytes).map(|parsed| (described, parsed)) + } + Err(error) => Err(error), + }; + let (described, parsed) = match parsed { + 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_result = if parsed.representable { + middleware + .runner + .preflight_http_response(middleware.chain, input) + .await + } else { + Ok(middleware + .runner + .http_response_input_unrepresentable(&described)) + }; + let preflight = match preflight_result { + 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 Some(mut session) = preflight.session else { + if preflight.headers == original_headers { + return Ok(None); + } + let status_line = response_status_line(header_bytes)?; + 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 status_line = response_status_line(header_bytes)?; + 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" - )); - } - 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 whole_body = session.requires_whole_body(); + if whole_body { + session.start_whole_body_deadline(middleware.whole_body_timeout); } - - 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 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(); } - (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) + 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 { - Err(miette!( - "upstream negotiated WebSocket extension that was not offered" - )) + 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)); } - } -} + }; -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")); + if let Some(guard) = middleware.generation_guard + && let Err(error) = guard.ensure_current() + { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) + .await; + return Err(error); } - 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 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 { + 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); } - 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 whole_body && !committed { + let mut headers = preflight.headers; + if finish.strip_stale_integrity_headers { + strip_response_integrity_headers(&mut headers); } - 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; + 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 { + 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()?; + } + } + } 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 { + representable: bool, + 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, - } - - impl CountingReader { - fn new(bytes: Vec) -> Self { - Self { - bytes, - position: 0, - reads: 0, +fn parse_response_head_for_middleware(header_bytes: &[u8]) -> Result { + // Lossy decoding preserves ASCII syntax and control bytes for validation. + // Never pass replacement text to middleware or use it for delivery. + let header = String::from_utf8_lossy(header_bytes); + let representable = std::str::from_utf8(header_bytes).is_ok(); + if parse_status_code(&header).is_none() { + return Err(miette!("HTTP response status line is malformed")); + } + 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(())) + 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 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, - }, - )), + 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 { + representable, + headers: if representable { headers } else { Vec::new() }, + connection_nominated, + declared_trailers, + }) +} - fn remove_header(name: &str) -> HeaderMutation { - HeaderMutation { - operation: Some(header_mutation::Operation::Remove( - openshell_core::proto::RemoveHeader { name: name.into() }, - )), - } +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(()) +} - #[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_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(()) +} - #[derive(Debug)] - struct CapturedFrame { - fin_opcode: u8, - masked: bool, - payload: Vec, - } +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_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 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 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") - } +#[derive(Clone, Copy)] +enum ResponseFraming { + Preserve(BodyLength), + ContentLength(u64), + Chunked, +} - 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(); - }); +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 mut received = Vec::new(); - client.read_to_end(&mut received).await.unwrap(); - task.await.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" + ) + }); +} - let response = String::from_utf8(received).unwrap(); - let (_, body) = response.split_once("\r\n\r\n").unwrap(); - serde_json::from_str(body).unwrap() - } +struct BufferedResponseReader<'a, R> { + upstream: &'a mut R, + buffered: &'a [u8], + position: usize, + exact_buffer: Vec, + exact_target: Option, + line_buffer: Vec, +} - 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 +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 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 + 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 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_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 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_line(&mut self) -> Result> { + loop { + 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")); + } + 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)); + } } } +} - fn compress_test_permessage_deflate(payload: &[u8]) -> Vec { - let mut compressor = Compress::new(Compression::fast(), false); +#[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_response_payload_with_deadline( + reader, + length, + !pending.is_empty(), + session, + client, + &mut framing, + ) + .await?; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + let partial = unit.len() < length; + remaining -= unit.len() as u64; + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + &mut framing, + unit_limit, + ) + .await?; + if partial { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + &mut framing, + ) + .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_response_payload_with_deadline( + reader, + length, + !pending.is_empty(), + session, + client, + &mut framing, + ) + .await?; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + let partial = unit.len() < length; + remaining -= unit.len(); + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + &mut framing, + unit_limit, + ) + .await?; + if partial { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + &mut framing, + ) + .await?; + } + } + let terminator = if let Ok(result) = + tokio::time::timeout(RESPONSE_UNIT_COALESCE_TIMEOUT, reader.read_exact_vec(2)) + .await + { + result? + } else { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + &mut framing, + ) + .await?; + read_exact_response_with_deadline(reader, 2, session, client, &mut framing) + .await? + }; + if terminator != 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, reader.read_line()).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 { + // Cancel only input acquisition, never expiry or client writes. + let wait = if pending.is_empty() { + if event_stream { + None + } else { + Some(RELAY_EOF_IDLE_TIMEOUT) + } + } else { + Some(RESPONSE_UNIT_COALESCE_TIMEOUT) + }; + let read_deadline = wait.map(|wait| tokio::time::Instant::now() + wait); + let whole_deadline = session.whole_body_deadline(); + let deadline = match (read_deadline, whole_deadline) { + (Some(a), Some(b)) => Some(a.min(b)), + (a, b) => a.or(b), + }; + let next = if let Some(deadline) = deadline { + if let Ok(result) = + tokio::time::timeout_at(deadline, reader.read_some(unit_limit)).await + { + result? + } else { + if whole_deadline.is_some_and(|d| d <= tokio::time::Instant::now()) { + expire_whole_body_deadline(session, client, &mut framing).await?; + } else if pending.is_empty() { + return Ok(Vec::new()); + } else { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + &mut framing, + ) + .await?; + } + continue; + } + } else { + reader.read_some(unit_limit).await? + }; + 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 +} + +// Unit size is a maximum. Coalesce available payload without waiting for +// the rest of a transfer chunk, and finish expiry outside read timeouts. +async fn read_response_payload_with_deadline( + reader: &mut BufferedResponseReader<'_, R>, + limit: usize, + has_pending: bool, + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + framing: &mut ResponseOutputState<'_>, +) -> Result> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + let mut payload = Vec::new(); + let mut coalesce_deadline = + has_pending.then(|| tokio::time::Instant::now() + RESPONSE_UNIT_COALESCE_TIMEOUT); + loop { + let whole_deadline = session.whole_body_deadline(); + let deadline = match (coalesce_deadline, whole_deadline) { + (Some(a), Some(b)) => Some(a.min(b)), + (a, b) => a.or(b), + }; + let read = reader.read_some(limit - payload.len()); + let result = if let Some(deadline) = deadline { + if let Ok(result) = tokio::time::timeout_at(deadline, read).await { + result + } else { + if whole_deadline.is_some_and(|d| d <= tokio::time::Instant::now()) { + expire_whole_body_deadline(session, client, framing).await?; + } + if coalesce_deadline.is_some_and(|d| d <= tokio::time::Instant::now()) { + return Ok(payload); + } + continue; + } + } else { + read.await + }; + let data = result?.ok_or_else(|| miette!("HTTP response body ended unexpectedly"))?; + payload.extend_from_slice(&data); + if payload.len() == limit { + return Ok(payload); + } + coalesce_deadline + .get_or_insert_with(|| tokio::time::Instant::now() + RESPONSE_UNIT_COALESCE_TIMEOUT); + } +} + +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, + SlowStream, + InvalidBodySequence, + InvalidWholeBodySequence, + } + + struct ResponseRelayService { + script: ResponseRelayScript, + request_only: bool, + } + + #[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: if self.request_only { + SupervisorMiddlewareOperation::HttpRequest + } else { + SupervisorMiddlewareOperation::HttpResponse + } as i32, + phase: if self.request_only { + SupervisorMiddlewarePhase::PreCredentials + } else { + 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, + > { + assert!( + !self.request_only, + "request-only service received a response" + ); + let mut 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(preflight) => { + if preflight + .config + .as_ref() + .is_some_and(|config| config.fields.contains_key("whole_body")) + { + script = ResponseRelayScript::WholeBody; + } + 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::SlowStream + | 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::SlowStream + | ResponseRelayScript::BlockStream + | ResponseRelayScript::InvalidBodySequence => { + data.to_ascii_uppercase() + } + ResponseRelayScript::HeadersOnly + | ResponseRelayScript::BlockPreflight => break, + }; + if matches!( + script, + ResponseRelayScript::SlowWholeBody + | ResponseRelayScript::SlowStream + ) { + 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(); @@ -4311,1196 +6267,2258 @@ mod tests { } #[test] - fn buffered_request_parser_rejects_missing_header_terminator() { - let err = request_from_buffered_http( - "GET", - "/v1/items", - "/v1/items", - b"GET /v1/items HTTP/1.1\r\nHost: api.example.com\r\n".to_vec(), - ) - .expect_err("unterminated headers must be rejected"); - - assert!(err.to_string().contains("missing the CRLF terminator")); + fn buffered_request_parser_rejects_missing_header_terminator() { + let err = request_from_buffered_http( + "GET", + "/v1/items", + "/v1/items", + b"GET /v1/items HTTP/1.1\r\nHost: api.example.com\r\n".to_vec(), + ) + .expect_err("unterminated headers must be rejected"); + + assert!(err.to_string().contains("missing the CRLF terminator")); + } + + #[test] + fn buffered_request_parser_rejects_malformed_header_fields() { + for raw in [ + b"GET /v1/items HTTP/1.1\r\nX-Test: first\r\n continued\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nX-Test value\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nX-Test : value\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nX@Test: value\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nX-Test: before\0after\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nX-Test: before\x7fafter\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nX-Test: before\rafter\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nConnection: x guard\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nConnection: content-length\r\nContent-Length: 0\r\n\r\n" + .as_slice(), + ] { + request_from_buffered_http("GET", "/v1/items", "/v1/items", raw.to_vec()) + .expect_err("malformed buffered header fields must be rejected"); + } + } + + #[test] + fn buffered_request_parser_rejects_malformed_request_lines() { + for raw in [ + b"GET /v1/items HTTP/1.1 extra\r\nHost: api.example.com\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), + b"GET\t/v1/items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), + b"GE(T /v1/items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), + b"GET /v1/\0items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), + b"GET /v1/items HTTP/2\r\nHost: api.example.com\r\n\r\n".as_slice(), + b"GET /v1/items\r\nHost: api.example.com\r\n\r\n".as_slice(), + ] { + request_from_buffered_http("GET", "/v1/items", "/v1/items", raw.to_vec()) + .expect_err("malformed buffered request lines must be rejected"); + } + } + + #[test] + fn parse_chunked() { + let headers = + "POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n"; + match parse_body_length(headers).unwrap() { + BodyLength::Chunked => {} + other => panic!("Expected Chunked, got {other:?}"), + } + } + + #[test] + fn parse_no_body() { + let headers = "GET /api HTTP/1.1\r\nHost: example.com\r\n\r\n"; + match parse_body_length(headers).unwrap() { + BodyLength::None => {} + other => panic!("Expected None, got {other:?}"), + } + } + + #[test] + fn parse_target_query_parses_duplicate_values() { + let (path, query) = parse_target_query("/download?tag=a&tag=b").expect("parse"); + assert_eq!(path, "/download"); + assert_eq!( + query.get("tag").cloned(), + Some(vec!["a".into(), "b".into()]) + ); + } + + #[test] + fn parse_target_query_decodes_percent_and_plus() { + let (path, query) = parse_target_query("/download?slug=my%2Fskill&name=Foo+Bar").unwrap(); + assert_eq!(path, "/download"); + assert_eq!( + query.get("slug").cloned(), + Some(vec!["my/skill".to_string()]) + ); + // `+` is decoded as space per application/x-www-form-urlencoded. + // Literal `+` should be sent as `%2B`. + assert_eq!( + query.get("name").cloned(), + Some(vec!["Foo Bar".to_string()]) + ); + } + + #[test] + fn parse_target_query_literal_plus_via_percent_encoding() { + let (_path, query) = parse_target_query("/search?q=a%2Bb").unwrap(); + assert_eq!( + query.get("q").cloned(), + Some(vec!["a+b".to_string()]), + "%2B should decode to literal +" + ); + } + + #[test] + fn parse_target_query_empty_value() { + let (_path, query) = parse_target_query("/api?tag=").unwrap(); + assert_eq!( + query.get("tag").cloned(), + Some(vec!["".to_string()]), + "key with empty value should produce empty string" + ); + } + + #[test] + fn parse_target_query_key_without_value() { + let (_path, query) = parse_target_query("/api?verbose").unwrap(); + assert_eq!( + query.get("verbose").cloned(), + Some(vec!["".to_string()]), + "key without = should produce empty string value" + ); + } + + #[test] + fn parse_target_query_unicode_after_decoding() { + // "café" = c a f %C3%A9 + let (_path, query) = parse_target_query("/search?q=caf%C3%A9").unwrap(); + assert_eq!( + query.get("q").cloned(), + Some(vec!["café".to_string()]), + "percent-encoded UTF-8 should decode correctly" + ); + } + + #[test] + fn parse_target_query_empty_query_string() { + let (path, query) = parse_target_query("/api?").unwrap(); + assert_eq!(path, "/api"); + assert!( + query.is_empty(), + "empty query after ? should produce empty map" + ); + } + + #[test] + fn parse_target_query_rejects_malformed_percent_encoding() { + let err = parse_target_query("/download?slug=bad%2").expect_err("expected parse error"); + assert!( + err.to_string().contains("percent-encoding"), + "unexpected error: {err}" + ); + } + + /// SEC-009: Reject requests with both Content-Length and Transfer-Encoding + /// to prevent CL/TE request smuggling (RFC 7230 Section 3.3.3). + #[test] + fn reject_dual_content_length_and_transfer_encoding() { + let headers = "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 5\r\nTransfer-Encoding: chunked\r\n\r\n"; + assert!( + parse_body_length(headers).is_err(), + "Must reject request with both CL and TE" + ); + } + + /// SEC-009: Same rejection regardless of header order. + #[test] + fn reject_dual_transfer_encoding_and_content_length() { + let headers = "POST /api HTTP/1.1\r\nHost: x\r\nTransfer-Encoding: chunked\r\nContent-Length: 5\r\n\r\n"; + assert!( + parse_body_length(headers).is_err(), + "Must reject request with both TE and CL" + ); + } + + /// SEC: Reject differing duplicate Content-Length headers. + #[test] + fn reject_differing_duplicate_content_length() { + let headers = + "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 0\r\nContent-Length: 50\r\n\r\n"; + assert!( + parse_body_length(headers).is_err(), + "Must reject differing duplicate Content-Length" + ); + } + + /// SEC: Accept identical duplicate Content-Length headers. + #[test] + fn accept_identical_duplicate_content_length() { + let headers = + "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 42\r\nContent-Length: 42\r\n\r\n"; + match parse_body_length(headers).unwrap() { + BodyLength::ContentLength(42) => {} + other => panic!("Expected ContentLength(42), got {other:?}"), + } + } + + /// SEC: Reject non-numeric Content-Length values. + #[test] + fn reject_non_numeric_content_length() { + let headers = "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: abc\r\n\r\n"; + assert!( + parse_body_length(headers).is_err(), + "Must reject non-numeric Content-Length" + ); + } + + /// SEC: Reject when second Content-Length is non-numeric (bypass test). + #[test] + fn reject_valid_then_invalid_content_length() { + let headers = + "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 42\r\nContent-Length: abc\r\n\r\n"; + assert!( + parse_body_length(headers).is_err(), + "Must reject when any Content-Length is non-numeric" + ); + } + + /// SEC: Unsupported transfer codings must not be silently treated as no body. + #[test] + fn reject_unsupported_transfer_encoding_sequences() { + for value in [ + "gzip", + "gzip, chunked", + "chunked, gzip", + "chunked, chunked", + "chunkedx", + "chunked,", + ] { + let headers = + format!("POST /api HTTP/1.1\r\nHost: x\r\nTransfer-Encoding: {value}\r\n\r\n"); + assert!( + parse_body_length(&headers).is_err(), + "unsupported transfer coding must be rejected: {value}" + ); + } + } + + #[test] + fn reject_multiple_chunked_transfer_encoding_fields() { + let headers = "POST /api HTTP/1.1\r\nHost: x\r\nTransfer-Encoding: chunked\r\nTransfer-Encoding: chunked\r\n\r\n"; + assert!(parse_body_length(headers).is_err()); + } + + #[test] + fn reject_content_length_with_unsupported_transfer_encoding() { + let headers = + "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 4\r\nTransfer-Encoding: gzip\r\n\r\n"; + assert!(parse_body_length(headers).is_err()); } #[test] - fn buffered_request_parser_rejects_malformed_header_fields() { - for raw in [ - b"GET /v1/items HTTP/1.1\r\nX-Test: first\r\n continued\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nX-Test value\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nX-Test : value\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nX@Test: value\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nX-Test: before\0after\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nX-Test: before\x7fafter\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nX-Test: before\rafter\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nConnection: x guard\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nConnection: content-length\r\nContent-Length: 0\r\n\r\n" - .as_slice(), - ] { - request_from_buffered_http("GET", "/v1/items", "/v1/items", raw.to_vec()) - .expect_err("malformed buffered header fields must be rejected"); - } + fn strip_connection_nominated_headers_before_forwarding() { + let raw = b"GET /api HTTP/1.1\r\nHost: x\r\nX-Guard: hidden\r\nConnection: keep-alive, x-guard\r\nKeep-Alive: timeout=5\r\nX-Visible: yes\r\n\r\n"; + let sanitized = + strip_connection_nominated_headers(raw, false).expect("sanitize request headers"); + let sanitized = String::from_utf8(sanitized).unwrap(); + + assert!(!sanitized.to_ascii_lowercase().contains("x-guard:")); + assert!(!sanitized.to_ascii_lowercase().contains("keep-alive:")); + assert!(!sanitized.to_ascii_lowercase().contains("connection:")); + assert!(sanitized.contains("X-Visible: yes\r\n")); } #[test] - fn buffered_request_parser_rejects_malformed_request_lines() { - for raw in [ - b"GET /v1/items HTTP/1.1 extra\r\nHost: api.example.com\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), - b"GET\t/v1/items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), - b"GE(T /v1/items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), - b"GET /v1/\0items HTTP/1.1\r\nHost: api.example.com\r\n\r\n".as_slice(), - b"GET /v1/items HTTP/2\r\nHost: api.example.com\r\n\r\n".as_slice(), - b"GET /v1/items\r\nHost: api.example.com\r\n\r\n".as_slice(), + fn connection_sanitization_preserves_only_websocket_upgrade_exception() { + let raw = b"GET /ws HTTP/1.1\r\nHost: x\r\nUpgrade: h2c\r\nUpgrade: h2c, websocket\r\nConnection: keep-alive, Upgrade, x-guard\r\nX-Guard: hidden\r\n\r\n"; + let sanitized = + strip_connection_nominated_headers(raw, true).expect("sanitize websocket headers"); + let sanitized = String::from_utf8(sanitized).unwrap(); + + assert_eq!(sanitized.matches("Upgrade: websocket\r\n").count(), 1); + assert_eq!(sanitized.matches("Connection: Upgrade\r\n").count(), 1); + assert!(!sanitized.to_ascii_lowercase().contains("upgrade: h2c")); + assert!(!sanitized.to_ascii_lowercase().contains("x-guard:")); + assert!(!sanitized.contains("keep-alive")); + } + + #[tokio::test] + async fn middleware_fixed_read_ahead_consumes_expect_continue() { + for (already_read, remaining, should_acknowledge) in [ + (b"hello".as_slice(), b"".as_slice(), false), + (b"he".as_slice(), b"llo".as_slice(), true), ] { - request_from_buffered_http("GET", "/v1/items", "/v1/items", raw.to_vec()) - .expect_err("malformed buffered request lines must be rejected"); + let mut raw = b"POST /api HTTP/1.1\r\nHost: example.com\r\nContent-Length: 5\r\nExpect: 100-continue\r\n\r\n".to_vec(); + raw.extend_from_slice(already_read); + let req = L7Request { + action: "POST".into(), + target: "/api".into(), + query_params: HashMap::new(), + raw_header: raw, + body_length: BodyLength::ContentLength(5), + }; + let (mut client, mut peer) = tokio::io::duplex(128); + peer.write_all(remaining).await.unwrap(); + + let result = buffer_request_body_for_middleware(&req, &mut client, None, 1024) + .await + .expect("fixed body should buffer"); + let BufferResult::Buffered(buffered) = result else { + panic!("fixed body unexpectedly exceeded capacity") + }; + assert_eq!(buffered.body, b"hello"); + assert!( + !String::from_utf8_lossy(&buffered.headers) + .to_ascii_lowercase() + .contains("expect:") + ); + + let mut response = [0u8; 64]; + let read = tokio::time::timeout( + std::time::Duration::from_millis(20), + peer.read(&mut response), + ) + .await; + if should_acknowledge { + let count = read.expect("partial body should be acknowledged").unwrap(); + assert_eq!(&response[..count], b"HTTP/1.1 100 Continue\r\n\r\n"); + } else { + assert!( + read.is_err(), + "complete read-ahead should not be acknowledged" + ); + } } } - #[test] - fn parse_chunked() { - let headers = - "POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n"; - match parse_body_length(headers).unwrap() { - BodyLength::Chunked => {} - other => panic!("Expected Chunked, got {other:?}"), + #[tokio::test] + async fn middleware_chunked_read_ahead_consumes_expect_continue() { + for (already_read, remaining, should_acknowledge) in [ + (b"5\r\nhello\r\n0\r\n\r\n".as_slice(), b"".as_slice(), false), + (b"5\r\nhe".as_slice(), b"llo\r\n0\r\n\r\n".as_slice(), true), + ] { + let mut raw = b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\nExpect: 100-continue\r\n\r\n".to_vec(); + raw.extend_from_slice(already_read); + let req = L7Request { + action: "POST".into(), + target: "/api".into(), + query_params: HashMap::new(), + raw_header: raw, + body_length: BodyLength::Chunked, + }; + let (mut client, mut peer) = tokio::io::duplex(128); + peer.write_all(remaining).await.unwrap(); + + let result = buffer_request_body_for_middleware(&req, &mut client, None, 1024) + .await + .expect("chunked body should buffer"); + let BufferResult::Buffered(buffered) = result else { + panic!("chunked body unexpectedly exceeded capacity") + }; + assert_eq!(buffered.body, b"hello"); + assert!( + !String::from_utf8_lossy(&buffered.headers) + .to_ascii_lowercase() + .contains("expect:") + ); + + let mut response = [0u8; 64]; + let read = tokio::time::timeout( + std::time::Duration::from_millis(20), + peer.read(&mut response), + ) + .await; + if should_acknowledge { + let count = read.expect("partial body should be acknowledged").unwrap(); + assert_eq!(&response[..count], b"HTTP/1.1 100 Continue\r\n\r\n"); + } else { + assert!( + read.is_err(), + "complete read-ahead should not be acknowledged" + ); + } } } - #[test] - fn parse_no_body() { - let headers = "GET /api HTTP/1.1\r\nHost: example.com\r\n\r\n"; - match parse_body_length(headers).unwrap() { - BodyLength::None => {} - other => panic!("Expected None, got {other:?}"), - } + #[tokio::test] + async fn collect_chunked_body_decodes_payload_bytes() { + let mut client = tokio::io::empty(); + let body = collect_chunked_body( + &mut client, + b"5\r\nhello\r\n6;ext=value\r\n world\r\n0\r\n\r\n", + None, + None, + ) + .await + .expect("chunked body should decode"); + + assert_eq!(body, b"hello world"); } - #[test] - fn parse_target_query_parses_duplicate_values() { - let (path, query) = parse_target_query("/download?tag=a&tag=b").expect("parse"); - assert_eq!(path, "/download"); - assert_eq!( - query.get("tag").cloned(), - Some(vec!["a".into(), "b".into()]) + #[tokio::test] + async fn middleware_chunked_request_with_trailers_is_rejected() { + let mut raw = b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\nTrailer: X-Checksum\r\n\r\n".to_vec(); + raw.extend_from_slice(b"5\r\nhello\r\n0\r\nX-Checksum: digest\r\n\r\n"); + let req = L7Request { + action: "POST".into(), + target: "/api".into(), + query_params: HashMap::new(), + raw_header: raw, + body_length: BodyLength::Chunked, + }; + + let error = + buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + .await + .expect_err("middleware 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}" ); } - #[test] - fn parse_target_query_decodes_percent_and_plus() { - let (path, query) = parse_target_query("/download?slug=my%2Fskill&name=Foo+Bar").unwrap(); - assert_eq!(path, "/download"); - assert_eq!( - query.get("slug").cloned(), - Some(vec!["my/skill".to_string()]) - ); - // `+` is decoded as space per application/x-www-form-urlencoded. - // Literal `+` should be sent as `%2B`. - assert_eq!( - query.get("name").cloned(), - Some(vec!["Foo Bar".to_string()]) + #[tokio::test] + async fn credential_rewrite_rejects_aggregate_chunk_extension_overflow() { + let mut wire = Vec::new(); + let chunk = format!("1;pad={}\r\nx\r\n", "a".repeat(64)); + while wire.len() <= MAX_REWRITE_BODY_BYTES { + wire.extend_from_slice(chunk.as_bytes()); + } + wire.extend_from_slice(b"0\r\n\r\n"); + + 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\n\r\n"; + let result = collect_and_rewrite_request_body( + &req, + &mut tokio::io::empty(), + headers, + std::str::from_utf8(headers).expect("headers"), + &wire, + None, + None, + ) + .await; + let Err(error) = result else { + panic!("aggregate chunk extensions must be bounded") + }; + assert!( + error + .to_string() + .contains("chunked body wire representation exceeds configured buffer limit"), + "unexpected error: {error}" ); } - #[test] - fn parse_target_query_literal_plus_via_percent_encoding() { - let (_path, query) = parse_target_query("/search?q=a%2Bb").unwrap(); - assert_eq!( - query.get("q").cloned(), - Some(vec!["a+b".to_string()]), - "%2B should decode to literal +" + #[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}" ); } - #[test] - fn parse_target_query_empty_value() { - let (_path, query) = parse_target_query("/api?tag=").unwrap(); - assert_eq!( - query.get("tag").cloned(), - Some(vec!["".to_string()]), - "key with empty value should produce empty string" + #[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 ); } - #[test] - fn parse_target_query_key_without_value() { - let (_path, query) = parse_target_query("/api?verbose").unwrap(); - assert_eq!( - query.get("verbose").cloned(), - Some(vec!["".to_string()]), - "key without = should produce empty string value" - ); + #[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 } + )); } - #[test] - fn parse_target_query_unicode_after_decoding() { - // "café" = c a f %C3%A9 - let (_path, query) = parse_target_query("/search?q=caf%C3%A9").unwrap(); - assert_eq!( - query.get("q").cloned(), - Some(vec!["café".to_string()]), - "percent-encoded UTF-8 should decode correctly" - ); + #[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); } - #[test] - fn parse_target_query_empty_query_string() { - let (path, query) = parse_target_query("/api?").unwrap(); - assert_eq!(path, "/api"); + #[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!( - query.is_empty(), - "empty query after ? should produce empty map" + matches!(error, CollectChunkedError::OverCapacity), + "over-cap wire body must be OverCapacity, got {error:?}" ); } - #[test] - fn parse_target_query_rejects_malformed_percent_encoding() { - let err = parse_target_query("/download?slug=bad%2").expect_err("expected parse error"); - assert!( - err.to_string().contains("percent-encoding"), - "unexpected error: {err}" - ); + #[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); } - /// SEC-009: Reject requests with both Content-Length and Transfer-Encoding - /// to prevent CL/TE request smuggling (RFC 7230 Section 3.3.3). - #[test] - fn reject_dual_content_length_and_transfer_encoding() { - let headers = "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 5\r\nTransfer-Encoding: chunked\r\n\r\n"; + #[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!( - parse_body_length(headers).is_err(), - "Must reject request with both CL and TE" + err.to_string().contains("Invalid chunk size token"), + "unexpected error: {err}" ); - } - - /// SEC-009: Same rejection regardless of header order. - #[test] - fn reject_dual_transfer_encoding_and_content_length() { - let headers = "POST /api HTTP/1.1\r\nHost: x\r\nTransfer-Encoding: chunked\r\nContent-Length: 5\r\n\r\n"; assert!( - parse_body_length(headers).is_err(), - "Must reject request with both TE and CL" + !err.to_string().contains("over_capacity") + && !err.to_string().contains("exceeds configured buffer limit"), + "protocol errors must not be reported as over-capacity: {err}" ); } - /// SEC: Reject differing duplicate Content-Length headers. - #[test] - fn reject_differing_duplicate_content_length() { - let headers = - "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 0\r\nContent-Length: 50\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) + .await + .expect("over-capacity is a BufferResult, not an Err"); + assert!( - parse_body_length(headers).is_err(), - "Must reject differing duplicate Content-Length" + matches!(result, BufferResult::OverCapacity { recoverable: false }), + "expected OverCapacity, got {result:?}" ); } - /// SEC: Accept identical duplicate Content-Length headers. - #[test] - fn accept_identical_duplicate_content_length() { - let headers = - "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 42\r\nContent-Length: 42\r\n\r\n"; - match parse_body_length(headers).unwrap() { - BodyLength::ContentLength(42) => {} - other => panic!("Expected ContentLength(42), got {other:?}"), - } - } + #[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, + }; - /// SEC: Reject non-numeric Content-Length values. - #[test] - fn reject_non_numeric_content_length() { - let headers = "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: abc\r\n\r\n"; - assert!( - parse_body_length(headers).is_err(), - "Must reject non-numeric Content-Length" - ); - } + 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"); - /// SEC: Reject when second Content-Length is non-numeric (bypass test). - #[test] - fn reject_valid_then_invalid_content_length() { - let headers = - "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 42\r\nContent-Length: abc\r\n\r\n"; assert!( - parse_body_length(headers).is_err(), - "Must reject when any Content-Length is non-numeric" + err.to_string().contains("no body framing"), + "unexpected error: {err}" ); } - /// SEC: Unsupported transfer codings must not be silently treated as no body. - #[test] - fn reject_unsupported_transfer_encoding_sequences() { - for value in [ - "gzip", - "gzip, chunked", - "chunked, gzip", - "chunked, chunked", - "chunkedx", - "chunked,", - ] { - let headers = - format!("POST /api HTTP/1.1\r\nHost: x\r\nTransfer-Encoding: {value}\r\n\r\n"); - assert!( - parse_body_length(&headers).is_err(), - "unsupported transfer coding must be rejected: {value}" - ); + #[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) + .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:?}") + } } } - #[test] - fn reject_multiple_chunked_transfer_encoding_fields() { - let headers = "POST /api HTTP/1.1\r\nHost: x\r\nTransfer-Encoding: chunked\r\nTransfer-Encoding: chunked\r\n\r\n"; - assert!(parse_body_length(headers).is_err()); + /// 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, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await; + assert!(result.is_err(), "Must reject headers with bare LF"); } - #[test] - fn reject_content_length_with_unsupported_transfer_encoding() { - let headers = - "POST /api HTTP/1.1\r\nHost: x\r\nContent-Length: 4\r\nTransfer-Encoding: gzip\r\n\r\n"; - assert!(parse_body_length(headers).is_err()); + /// 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"); } - #[test] - fn strip_connection_nominated_headers_before_forwarding() { - let raw = b"GET /api HTTP/1.1\r\nHost: x\r\nX-Guard: hidden\r\nConnection: keep-alive, x-guard\r\nKeep-Alive: timeout=5\r\nX-Visible: yes\r\n\r\n"; - let sanitized = - strip_connection_nominated_headers(raw, false).expect("sanitize request headers"); - let sanitized = String::from_utf8(sanitized).unwrap(); + #[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(), + ), + ]; - assert!(!sanitized.to_ascii_lowercase().contains("x-guard:")); - assert!(!sanitized.to_ascii_lowercase().contains("keep-alive:")); - assert!(!sanitized.to_ascii_lowercase().contains("connection:")); - assert!(sanitized.contains("X-Visible: yes\r\n")); + 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"); + } } - #[test] - fn connection_sanitization_preserves_only_websocket_upgrade_exception() { - let raw = b"GET /ws HTTP/1.1\r\nHost: x\r\nUpgrade: h2c\r\nUpgrade: h2c, websocket\r\nConnection: keep-alive, Upgrade, x-guard\r\nX-Guard: hidden\r\n\r\n"; - let sanitized = - strip_connection_nominated_headers(raw, true).expect("sanitize websocket headers"); - let sanitized = String::from_utf8(sanitized).unwrap(); - - assert_eq!(sanitized.matches("Upgrade: websocket\r\n").count(), 1); - assert_eq!(sanitized.matches("Connection: Upgrade\r\n").count(), 1); - assert!(!sanitized.to_ascii_lowercase().contains("upgrade: h2c")); - assert!(!sanitized.to_ascii_lowercase().contains("x-guard:")); - assert!(!sanitized.contains("keep-alive")); + /// 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_fixed_read_ahead_consumes_expect_continue() { - for (already_read, remaining, should_acknowledge) in [ - (b"hello".as_slice(), b"".as_slice(), false), - (b"he".as_slice(), b"llo".as_slice(), true), - ] { - let mut raw = b"POST /api HTTP/1.1\r\nHost: example.com\r\nContent-Length: 5\r\nExpect: 100-continue\r\n\r\n".to_vec(); - raw.extend_from_slice(already_read); - let req = L7Request { - action: "POST".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw, - body_length: BodyLength::ContentLength(5), - }; - let (mut client, mut peer) = tokio::io::duplex(128); - peer.write_all(remaining).await.unwrap(); - - let result = buffer_request_body_for_middleware(&req, &mut client, None, 1024) + 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 - .expect("fixed body should buffer"); - let BufferResult::Buffered(buffered) = result else { - panic!("fixed body unexpectedly exceeded capacity") - }; - assert_eq!(buffered.body, b"hello"); - assert!( - !String::from_utf8_lossy(&buffered.headers) - .to_ascii_lowercase() - .contains("expect:") - ); - - let mut response = [0u8; 64]; - let read = tokio::time::timeout( - std::time::Duration::from_millis(20), - peer.read(&mut response), - ) - .await; - if should_acknowledge { - let count = read.expect("partial body should be acknowledged").unwrap(); - assert_eq!(&response[..count], b"HTTP/1.1 100 Continue\r\n\r\n"); - } else { - assert!( - read.is_err(), - "complete read-ahead should not be acknowledged" - ); - } - } + .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" + ); } #[tokio::test] - async fn middleware_chunked_read_ahead_consumes_expect_continue() { - for (already_read, remaining, should_acknowledge) in [ - (b"5\r\nhello\r\n0\r\n\r\n".as_slice(), b"".as_slice(), false), - (b"5\r\nhe".as_slice(), b"llo\r\n0\r\n\r\n".as_slice(), true), - ] { - let mut raw = b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\nExpect: 100-continue\r\n\r\n".to_vec(); - raw.extend_from_slice(already_read); - let req = L7Request { - action: "POST".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw, - body_length: BodyLength::Chunked, - }; - let (mut client, mut peer) = tokio::io::duplex(128); - peer.write_all(remaining).await.unwrap(); + 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 result = buffer_request_body_for_middleware(&req, &mut client, None, 1024) - .await - .expect("chunked body should buffer"); - let BufferResult::Buffered(buffered) = result else { - panic!("chunked body unexpectedly exceeded capacity") - }; - assert_eq!(buffered.body, b"hello"); + let error = parse_http_request( + &mut client, + &crate::l7::path::CanonicalizeOptions::default(), + ) + .await + .expect_err("absolute-form authority mismatch must fail closed"); + assert!( + error + .to_string() + .contains("request authority does not match the Host header"), + "{error}" + ); + } + + #[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!( - !String::from_utf8_lossy(&buffered.headers) - .to_ascii_lowercase() - .contains("expect:") + 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 mut response = [0u8; 64]; - let read = tokio::time::timeout( - std::time::Duration::from_millis(20), - peer.read(&mut response), - ) - .await; - if should_acknowledge { - let count = read.expect("partial body should be acknowledged").unwrap(); - assert_eq!(&response[..count], b"HTTP/1.1 100 Continue\r\n\r\n"); - } else { - assert!( - read.is_err(), - "complete read-ahead should not be acknowledged" - ); - } + 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 collect_chunked_body_decodes_payload_bytes() { - let mut client = tokio::io::empty(); - let body = collect_chunked_body( - &mut client, - b"5\r\nhello\r\n6;ext=value\r\n world\r\n0\r\n\r\n", - None, - None, + 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 - .expect("chunked body should decode"); + .unwrap(); - assert_eq!(body, b"hello world"); + 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_request_with_trailers_is_rejected() { - let mut raw = b"POST /api HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\nTrailer: X-Checksum\r\n\r\n".to_vec(); - raw.extend_from_slice(b"5\r\nhello\r\n0\r\nX-Checksum: digest\r\n\r\n"); - let req = L7Request { - action: "POST".into(), - target: "/api".into(), - query_params: HashMap::new(), - raw_header: raw, - body_length: BodyLength::Chunked, - }; - - let error = - buffer_request_body_for_middleware(&req, &mut tokio::io::empty(), None, 64 * 1024) + 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_err("middleware 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}" + .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 credential_rewrite_rejects_aggregate_chunk_extension_overflow() { - let mut wire = Vec::new(); - let chunk = format!("1;pad={}\r\nx\r\n", "a".repeat(64)); - while wire.len() <= MAX_REWRITE_BODY_BYTES { - wire.extend_from_slice(chunk.as_bytes()); - } - wire.extend_from_slice(b"0\r\n\r\n"); + 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(), + ) + .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", + ); + } - 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\n\r\n"; - let result = collect_and_rewrite_request_body( - &req, - &mut tokio::io::empty(), - headers, - std::str::from_utf8(headers).expect("headers"), - &wire, - None, - None, + #[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; - let Err(error) = result else { - panic!("aggregate chunk extensions must be bounded") - }; assert!( - error - .to_string() - .contains("chunked body wire representation exceeds configured buffer limit"), - "unexpected error: {error}" + result.is_err(), + "a target that escapes the path root must be rejected at the parser" + ); + } + + #[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"); + } + + #[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" ); } #[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, + 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; - let Err(error) = result else { - panic!("credential rewriting must reject non-empty chunked trailers") - }; + .await + .unwrap() + .unwrap(); + assert_eq!(req.target, "/a/b"); assert!( - error.to_string().contains( - "chunked request trailers are not supported when buffering or transforming request bodies" - ), - "unexpected error: {error}" + 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 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( + 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, - &[], - None, - Some(openshell_supervisor_middleware::MAX_MIDDLEWARE_PAYLOAD_BYTES), + &crate::l7::path::CanonicalizeOptions::default(), ) .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 + .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()]) ); } + /// 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 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); + async fn parse_http_request_does_not_overread_next_request() { + let (mut client, mut writer) = tokio::io::duplex(4096); - let result = buffer_request_body_for_middleware( - &req, + 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, - None, - openshell_supervisor_middleware::MAX_MIDDLEWARE_PAYLOAD_BYTES, + &crate::l7::path::CanonicalizeOptions::default(), ) .await - .expect("oversized body should produce a capacity result"); + .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" + ); - assert!(matches!( - result, - BufferResult::OverCapacity { recoverable: true } + 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")); } - #[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); + #[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); + } - 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"); + #[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" + )); + } - assert_eq!(body.len(), payload_len); + #[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)); } #[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 response_middleware_unbound_chains_preserve_ordinary_headers() { + let mut many_headers = b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n".to_vec(); + for _ in 0..129 { + many_headers.extend_from_slice(b"Set-Cookie: a=b\r\n"); + } + many_headers.extend_from_slice(b"\r\n"); + let opaque_headers = + b"HTTP/1.1 200 OK\r\nContent-Length: 3\r\nX-Opaque: \xff\xfe\r\n\r\nabc".to_vec(); + let (_, chain) = response_middleware_fixture(ResponseRelayScript::HeadersOnly); + let request_only_runner = + openshell_supervisor_middleware::ChainRunner::new(Arc::new(ResponseRelayService { + script: ResponseRelayScript::HeadersOnly, + request_only: true, + })); + let empty_runner = openshell_supervisor_middleware::ChainRunner::default(); + for response in [many_headers, opaque_headers] { + assert!(response.len() < MAX_HEADER_BYTES); + for context in [ + None, + Some(response_middleware_context(&empty_runner, &[], "GET")), + Some(response_middleware_context( + &request_only_runner, + &chain, + "GET", + )), + ] { + let mut upstream = response.as_slice(); + let mut delivered = Vec::new(); + let outcome = relay_response( + "GET", + &mut upstream, + &mut delivered, + RelayResponseOptions::default(), + context, + ) + .await + .unwrap(); + assert!(matches!(outcome, RelayOutcome::Reusable)); + assert_eq!(delivered, response); + } + } + } + + #[tokio::test] + async fn response_middleware_selected_hook_enforces_header_limits() { + let mut response = b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n".to_vec(); + for _ in 0..129 { + response.extend_from_slice(b"Set-Cookie: a=b\r\n"); + } + response.extend_from_slice(b"\r\n"); + let (runner, chain) = response_middleware_fixture(ResponseRelayScript::HeadersOnly); + let mut upstream = response.as_slice(); + let mut delivered = Vec::new(); + let outcome = relay_response( + "GET", + &mut upstream, + &mut delivered, + RelayResponseOptions::default(), + Some(response_middleware_context(&runner, &chain, "GET")), + ) + .await + .unwrap(); + assert!(matches!(outcome, RelayOutcome::Consumed)); + assert!(delivered.starts_with(b"HTTP/1.1 502 ")); + } - let error = - collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) + #[tokio::test] + async fn response_middleware_flushes_partial_framed_payload_promptly() { + for chunked in [false, true] { + let (runner, chain) = response_middleware_fixture(ResponseRelayScript::Stream); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(8192); + // Small capacity forces the relay to complete partial downstream writes. + let (mut client_read, mut client_write) = tokio::io::duplex(7); + let head = if chunked { + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nContent-Type: text/event-stream\r\n\r\n1000\r\nabc".as_slice() + } else { + b"HTTP/1.1 200 OK\r\nContent-Length: 4096\r\nContent-Type: text/event-stream\r\n\r\nabc".as_slice() + }; + upstream_write.write_all(head).await.unwrap(); + let task = tokio::spawn(async move { + relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(response_middleware_context(&runner, &chain, "GET")), + ) .await - .expect_err("wire framing over the cap must be rejected"); + }); + let result = tokio::time::timeout(std::time::Duration::from_secs(2), async { + let mut head = Vec::new(); + while !head.ends_with(b"\r\n\r\n") { + head.push(client_read.read_u8().await.unwrap()); + } + let mut first = [0; 8]; + client_read.read_exact(&mut first).await.unwrap(); + assert_eq!(&first, b"3\r\nABC\r\n"); + // Complete the same framed payload only after its first transformed + // bytes have reached the consumer. + upstream_write.write_all(&vec![b'd'; 4093]).await.unwrap(); + if chunked { + for fragment in [b"\r".as_slice(), b"\n0\r", b"\n\r", b"\n"] { + upstream_write.write_all(fragment).await.unwrap(); + tokio::task::yield_now().await; + } + } + drop(upstream_write); + let mut rest = Vec::new(); + client_read.read_to_end(&mut rest).await.unwrap(); + assert!(rest.ends_with(b"0\r\n\r\n")); + let body = collect_chunked_body(&mut tokio::io::empty(), &rest, None, None) + .await + .unwrap(); + assert_eq!(body, vec![b'D'; 4093]); + }) + .await; + if result.is_err() { + task.abort(); + } + let relay = task.await; + assert!(result.is_ok(), "partial payload stalled, chunked={chunked}"); + assert!(relay.unwrap().is_ok()); + } + } - assert!( - matches!(error, CollectChunkedError::OverCapacity), - "over-cap wire body must be OverCapacity, got {error:?}" - ); + fn response_middleware_fixture( + script: ResponseRelayScript, + ) -> ( + openshell_supervisor_middleware::ChainRunner, + Vec, + ) { + response_middleware_fixture_with_error( + script, + openshell_supervisor_middleware::OnError::FailClosed, + ) } - #[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"); + 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, + request_only: false, + })); + 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: DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT, + } + } + + async fn run_response_middleware_relay( + response: &[u8], + method: &str, + script: ResponseRelayScript, + ) -> (Result, Vec) { + run_response_middleware_relay_with_error( + response, + method, + script, + openshell_supervisor_middleware::OnError::FailClosed, + ) + .await + } - let body = collect_chunked_body(&mut tokio::io::empty(), &wire, None, Some(max_body_bytes)) - .await - .expect("middleware cap should control chunked body collection"); + async fn run_response_middleware_relay_with_error( + response: &[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_mins(2), + ) + .await + } - assert_eq!(body.len(), payload_len); + async fn run_response_middleware_relay_with_timeout( + response: &[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); + let response = response.to_vec(); + 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; + drop(client_write); + let mut delivered = Vec::new(); + client_read.read_to_end(&mut delivered).await.unwrap(); + (outcome, delivered) } #[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 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!( - err.to_string().contains("Invalid chunk size token"), - "unexpected error: {err}" + delivered.contains("cache-control: private\r\n"), + "{delivered}" ); + assert!(delivered.contains("Content-Length: 5\r\n"), "{delivered}"); 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}" + !delivered.to_ascii_lowercase().contains("transfer-encoding"), + "{delivered}" ); + assert!(delivered.ends_with("\r\n\r\nhello"), "{delivered}"); } #[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) - .await - .expect("over-capacity is a BufferResult, not an Err"); - + 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!( - matches!(result, BufferResult::OverCapacity { recoverable: false }), - "expected OverCapacity, got {result:?}" + delivered.contains("Transfer-Encoding: chunked\r\n"), + "{delivered}" ); - } - - #[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) - .await - .expect_err("read-ahead leftovers must not become a request body"); - assert!( - err.to_string().contains("no body framing"), - "unexpected error: {err}" + delivered.ends_with("2;ext=yes\r\nhe\r\n3\r\nllo\r\n0\r\n\r\n"), + "{delivered}" ); - } - - #[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) - .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:?}") - } - } + 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: 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, - &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 headers with bare LF"); + 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}"); } - /// 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(), + 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; - assert!(result.is_err(), "Must reject headers with invalid UTF-8"); + 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 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(), - ), - ]; - - 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"); - } + 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; + 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); } - /// 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_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; - assert!(result.is_err(), "Must reject unsupported HTTP version"); + 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("HTTP/1.1 200 OK"), "{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_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 - .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!(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_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_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 error = 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_err("absolute-form authority mismatch must fail closed"); + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); assert!( - error - .to_string() - .contains("request authority does not match the Host header"), - "{error}" + String::from_utf8(delivered) + .unwrap() + .contains("response_delivery_failed") ); } - #[test] - fn origin_form_targets_with_embedded_urls_use_host_authority() { - let host: http::uri::Authority = "api.example.test".parse().unwrap(); + #[tokio::test] + 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 { + 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 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; + 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(); + let (_, body) = delivered.split_once("\r\n\r\n").unwrap(); + let decoded = collect_chunked_body(&mut tokio::io::empty(), body.as_bytes(), None, None) + .await + .unwrap(); + assert_eq!(decoded, b"hello"); + assert!(!delivered.contains("whole:hello"), "{delivered}"); + } - for target in ["/fetch/http://example.test", "/?next=http://example.test"] { - assert!( - absolute_form_uri(target).unwrap().is_none(), - "{target} must remain origin-form" + #[tokio::test(start_paused = true)] + async fn response_middleware_expiry_preserves_bytes_through_slow_stream_and_client() { + for chunked in [true, false] { + let (runner, mut chain) = response_middleware_fixture_with_error( + ResponseRelayScript::SlowStream, + openshell_supervisor_middleware::OnError::FailOpen, ); - 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)) + let mut whole_body = chain[0].clone(); + whole_body.name = "whole-body".into(); + whole_body + .config + .fields + .insert("whole_body".into(), prost_types::Value::default()); + chain[0].order = 1; + chain.insert(0, whole_body); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(8192); + // Force write_all to make partial progress before each wait. + let (mut client_read, mut client_write) = tokio::io::duplex(7); + let producer = async move { + let head = if chunked { + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n800\r\n".as_slice() + } else { + b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".as_slice() + }; + upstream_write.write_all(head).await.unwrap(); + upstream_write.write_all(&vec![b'a'; 2048]).await.unwrap(); + if chunked { + upstream_write.write_all(b"\r\n").await.unwrap(); + } + // The first coalesced unit belongs to the whole-body stage. A new + // partial unit starts coalescing just before its deadline. + tokio::time::sleep(std::time::Duration::from_millis(9)).await; + upstream_write + .write_all(if chunked { b"1\r\nb\r\n" } else { b"b" }) + .await + .unwrap(); + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + if chunked { + upstream_write.write_all(b"0\r\n\r\n").await.unwrap(); + } + upstream_write.shutdown().await.unwrap(); + }; + let relay = async { + let mut context = response_middleware_context(&runner, &chain, "GET"); + context.whole_body_timeout = std::time::Duration::from_millis(10); + let outcome = relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(context), + ) + .await; + drop(client_write); + outcome + }; + let consumer = async move { + let mut delivered = Vec::new(); + let mut bytes = [0; 7]; + loop { + let count = client_read.read(&mut bytes).await.unwrap(); + if count == 0 { + break; + } + delivered.extend_from_slice(&bytes[..count]); + tokio::time::sleep(std::time::Duration::from_millis(1)).await; + } + delivered + }; + let ((), outcome, delivered) = tokio::time::timeout( + std::time::Duration::from_secs(3), + Box::pin(async { tokio::join!(producer, relay, consumer) }), + ) + .await + .expect("response relay stalled"); + assert!(outcome.is_ok(), "{outcome:?}"); + assert!(delivered.starts_with(b"HTTP/1.1 200 OK\r\n")); + let head_end = delivered + .windows(4) + .position(|bytes| bytes == b"\r\n\r\n") .unwrap() - .expect("origin-form request with Host must have an authority"); - assert_eq!(authority.authority, host); - assert_eq!(authority.effective_port, 443); + + 4; + let mut wire = &delivered[head_end..]; + let mut body = Vec::new(); + loop { + let end = wire.windows(2).position(|bytes| bytes == b"\r\n").unwrap(); + let size = + usize::from_str_radix(std::str::from_utf8(&wire[..end]).unwrap(), 16).unwrap(); + wire = &wire[end + 2..]; + if size == 0 { + assert_eq!(wire, b"\r\n"); + break; + } + body.extend_from_slice(&wire[..size]); + assert_eq!(&wire[size..size + 2], b"\r\n"); + wire = &wire[size + 2..]; + } + let mut expected = vec![b'A'; 2048]; + expected.push(b'B'); + assert_eq!(body, expected, "chunked={chunked}"); } } #[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_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(); + .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}"); + } - 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 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 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 - .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" - ); + 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 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(), - ) - .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", - ); + 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 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_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 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_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 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); + 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; + 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}"); + } - 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", + #[tokio::test] + async fn response_middleware_unrepresentable_input_obeys_failure_policy() { + let mut many_headers = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n".to_vec(); + for _ in 0..=openshell_supervisor_middleware::MAX_MIDDLEWARE_HEADERS { + many_headers.extend_from_slice(b"X-Extra: value\r\n"); + } + many_headers.extend_from_slice(b"\r\nok"); + for response in [ + b"HTTP/1.1 200 OK\r\nX-Legacy: \xff\r\nContent-Length: 2\r\n\r\nok".as_slice(), + many_headers.as_slice(), + ] { + for on_error in [ + openshell_supervisor_middleware::OnError::FailOpen, + openshell_supervisor_middleware::OnError::FailClosed, + ] { + let (outcome, delivered) = run_response_middleware_relay_with_error( + response, + "GET", + ResponseRelayScript::HeadersOnly, + on_error, ) - .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" - ); + .await; + if on_error == openshell_supervisor_middleware::OnError::FailOpen { + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + assert_eq!(delivered, response); + } else { + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + assert!(delivered.starts_with(b"HTTP/1.1 502 Bad Gateway\r\n")); + assert!( + String::from_utf8(delivered) + .unwrap() + .contains("response_delivery_failed") + ); + } + } + } + } - 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()); + #[tokio::test] + async fn response_middleware_obs_text_does_not_bypass_unsafe_headers() { + for response in [ + b"HTTP/1.1 200 OK\r\nX-Legacy: \xff\r\nBad Name: value\r\nContent-Length: 2\r\n\r\nok".as_slice(), + b"HTTP/1.1 200 OK\r\nX-Legacy: \xff\x00\r\nContent-Length: 2\r\n\r\nok".as_slice(), + b"HTTP/1.1 200 OK\r\nX-Legacy: \xff\r\nTrailer: content-length\r\nContent-Length: 2\r\n\r\nok".as_slice(), + ] { + let (outcome, delivered) = run_response_middleware_relay_with_error( + response, "GET", ResponseRelayScript::HeadersOnly, + openshell_supervisor_middleware::OnError::FailOpen, + ).await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + assert!(delivered.starts_with(b"HTTP/1.1 502 Bad Gateway\r\n")); + } } - #[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 +8542,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5569,6 +8588,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5619,6 +8639,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5663,6 +8684,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5704,6 +8726,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5742,6 +8765,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5782,6 +8806,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5826,6 +8851,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5869,6 +8895,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5905,6 +8932,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5951,6 +8979,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5998,6 +9027,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -7949,4 +10979,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/proxy.rs b/crates/openshell-supervisor-network/src/proxy.rs index 70c90b0158..58579d8b58 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_mins(1), async { @@ -7685,7 +7833,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"); @@ -7782,7 +7930,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 { .. } => { @@ -7896,6 +8047,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, @@ -7980,6 +8296,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await?; @@ -8245,6 +8562,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await?; @@ -11101,7 +11419,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 { @@ -11403,6 +11721,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await @@ -11484,6 +11803,7 @@ network_policies: signing_region: "us-west-2", host: "api.example.com", port: 80, + response_middleware: None, }, ) .await @@ -11573,6 +11893,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await; @@ -11623,6 +11944,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await; diff --git a/docs/extensibility/gateway-interceptors.mdx b/docs/extensibility/gateway-interceptors.mdx index bf9656a5ed..bce7c31455 100644 --- a/docs/extensibility/gateway-interceptors.mdx +++ b/docs/extensibility/gateway-interceptors.mdx @@ -5,6 +5,7 @@ title: "Gateway Interceptors" sidebar-title: "Gateway Interceptors" description: "Extend OpenShell gateway operations with deployment-specific governance and business logic." keywords: "Generative AI, Cybersecurity, AI Agents, Gateway Interceptors, Extensibility, Governance" +position: 2 --- Gateway interceptors let operators add deployment-specific governance to OpenShell control-plane operations without modifying the gateway. An external gRPC service can modify or validate selected API writes before the gateway handles them, then observe successful responses after commit. diff --git a/docs/extensibility/overview.mdx b/docs/extensibility/overview.mdx new file mode 100644 index 0000000000..be9f14fe80 --- /dev/null +++ b/docs/extensibility/overview.mdx @@ -0,0 +1,29 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "Extensibility" +sidebar-title: "Overview" +description: "Add custom checks to sandbox traffic and gateway operations." +keywords: "Supervisor Middleware, Gateway Interceptors, Extensibility" +position: 0 +--- + +OpenShell lets you add custom checks and transformations to sandbox traffic and gateway API operations. You can use these extensions to connect your organization's content checks, policy rules, or audit services to OpenShell. + +Choose an extension based on what you need to control. Supervisor middleware handles network traffic between an agent and external services. Gateway interceptors handle selected API operations that create or change OpenShell resources. You can use both in the same deployment. + +## Supervisor middleware + +Use supervisor middleware when you need to check the content an agent sends or receives. For example, you might redact recognized API tokens from outgoing requests, block prohibited content, or remove sensitive content from an HTTP response before the agent sees it. + +Middleware runs in the sandbox's network request and response flow. You select destination hosts in sandbox policy and choose a built-in implementation or a service you operate. + +Start with [Supervisor Middleware](/extensibility/supervisor-middleware) to learn how it works and choose a setup or protocol guide. + +## Gateway interceptors + +Use gateway interceptors when you need to enforce rules on how people and applications manage OpenShell resources. For example, you might apply an approved policy to new sandboxes, reject unauthorized policy changes, or report completed operations to an audit service. + +Interceptors run as external services that the gateway calls for selected API operations. They can modify or reject an operation before it takes effect, or observe it after it succeeds. + +Start with [Gateway Interceptors](/extensibility/gateway-interceptors) to choose operations and connect an interceptor service. diff --git a/docs/extensibility/supervisor-middleware.mdx b/docs/extensibility/supervisor-middleware.mdx deleted file mode 100644 index 9f3c723b16..0000000000 --- a/docs/extensibility/supervisor-middleware.mdx +++ /dev/null @@ -1,234 +0,0 @@ ---- -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -title: "Supervisor Middleware" -sidebar-title: "Supervisor Middleware" -description: "Configure and operate built-in and operator-run middleware for sandbox HTTP requests and WebSocket messages." -keywords: "Generative AI, Cybersecurity, AI Agents, Supervisor Middleware, Extensibility, Request Filtering" ---- - -Supervisor middleware adds ordered processing stages to allowed HTTP and WebSocket egress. Middleware runs after network and L7 policy admit traffic and before OpenShell injects provider credentials. A stage can allow or deny an HTTP request or client WebSocket text message, replace its payload, add approved HTTP headers, and report audit-safe findings. - -Middleware selection is independent of the network policy rule that admitted the request. OpenShell matches middleware by destination host, so the same middleware applies consistently across broad, specific, user-authored, and provider-derived network policies. - -## Request Flow - -For each inspected HTTP request, the supervisor: - -1. Evaluates network and L7 policy. -2. Selects middleware whose host selectors match the admitted destination. -3. Buffers the request body using the largest body limit in the selected chain. -4. Runs matching middleware by ascending `order`. Policy validation rejects duplicate order values. -5. Re-checks body-aware protocol policy (GraphQL, JSON-RPC, MCP) after each stage that replaces the body. Every middleware receives a payload the policy admits, and a transformation cannot smuggle a denied or unparseable operation to a later stage or the upstream. -6. Applies allowed transformations, injects provider credentials, and forwards the request. - -For an RFC 6455 upgrade over `ws://` or `wss://`, the supervisor first finds every host-matched attachment, then selects only implementations that advertise `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. It opens one ordered, phase-specific `EvaluateWebSocketSession` stream per selected stage. OpenShell sends `WebSocketSessionEvent` values, while the service returns `WebSocketSessionEventResult` values only for preflight and message events; session start and end are notifications. Future upstream-to-client inspection uses the same RPC with `PRE_RETURN`; an implementation that advertises both phases receives two independent streams for the WebSocket session. An attachment without the selected binding can still inspect the HTTP upgrade request when it advertises the HTTP binding, but it is not a failed WebSocket stage. OpenShell allows post-upgrade traffic and emits an informational `binding_not_selected` coverage event for that attachment. - -1. A preflight before the upgrade is sent upstream. The stage chooses `INSPECT`, voluntary `SKIP`, or authoritative `DENY` and may return a bounded diagnostic reason, stable reason code, findings, and metadata. OpenShell runs selected preflights concurrently; any `DENY` rejects the upgrade regardless of `on_error`. -2. A session-start event after the upstream accepts the upgrade, including the negotiated subprotocol. -3. Complete client-to-upstream text messages in sequence order. OpenShell reassembles fragmented messages and decompresses negotiated `permessage-deflate` messages before evaluation. -4. A best-effort session-end event when the stage stream remains writable. OpenShell attempts at most one terminal event for each opened stream, including streams opened during a preflight that rejects the upgrade before session start. It then half-closes the request stream and briefly drains the response stream so the terminal event can leave the local transport before the RPC closes. Middleware services should finish their response stream after the request stream reaches EOF. - -The protobuf represents each logical message with a `text` or `binary` payload variant. Text uses the protobuf `string` type, so invalid UTF-8 cannot enter the middleware contract. Results use an optional matching replacement variant: absence preserves the input, while presence represents a replacement even when its content is empty. OpenShell rejects attempts to change the message type. Allowed replacements are re-framed, re-compressed when required, and forwarded. Binary messages, control frames, and upstream-to-client traffic remain uninspected. Binary messages pass through under both `on_error` modes. For each active selected stage, OpenShell emits an informational `unsupported_message_type` coverage event and advances the session-global sequence; the next text message can therefore reach the stage with a valid sequence gap. - -The network supervisor reserves process-wide assembly capacity before buffering every parsed WebSocket text message, even when no middleware is selected. At most 32 assemblies run while 64 additional callers wait without buffering payload bytes. When both bounds are full, OpenShell closes the WebSocket with code `1013` before reading the new message payload. A text message may contain at most 4,096 fragments, must make input progress within 30 seconds, and must finish assembly within 2 minutes. Forwarding the completed text frame must finish within another 2 minutes. The assembly budget lasts for the supervisor process lifetime, so policy reloads do not reset its capacity. - -Active middleware sessions additionally reserve shared middleware capacity before buffering WebSocket text, and HTTP middleware reserves the same capacity before buffering request bodies; at most 32 evaluations run and 64 additional unbuffered callers wait for capacity. When both middleware bounds are full, OpenShell sheds an HTTP request with `503 Service Unavailable` before reading its body. Persistent middleware streams use a separate process-wide budget of 32 sessions. WebSocket session admission does not wait: if the budget is full, OpenShell applies each selected config's `on_error` behavior before opening a stream. - -Because each transformed body is re-checked before the next stage runs, a middleware hook always receives a request that satisfies the sandbox policy. A stage whose output the policy rejects stops the chain; under `enforcement: audit` the rejection is logged and the request proceeds. - -If post-transformation policy evaluation itself fails, OpenShell denies the request and emits a high-severity detection finding. This failure is separate from middleware `on_error` because the middleware completed successfully; the sandbox policy could not validate its output. - -Middleware receives the request before credential injection. Operator-run services cannot inspect OpenShell-managed credentials. Middleware-visible request headers are delivered in wire order and repeated header names are preserved as separate entries. OpenShell filters credential, routing, framing, and hop-by-hop headers before invoking middleware. It rejects malformed request headers and unsupported transfer-coding sequences before middleware or policy dispatch. Headers named by a request's `Connection` field are omitted from middleware input and removed before forwarding, except for the validated WebSocket upgrade pair. - -The request context identifies the originating sandbox to operator-run services. It carries the sandbox ID (`sandbox_id`), the sandbox name (`sandbox_name`), and the workspace (`workspace`), letting audit and approval interfaces show a human-readable name and its workspace instead of an opaque ID. `sandbox_name` and `workspace` are for display and logging only: names are workspace-scoped and may be reused for different sandbox instances, so services must use `sandbox_id` for authorization, persistence, durable correlation, and identity. `sandbox_id` is always present on middleware requests. `sandbox_name` and `workspace` are best-effort: a supervisor that cannot resolve a value, or an older supervisor that predates a field, sends an empty string. Services should fall back to the sandbox ID when the name or workspace is empty. - -## Choose a Middleware Type - -| Type | Registration | Payload limit | Deployment | -| --- | --- | --- | --- | -| Built-in | None | Defined by OpenShell | Runs inside the supervisor | -| Operator-run service | Required in gateway TOML | Set by the operator, up to the service capability | Runs as a separate service reachable by the gateway and supervisors | - -`openshell/regex` is an example built-in middleware. It replaces only simple, self-contained token patterns in UTF-8 HTTP bodies and client WebSocket text messages; the initial pattern recognizes `sk-` tokens. It does not infer values from keyword assignments such as JSON `password` fields. This best-effort text transformation is not parser-aware and does not guarantee that it will detect or fully remove sensitive values. Its `config` accepts one field, `mode: redact`, which is also the default when the field is omitted. Unknown config fields and non-string values are rejected at policy validation. Custom expressions are not configurable yet. - -Operator-run services expose bindings for supported operation and phase pairs. A binding is identified by its operation and phase. V1 supports `HttpRequest/pre_credentials` and `WebSocketMessage/pre_credentials`; a service may expose either or both. Policies attach the complete middleware by its operator-owned gateway registration name. - -## Register a Middleware Service - -Start an operator-run service before starting the gateway, then add a registration to the local gateway TOML: - -```toml -[[openshell.supervisor.middleware]] -name = "local-content-guard" -grpc_endpoint = "https://content-guard.example:50051" -tls_ca_cert_path = "/etc/openshell/content-guard-ca.pem" -audience = "urn:example:content-guard" -max_payload_bytes = 262144 -timeout = "500ms" -``` - -| Field | Description | -| --- | --- | -| `name` | Operator-owned registration name used by policy attachments and diagnostics. Names must be unique, and `openshell/` is reserved for built-ins. | -| `grpc_endpoint` | Service address reachable from both the gateway and sandbox supervisors. Authenticated extensions use TLS `https://`. | -| `tls_ca_cert_path` | Optional PEM trust roots for a private HTTPS service. Custom roots replace platform roots and retain hostname verification. | -| `audience` | Exact audience expected by the service. Defaults to `urn:openshell:extension:middleware:`. | -| `allow_insecure_transport` | Opt this registration out of extension authentication, permitting a plaintext `http://` endpoint with no bearer credential. Defaults to `false`. Development and trusted-network deployments only. | -| `max_payload_bytes` | Shared operator limit applied to inspectable logical payloads across every binding exposed by the service, up to the 4 MiB platform maximum. It caps HTTP bodies and complete WebSocket text messages. | -| `timeout` | Optional service-wide RPC timeout using an integer with an `ms` or `s` suffix. Defaults to `500ms`; valid values range from `10ms` through `30s`. | - -Each binding returned by `Describe` may advertise a shorter `timeout` using the same syntax and bounds. The operator-configured service timeout is a ceiling: OpenShell uses the smaller of the binding and service values. An omitted binding timeout inherits the service setting, and an omitted service setting uses the 500 ms platform default. OpenShell rejects an invalid timeout before accepting the manifest. The operator-configured service timeout applies to `Describe` and `ValidateConfig`. The effective binding timeout applies only to `EvaluateHttpRequest`, WebSocket preflight, and each WebSocket message. WebSocket streams have no connection-wide deadline. - -The gateway connects to every registered service and verifies its capabilities before accepting traffic. Gateway startup fails when a service is unavailable, reports an invalid capability, or exposes more than one binding for the same operation and phase. The manifest `name` is diagnostic metadata and does not need to match the operator registration name. Operator-run registration names cannot claim the reserved `openshell/` namespace. - -Registration is static. Restart the gateway after adding, removing, or changing a service. See [Gateway Configuration](/reference/gateway-config#supervisor-middleware-services) for the complete gateway TOML context. - -### Authenticate OpenShell Callers - -When gateway JWT signing is configured, OpenShell attaches a short-lived EdDSA bearer token to every remote middleware RPC. Gateway calls use `caller_kind: gateway`; sandbox supervisor calls use `caller_kind: supervisor` and include the sandbox ID. Supervisors request credentials by registration name through `RefreshSandboxToken`. The gateway derives the audience from operator-owned configuration and authorizes each name against the sandbox's effective policy. - -Return your expected audience in the `expected_audience` field of your `Describe` manifest. After authenticated `Describe` succeeds, OpenShell compares the advertised value with its operator-configured audience and refuses to start when they differ. This is a post-authentication consistency assertion, not audience discovery: a strict verifier may reject an incorrect audience before returning the manifest, in which case startup reports an authentication failure. Leave the field empty to skip the consistency check. - -Provision the trusted gateway URL, expected gateway ID, and public key or JWKS through the deployment. This operator-provisioned key material is the authoritative cold-start trust anchor. The expected issuer is exactly `openshell-gateway:`; fetching JWKS does not establish that identity by itself. After initial trust is established, `GET /.well-known/openid-configuration` and its `jwks_uri` provide steady-state key refresh and operational convenience. The document is OIDC-shaped rather than OIDC-compliant: `issuer` is the gateway identity, not the URL serving the document, so compare `iss` against the configured value and fetch updates only over authenticated TLS at the trusted gateway URL. - -Cache keys by `kid`. Validate, at minimum: - -- `typ` is exactly `openshell-ext+jwt`. Extension tokens and sandbox-to-gateway bootstrap tokens share a signing key and differ only in audience; this header is a second, independent discriminator. -- `alg` is pinned to `EdDSA`. Never select the algorithm from the token. -- Signature, expected issuer, exact audience, and positive expiry. -- `caller_kind`, and the sandbox identity when your service scopes behavior per sandbox. - -A sandbox-to-gateway JWT is not an extension credential even though both token types use the same signing key. - -Each token carries a unique `jti` that identifies that token instance for correlation and future explicit revocation. OpenShell reuses a token across calls until rotation and does not track `jti`, so rejecting a repeated `jti` would reject legitimate requests. Per-request replay resistance requires a request nonce or signature, channel binding, or another proof-of-possession mechanism. - -### Run Without Extension Authentication - -Set `allow_insecure_transport = true` on a registration to keep a plaintext `http://` endpoint working. OpenShell then attaches no credential to that service, supervisors do not request one, and the gateway refuses to mint one if asked. The gateway logs a warning naming the registration at every startup. - -The service cannot distinguish OpenShell from any other client that can reach it. Use this only where the network already provides that guarantee, and prefer `https://` everywhere else. - -## Apply Middleware with Policy - -Add middleware configs to the top-level `network_middlewares` map. Each key is the policy-local config name: - -```yaml -network_middlewares: - regex-redactor: - name: Redact API tokens - middleware: openshell/regex - order: 10 - config: - mode: redact - on_error: fail_closed - endpoints: - include: ["*.example.com"] - exclude: ["trusted.example.com"] -``` - -Each config has a stable policy-local identity from its map key, an optional human-readable `name` that defaults to that key, a built-in or operator-owned registration name in `middleware`, an integer `order`, implementation-owned `config`, failure behavior, and host selectors. The optional name does not replace the map key for attachment or future keyed updates. A policy accepts at most 10 middleware configs. - -`include` selects destination hosts. `exclude` takes precedence and removes hosts from that selection. Each config accepts at most 32 combined include and exclude patterns. Matching is case-insensitive and uses the same exact-host and DNS glob behavior as network policy endpoints: `*` matches exactly one DNS label, `**` matches one or more labels, and intra-label patterns like `*-api.example.com` work. Brace alternates such as `{prod,staging}` are rejected at validation; list each host pattern separately. - -Matching configs run once each by ascending `order`; lower values run first. Order values must be unique across the complete policy, even when endpoint selectors do not overlap. The default order is `0`, so policies with multiple configs normally set explicit values. Different map keys may attach the same middleware and run as separate stages. Map keys are structurally unique. Runtime selection defensively rejects chains with more than 10 stages. - -See [Policy Schema](/reference/policy-schema#network-middleware) for the complete field reference. - -## Configure Failure Behavior - -`on_error` controls what happens after an operation binding is selected and middleware is unavailable, rejects its configuration, returns an invalid result, or exceeds the selected binding's payload limit. It does not turn an unadvertised operation or an unsupported WebSocket message class into a middleware failure. - -| Value | Behavior | -| --- | --- | -| `fail_closed` | Denies the HTTP request or closes the WebSocket when the stage fails. This is the default. | -| `fail_open` | Skips the failed HTTP stage. For a broken WebSocket stage stream, disables that stage for the rest of the connection and continues the remaining chain. | - -Use `fail_open` only when bypassing the middleware preserves the intended security policy. OpenShell emits a detection finding when a failed stage is bypassed and a separate state-change finding when a WebSocket stage is disabled for the session. - -Capability coverage is separate from failure handling. A host-matched HTTP-only attachment does not join the WebSocket chain, regardless of `on_error`. Binary messages are outside the V1 text-message binding and pass through even when a selected stage is `fail_closed`. OpenShell records both states as informational coverage events so operators do not mistake pass-through traffic for inspected traffic. If a deployment requires all WebSocket message classes to be inspected, V1 cannot express that requirement. - -An explicit deny decision always stops the chain and denies the request or WebSocket upgrade, regardless of `on_error`. A WebSocket preflight `DENY` is a successful policy decision, not a middleware failure; OpenShell rejects the upgrade before upstream contact and ends each still-writable stream opened by a successful preflight decision with `MIDDLEWARE_DENIAL`. The HTTP response uses `error: middleware_denied`, identifies the policy-local middleware config, and omits policy-advisor remediation because the network and L7 allow rules already matched. OpenShell never copies the free-form middleware `reason` into the response or security logs. HTTP results, WebSocket preflight decisions, and WebSocket message results can instead return an optional stable `reason_code`: 1–64 bytes, starting with a lowercase ASCII letter and containing only lowercase ASCII letters, digits, and underscores. Invalid codes make the result a middleware failure governed by `on_error`. Preflight findings and metadata use the same bounds and audit-safe handling as message results. - -```json -{ - "error": "middleware_denied", - "detail": "Request rejected by configured middleware", - "policy": "api-policy", - "middleware": "prototype-content-guard", - "reason_code": "content_match" -} -``` - -A failed `fail_closed` stage uses `error: middleware_failed` and a platform-owned `detail`. It also omits `rule_missing`, `next_steps`, and `agent_guidance`: the failure did not result from a missing network or L7 policy rule, and changing policy cannot repair it. Runtime diagnostic text is available only through sanitized operator telemetry. - -Middleware decisions are enforced regardless of the endpoint's `enforcement` mode. `enforcement: audit` applies to an endpoint's network and L7 policy rules and does not bypass middleware: a middleware deny, or a failed `fail_closed` stage, blocks the request even on an audit endpoint. A middleware service that needs to observe traffic without blocking should return an allow decision with findings, which OpenShell emits as detection findings. - -## Set Payload Limits - -Every middleware binding declares the largest logical payload or replacement it supports through `max_payload_bytes`. For `HTTP_REQUEST`, that payload is one request body. For `WEBSOCKET_MESSAGE`, it is one complete message rather than the whole session. - -- Built-in middleware uses its OpenShell-defined limit. -- Each operator-run registration sets one `max_payload_bytes` ceiling no higher than any binding's advertised `max_payload_bytes` capability. -- A selected chain buffers using its largest stage limit, so every stage that can process the body receives it. -- The same per-stage limit applies to request bodies and replacement bodies. - -The gateway rejects a registration whose operator limit exceeds the service capability or the 4 MiB platform maximum instead of silently clamping it. OpenShell also bounds the non-payload protobuf components: 64 KiB for service config, 4 KiB for request context, 32 KiB for the target, and 128 request header lines totaling at most 64 KiB encoded. Results allow a 4 KiB discarded free-form reason, a 64-byte validated reason code, 64 header mutations totaling at most 64 KiB encoded, 32 findings of at most 4 KiB encoded each, and 64 metadata entries totaling at most 32 KiB. Middleware gRPC servers should configure request and response message limits to at least 4 MiB plus 293 KiB so every platform-valid envelope fits. - -At request time, exceeding a selected stage's limit is a middleware failure for that stage alone and follows that config's `on_error` behavior; other stages in the chain still run against their own limits. OpenShell can apply `fail_open` to an oversized `Content-Length` before consuming body bytes. A chunked body can cross the limit only after bytes have been consumed, so OpenShell denies that request because it cannot safely resume the original stream. - -For a WebSocket binding, `max_payload_bytes` covers complete client text messages and replacements. Exceeding a selected stage's effective text-message limit follows that stage's `on_error`. The 4 MiB parsed-text platform cap and other protocol-safety limits are independent of middleware failure policy. Binary messages are not delivered to middleware, so the operator ceiling does not become a binary relay limit; individual raw binary frames retain the 16 MiB relay-safety bound. Oversized parsed text closes the connection with code `1009`; invalid UTF-8 uses `1007`; protocol errors use `1002`; middleware or policy denials use `1008`; and policy reload uses `1012`. - -## Mutate Request Headers - -A middleware result can return ordered header mutations before OpenShell injects credentials. A `write` mutation adds a value when the case-insensitive header name is absent and selects one behavior when it is already present: - -- `append` adds another field value. -- `overwrite` removes every existing value before adding the new value. -- `skip` leaves existing values unchanged. - -A `remove` mutation removes every value for a case-insensitive header name. OpenShell applies each successful stage's mutations before invoking the next middleware, so later stages observe the accumulated header state. - -Writes and removals may target middleware-visible end-to-end request headers. Protected credential, routing, framing, and hop-by-hop headers are always rejected. Header values must not contain control characters or OpenShell credential placeholder syntax. Middleware runs before credential injection, but it cannot introduce a value that the later injection step would resolve. - -OpenShell validates and applies each stage's mutations atomically. An invalid operation discards every mutation from that stage and follows its `on_error` behavior. Built-in failures can name the offending header. Operator-run failures use a platform-owned error code so request-derived header text cannot reach logs or denied responses. - -## Operate Middleware Services - -Plan startup and updates around these boundaries: - -- Start registered services before the gateway. The gateway validates every registration during startup. -- Keep service endpoints reachable from both the gateway and sandbox supervisors. The supervisors call operator-run services directly on the request path. -- Restart the gateway after changing registrations. -- Keep required services available before creating or updating policies. The gateway validates implementation-owned config before persisting a policy. -- Treat `fail_open` as an explicit availability-over-enforcement decision. - -When the effective sandbox configuration changes, a running supervisor validates the new service registry before installing it. If the reload fails, the supervisor keeps its last-known-good registry and emits a configuration failure event. - -## Observe Middleware - -Middleware activity is emitted through OpenShell's OCSF logging: - -- Each invocation records its policy-local config name, attached middleware name, decision, transformation state, and failure state. -- A denied invocation records a platform-owned reason derived from the policy-local config name and optional validated reason code. OpenShell does not record service-provided free-form reason text. -- A bypass under `fail_open` emits a detection finding. -- A required stage that fails closed emits a high-severity detection finding. -- A host-matched attachment without a WebSocket binding emits an informational `binding_not_selected` coverage event. -- A binary message encountered by an active WebSocket stage emits an informational `unsupported_message_type` coverage event with message type, sequence, and byte count. It is not reported as an invocation or failure. -- Built-in findings include their type, label, and aggregate count. Operator-run findings use the operator-owned registration name and a platform label plus the aggregate count; OpenShell does not log service-provided finding text or diagnostic metadata. A stage can return at most 32 findings. Exceeding the per-stage cap is an invalid response handled through `on_error`. A maximum 10-stage chain retains and emits up to 320 findings without silently dropping findings from later stages. -- Registry reload success and failure are emitted as configuration state changes. - -See [Logging](/observability/logging) for log access and [OCSF JSON Export](/observability/ocsf-json-export) for structured export. - -## Current Limitations - -- Middleware applies only through operation bindings advertised by each implementation. For protocols that have no supported middleware operation at all, such as HTTP/2 prior knowledge or non-HTTP TCP, the existing uninspectable-traffic gate denies a host match containing `fail_closed` and relays an all-`fail_open` match with a detection finding. -- The typed operation and phase pairs are `HTTP_REQUEST/PRE_CREDENTIALS` and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. -- A host match does not imply every advertised operation: an HTTP-only attachment can inspect the upgrade GET, then post-upgrade traffic passes with `binding_not_selected` coverage. -- The V1 WebSocket binding inspects complete client text messages only. Binary messages pass with `unsupported_message_type` coverage for active stages; control frames and upstream-to-client messages remain outside the middleware operation. -- Selection uses destination host include and exclude patterns. -- A fail-closed middleware cannot cover `tls: skip` endpoints because OpenShell cannot inspect that traffic. An all-`fail_open` match may cover the endpoint; OpenShell bypasses the middleware and emits a detection finding. -- Operator-run services use TLS `https://` when gateway JWT signing is enabled, unless the registration sets `allow_insecure_transport`. Certificates must chain to the configured custom CA or platform roots, and the endpoint hostname must match. -- Extension tokens and sandbox-to-gateway tokens are signed by the same key. They are separated by audience and by `typ`, but the extension credential path cannot yet be rotated or revoked independently of sandbox admission. -- OpenShell does not track or revoke `jti`; bearer tokens can be replayed until expiry. Per-request replay resistance requires proof of possession or request binding. -- mTLS client authentication, health checks, runtime registration, and overlapping signing-key rotation are not available. diff --git a/docs/extensibility/supervisor-middleware/configure.mdx b/docs/extensibility/supervisor-middleware/configure.mdx new file mode 100644 index 0000000000..a69d2a8f08 --- /dev/null +++ b/docs/extensibility/supervisor-middleware/configure.mdx @@ -0,0 +1,126 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "Configure Supervisor Middleware" +sidebar-title: "Configure Middleware" +description: "Choose, register, and attach supervisor middleware to sandbox traffic." +keywords: "Supervisor Middleware, Configuration, Network Policy, Extension Authentication" +position: 2 +--- + +Configure middleware in two places. Register operator-run services in gateway TOML, then attach built-in or registered middleware to destination hosts in sandbox policy. + +## Choose a middleware type + +| Type | Registration | Payload limit | Deployment | +| --- | --- | --- | --- | +| Built-in | None. | Defined by OpenShell. | Runs inside the supervisor. | +| Operator-run service | Required in gateway TOML. | Set by the operator, up to the service capability. | Runs as a separate service reachable by the gateway and supervisors. | + +`openshell/regex` is an example built-in middleware. It replaces simple, self-contained token patterns in UTF-8 HTTP bodies and client WebSocket text messages. The initial pattern recognizes `sk-` tokens. It does not infer values from fields such as a JSON `password` property, and it does not guarantee that it will detect or remove every sensitive value. + +Its `config` accepts `mode: redact`, which is also the default. Policy validation rejects unknown fields and non-string values. Custom expressions are not configurable. + +Operator-run services advertise the operation and phase pairs they support. V1 defines `HTTP_REQUEST/PRE_CREDENTIALS`, `HTTP_RESPONSE/PRE_RETURN`, and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. A service may advertise any combination. + +## Register an operator-run service + +Start the service before the gateway, then add its registration to the local gateway TOML: + +```toml +[[openshell.supervisor.middleware]] +name = "local-content-guard" +grpc_endpoint = "https://content-guard.example:50051" +tls_ca_cert_path = "/etc/openshell/content-guard-ca.pem" +audience = "urn:example:content-guard" +max_payload_bytes = 262144 +timeout = "500ms" +``` + +| Field | Description | +| --- | --- | +| `name` | Operator-owned name used by policy attachments and diagnostics. Names must be unique. The `openshell/` namespace is reserved for built-ins. | +| `grpc_endpoint` | Service address reachable from the gateway and sandbox supervisors. Authenticated extensions use TLS `https://`. | +| `tls_ca_cert_path` | Optional PEM trust roots for a private HTTPS service. Custom roots replace platform roots and retain hostname verification. | +| `audience` | Exact audience expected by the service. Defaults to `urn:openshell:extension:middleware:`. | +| `allow_insecure_transport` | Allows plaintext `http://` without a bearer credential. Defaults to `false`. Use it only on a network that already authenticates callers. | +| `max_payload_bytes` | Operator ceiling for inspectable logical payloads across all advertised bindings, up to the 4 MiB platform maximum. | +| `timeout` | Optional service-wide RPC timeout from `10ms` through `30s`. Defaults to `500ms`. | + +Each binding returned by `Describe` may advertise a shorter timeout. OpenShell uses the smaller of the binding and service values. The service timeout applies to `Describe` and `ValidateConfig`. The effective binding timeout applies to request evaluation, response preflight and unit results, WebSocket preflight, and each WebSocket message. Response and WebSocket streams have no connection-wide deadline. + +The gateway calls `Describe` and validates every registration before it accepts traffic. Startup fails if a service is unavailable, returns an invalid capability, or exposes duplicate bindings for one operation and phase. The manifest `name` is diagnostic metadata and does not need to match the operator registration name. + +Registration is static. Restart the gateway after you add, remove, or change a service. See [Gateway Configuration](/reference/gateway-config#supervisor-middleware-services) for the full TOML context. + +## Authenticate OpenShell callers + +When gateway JWT signing is configured, OpenShell attaches a short-lived EdDSA bearer token to every remote middleware RPC. Gateway calls use `caller_kind: gateway`. Sandbox supervisor calls use `caller_kind: supervisor` and include the sandbox ID. Supervisors request credentials by registration name through `RefreshSandboxToken`; the gateway authorizes that name against the sandbox's effective policy. + +Return the expected audience in the `expected_audience` field of the `Describe` manifest. After authentication succeeds, OpenShell compares this field with the operator-configured audience and refuses to start on a mismatch. A strict service may reject the token before returning its manifest, in which case startup reports an authentication failure. Leave the field empty to skip the consistency check. + +Provision the trusted gateway URL, expected gateway ID, and public key or JWKS with the service. The expected issuer is exactly `openshell-gateway:`. After initial trust is established, `GET /.well-known/openid-configuration` and its `jwks_uri` provide key refresh. This document is OIDC-shaped, but its `issuer` is the gateway identity rather than the URL that serves the document. Compare `iss` with the configured value and fetch updates only over authenticated TLS from the trusted gateway URL. + +Cache keys by `kid` and validate: + +- `typ` is exactly `openshell-ext+jwt`. +- `alg` is pinned to `EdDSA` rather than selected from the token. +- The signature, expected issuer, exact audience, and positive expiry. +- `caller_kind`, plus the sandbox identity when the service scopes behavior per sandbox. + +Extension tokens and sandbox-to-gateway bootstrap tokens use the same signing key but have different audiences and `typ` values. Each extension token also carries a `jti`. OpenShell reuses a token until rotation and does not track `jti`, so do not reject a repeated value as a replay. Per-request replay resistance requires proof of possession or request binding. + +Set `allow_insecure_transport = true` only when the surrounding network authenticates callers. OpenShell then sends no credential to the service, and the gateway refuses to mint one if asked. The gateway logs a warning for the registration at every startup. + +## Attach middleware in policy + +Add configurations to the top-level `network_middlewares` map. Each map key is a stable policy-local identity: + +```yaml +network_middlewares: + regex-redactor: + name: Redact API tokens + middleware: openshell/regex + order: 10 + config: + mode: redact + on_error: fail_closed + endpoints: + include: ["*.example.com"] + exclude: ["trusted.example.com"] +``` + +The optional `name` defaults to the map key. `middleware` names a built-in or operator-run registration. A policy accepts at most 10 configurations. + +`include` selects destination hosts. `exclude` takes precedence. Each configuration accepts at most 32 combined patterns. Matching is case-insensitive and uses the same exact-host and DNS glob behavior as network policy endpoints. `*` matches one DNS label, `**` matches one or more labels, and an intra-label pattern such as `*-api.example.com` works. Policy validation rejects brace alternates such as `{prod,staging}`. + +Matching configurations run once each by ascending `order`. Order values must be unique across the complete policy, even when selectors do not overlap. The default is `0`, so set explicit values when a policy has multiple configurations. + +See [Policy Schema](/reference/policy-schema#network-middleware) for the complete field reference. + +## Choose failure behavior + +`on_error` applies after OpenShell selects a supported binding and that stage fails. Failures include an unavailable service, rejected configuration, invalid result, timeout, or payload over the stage limit. + +| Value | Request behavior | Response behavior | WebSocket behavior | +| --- | --- | --- | --- | +| `fail_closed` | Denies the request. This is the default. | Returns `502 response_delivery_failed` before commitment or aborts delivery after commitment. | Rejects the upgrade or closes the connection. | +| `fail_open` | Skips the failed stage. | Disables the stage and continues with the retained input. | Disables a broken stage for the rest of the connection and continues the remaining chain. | + +A valid upstream response can still exceed middleware envelope limits or contain header bytes that the middleware protocol cannot represent. OpenShell relays the original response when every selected response stage uses `fail_open`. Any selected `fail_closed` stage causes the canonical `502` delivery failure. Malformed or unsafe HTTP does not qualify for this bypass. + +An unsupported binding or message class is a coverage gap, not a middleware failure. `on_error` does not make an HTTP-only service inspect WebSocket messages, and it does not make V1 inspect binary messages. + +For protocols without a supported middleware operation, OpenShell uses its uninspectable-traffic behavior. A matching `fail_closed` configuration blocks that traffic. If all matching configurations use `fail_open`, OpenShell bypasses middleware and emits a detection finding. + +Use `fail_open` only when bypassing the stage preserves the intended security policy. OpenShell emits a detection finding for a bypass and a separate state-change finding when it disables a WebSocket stage. + +An explicit deny result always stops the chain, regardless of `on_error`. Middleware decisions also remain enforced when the endpoint uses `enforcement: audit`. To observe traffic without blocking it, return an allow decision with findings. + +Continue with the [HTTP request guide](/extensibility/supervisor-middleware/http) or [WebSocket session guide](/extensibility/supervisor-middleware/websocket) for protocol-specific results and limits. + +## Run the content guard example + +The [content guard example](https://github.com/NVIDIA/OpenShell/tree/main/examples/supervisor-middleware-content-guard) implements request, response, and WebSocket bindings in one service. It matches configured literal terms in UTF-8 request bodies, complete response bodies, and client WebSocket text messages. It supports redaction and denial but is not a general PII detector. + +The example includes a policy, local fixture, and smoke launcher. diff --git a/docs/extensibility/supervisor-middleware/http.mdx b/docs/extensibility/supervisor-middleware/http.mdx new file mode 100644 index 0000000000..9d953e9e50 --- /dev/null +++ b/docs/extensibility/supervisor-middleware/http.mdx @@ -0,0 +1,168 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "HTTP Supervisor Middleware" +sidebar-title: "HTTP Requests" +description: "Understand how supervisor middleware inspects HTTP requests and transforms upstream responses." +keywords: "Supervisor Middleware, HTTP, Request Body, Header Mutation, Payload Limits" +position: 3 +--- + +HTTP middleware evaluates admitted requests through `HTTP_REQUEST/PRE_CREDENTIALS` and final upstream responses through `HTTP_RESPONSE/PRE_RETURN`. Request stages run before credential injection. Response stages run before OpenShell delivers the result to the sandbox. + +## Request flow + +```mermaid +sequenceDiagram + participant App as Sandbox process + participant Supervisor as Network supervisor + participant Policy as Network and L7 policy + participant Chain as Ordered middleware chain + participant Upstream as Upstream service + + App->>Supervisor: HTTP request + Supervisor->>Policy: Evaluate destination, method, and path + Policy-->>Supervisor: Admit + Supervisor->>Supervisor: Select host matches and buffer body + loop Each stage by ascending order + Supervisor->>Chain: EvaluateHttpRequest + Chain-->>Supervisor: Allow, deny, or replace + opt Body replaced + Supervisor->>Policy: Re-check body-aware protocol policy + Policy-->>Supervisor: Admit or deny transformed body + end + end + alt A stage denies, fails closed, or produces a denied body + Supervisor-->>App: Deny before credential injection + else All selected stages complete + Supervisor->>Supervisor: Inject provider credentials + Supervisor->>Upstream: Forward request + end +``` + +The supervisor follows this sequence: + +1. It evaluates network and L7 policy. +2. It selects middleware whose host patterns match the admitted destination. +3. It buffers the body using the largest payload limit in the selected chain. +4. It runs stages by ascending `order`. +5. It re-checks GraphQL, JSON-RPC, or MCP policy after each body replacement. +6. It applies allowed changes, injects provider credentials, and forwards the request. + +Each stage receives a request that the current policy admits. A replacement cannot carry a denied or unparseable operation to a later stage or the upstream. If policy rejects a transformed body, the chain stops. With `enforcement: audit`, OpenShell logs the policy rejection and continues the request. + +If the post-transformation policy evaluation itself fails, OpenShell denies the request and emits a high-severity detection finding. This is not a middleware failure because the stage completed successfully. + +## Response flow + +```mermaid +sequenceDiagram + participant Upstream as Upstream service + participant Supervisor as Network supervisor + participant Chain as Response middleware chain + participant App as Sandbox process + + Upstream-->>Supervisor: Final HTTP response + loop Each stage by ascending order + Supervisor->>Chain: Preflight with status and safe headers + Chain-->>Supervisor: Skip or inspect with body mode + opt Inspect body + loop Normalized body units + Supervisor->>Chain: Body unit + Chain-->>Supervisor: Pass, transform, stop, or skip remaining + end + opt Response has trailers + Supervisor->>Chain: Normalized trailers + Chain-->>Supervisor: Trailer mutations + end + end + end + alt A stage stops delivery or fails closed + Supervisor--xApp: Return 502 before commitment or abort after commitment + else Response chain completes + Supervisor->>Supervisor: Repair framing and stale metadata + Supervisor-->>App: Deliver response + end +``` + +The supervisor relays interim `1xx` responses unchanged and retains the final non-`1xx` response. A `101` protocol upgrade bypasses generic response processing. + +For the final response, OpenShell runs `HttpResponsePreReturn.Evaluate` preflight once per matching stage in policy order. Each stage sees earlier response-header mutations and returns `SKIP` or `INSPECT`. An inspecting stage selects one body mode: + +| Mode | Behavior | +| --- | --- | +| `HEADERS_ONLY` | Inspects the status and safe response headers without receiving body or trailer events. | +| `WHOLE_BODY_BYTES` | Buffers the normalized body and sends one bounded body unit before committing the response head. | +| `STREAM_BYTES` | Sends ordered normalized units of at most 64 KiB. Each result accounts for its complete input unit before OpenShell reads more data. | + +Bodyless responses, partial responses, non-identity content encodings, and responses with `Cache-Control: no-transform` allow headers-only inspection but reject body inspection according to the stage's `on_error` policy. + +Body-inspecting stages receive a final body marker followed by normalized trailers. A stage may add only trailer names it declared during preflight. After transformation, OpenShell removes stale range and integrity metadata and repairs downstream framing. Buffered output uses a recalculated `Content-Length` unless trailers require chunked framing. Streaming output uses middleware-owned chunked framing. + +OpenShell keeps one request ID across the request and response hooks for an exchange. Response streams have no whole-response deadline. Preflight and each unit result use the effective binding timeout, and each unit's complete middleware chain is capped at 30 seconds. + +## Middleware input + +Middleware receives the request body, target, filtered headers, and request context. The context always includes `sandbox_id`. It also includes best-effort `sandbox_name` and `workspace` values. + +Use `sandbox_id` for authorization, persistence, durable correlation, and identity. Names are workspace-scoped and may be reused. Use `sandbox_name` and `workspace` only for display or logging, and fall back to the ID when either value is empty. + +OpenShell delivers visible request headers in wire order and preserves repeated names as separate entries. It removes credential, routing, framing, and hop-by-hop headers before evaluation. It also removes headers named by `Connection`, except for a validated WebSocket upgrade pair. Malformed headers and unsupported transfer-coding sequences fail before middleware or policy dispatch. + +Response middleware receives the final status and safe end-to-end headers. OpenShell removes framing, hop-by-hop, and `Connection`-nominated fields. It does not expose upstream transfer coding, transfer chunks, or socket-read boundaries. + +## Request results and body replacement + +A stage may allow or deny the request, return a replacement body, mutate approved headers, and report findings or metadata. The same per-stage limit applies to the input and replacement body. + +An explicit denial returns a structured response such as: + +```json +{ + "error": "middleware_denied", + "detail": "Request rejected by configured middleware", + "policy": "api-policy", + "middleware": "prototype-content-guard", + "reason_code": "content_match" +} +``` + +OpenShell does not copy a service-provided free-form `reason` into the response or security logs. A result may provide a stable `reason_code` of 1 through 64 bytes. It must start with a lowercase ASCII letter and contain only lowercase ASCII letters, digits, and underscores. An invalid code makes the result a middleware failure governed by `on_error`. + +A failed `fail_closed` stage returns `error: middleware_failed` with a platform-owned `detail`. Middleware errors omit policy-advisor remediation because changing network policy cannot repair the stage. + +## Response delivery failures + +OpenShell retains each response input until the corresponding stage returns a valid result. A failed `fail_open` stage can therefore be disabled without losing bytes, and the retained input continues through later stages. + +Before response commitment, a failed `fail_closed` stage returns `502 Bad Gateway` with `error: response_delivery_failed`. A `HEAD` response carries the same headers and no body. After commitment, OpenShell aborts delivery without appending an error body, terminating chunk, or error trailer. It also does not reuse the upstream connection. Retrying the request may repeat upstream side effects. + +## Header mutations + +A request result or response preflight may return ordered header mutations. OpenShell applies successful mutations before the next stage, so later middleware sees the accumulated header state. + +A `write` mutation selects one behavior when a case-insensitive header name already exists: + +- `append` adds another value. +- `overwrite` removes all existing values before writing the new one. +- `skip` leaves the existing values unchanged. + +A `remove` mutation removes all values for a case-insensitive name. Writes and removals may target middleware-visible end-to-end fields without a required prefix. OpenShell rejects request credential and routing fields. For responses it also protects status, framing, connection control, authentication challenges, content coding, range metadata, and security policy fields. Hop-by-hop and `Connection`-nominated fields remain protected in every profile. + +Header values cannot contain control characters. Request middleware also cannot write OpenShell credential placeholder syntax. Response trailer mutations use the same validation rules. Middleware may change an existing safe trailer, but it may add a new trailer only when the response preflight declared that name. + +OpenShell validates and applies each stage's mutations atomically. If one operation is invalid, it discards every mutation from that stage and applies `on_error`. + +## Payload and capacity limits + +Each binding declares `max_payload_bytes`. An operator-run registration sets a ceiling no higher than the advertised capability or the 4 MiB platform maximum. A selected request chain buffers with its largest stage limit so that every stage that can process the body receives it. + +Exceeding one stage's limit fails that stage and follows its `on_error`; other stages keep their own limits. OpenShell can skip an oversized `Content-Length` request under `fail_open` before it consumes body bytes. If a chunked body crosses the limit after consumption begins, OpenShell denies the request because it cannot safely resume the original stream. + +HTTP bodies share process-wide middleware evaluation capacity with active WebSocket message evaluations. At most 32 evaluations run while 64 additional callers wait without buffering payload bytes. When both bounds are full, OpenShell returns `503 Service Unavailable` before reading the request body. + +For `HTTP_RESPONSE`, `max_payload_bytes` limits a complete body under `WHOLE_BODY_BYTES`, each replacement, and the largest unit a streaming stage accepts. The platform caps stream input units at 64 KiB. Whole-body inspection delays commitment. Headers-only and streaming inspection commit streaming-compatible framing after preflight. + +Response body work shares the same 32 active and 64 waiting evaluation budget. Response streams also share the separate process-wide budget of 32 persistent middleware sessions with WebSocket stages. Response stream admission does not wait when that budget is full. OpenShell applies each selected configuration's `on_error` before opening its stream. + +See [Configure middleware](/extensibility/supervisor-middleware/configure) for registration limits and `on_error`, and [Operate middleware](/extensibility/supervisor-middleware/operate) for logging and service lifecycle. diff --git a/docs/extensibility/supervisor-middleware/index.mdx b/docs/extensibility/supervisor-middleware/index.mdx new file mode 100644 index 0000000000..1d11c3d516 --- /dev/null +++ b/docs/extensibility/supervisor-middleware/index.mdx @@ -0,0 +1,65 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "Supervisor Middleware" +sidebar-title: "Supervisor Middleware" +slug: "extensibility/supervisor-middleware" +description: "Inspect, block, or change sandbox requests and responses with supervisor middleware." +keywords: "Generative AI, Cybersecurity, AI Agents, Supervisor Middleware, Extensibility, Request Filtering" +position: 1 +--- + +Supervisor middleware lets you inspect, block, or change the content that an agent sends and receives over the network. The sandbox's supervisor, which enforces its network policy, runs middleware as traffic passes between the agent and an external service. + +## Why use middleware + +A network policy can allow an agent to call an API, but you may also need to check what the agent sends to that API or what the API returns. Middleware lets you add those checks without changing the agent's code. + +For example, you can use middleware to: + +- Redact recognized API token patterns from outgoing request bodies. +- Reject requests that contain content your organization prohibits. +- Remove sensitive content from HTTP responses before the agent receives it. +- Check outgoing WebSocket text messages during a session. + +The checks depend on the middleware you choose. OpenShell includes a basic token-pattern redactor, and you can run your own service for custom checks and transformations. See [Configure middleware](/extensibility/supervisor-middleware/configure) for the built-in redactor's limits and an example content guard service. + +## How it fits into a request + +OpenShell checks network and application-layer policy first. If policy blocks a request, it never reaches middleware. + +For an allowed request, OpenShell selects middleware by destination host. It runs the selected checks in the order you configure, then adds provider credentials and sends the request to the external service. Request middleware cannot see the provider credentials that OpenShell adds afterward. + +When the service returns an HTTP response, response middleware can inspect or change it before OpenShell delivers it to the agent. An outgoing request check and an incoming response check are separate capabilities, so choose middleware that supports the direction you need. + +Middleware can also inspect outgoing WebSocket text messages. It does not inspect binary messages or messages returning from the service over a WebSocket connection. + +If middleware explicitly denies a request or message, OpenShell blocks it even when network policy allows it. You also choose what happens if a check fails or its service is unavailable. By default, OpenShell stops the affected traffic. You can configure it to continue without the failed check when that is acceptable. + +## Choose a guide + + + + + +Start here to choose middleware, connect your own service, and select which destinations it checks. + + + + +Check or change outgoing requests and incoming responses. Learn what middleware can read and modify, and how failures affect delivery. + + + + +Check outgoing text messages during a connection. Learn which messages middleware sees and when OpenShell closes a session. + + + + +Run middleware services, apply configuration changes, and use logs to understand decisions and failures. + + + + +For field definitions, see [Policy Schema](/reference/policy-schema#network-middleware) and [Gateway Configuration](/reference/gateway-config#supervisor-middleware-services). diff --git a/docs/extensibility/supervisor-middleware/operate.mdx b/docs/extensibility/supervisor-middleware/operate.mdx new file mode 100644 index 0000000000..d7f9335b98 --- /dev/null +++ b/docs/extensibility/supervisor-middleware/operate.mdx @@ -0,0 +1,67 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "Operate Supervisor Middleware" +sidebar-title: "Operate Middleware" +description: "Run and observe operator middleware services and understand shared platform limits." +keywords: "Supervisor Middleware, Operations, OCSF, Logging, Service Lifecycle" +position: 5 +--- + +Operator-run middleware sits on the request path. Plan service availability, gateway restarts, sandbox configuration reloads, and alerting around the failure behavior in each policy attachment. + +## Start and update services + +- Start registered services before the gateway. The gateway validates every registration during startup. +- Keep service endpoints reachable from the gateway and sandbox supervisors. Supervisors call services directly on the request path. +- Restart the gateway after changing registrations. +- Keep required services available before creating or updating policies. The gateway validates implementation-owned configuration before persisting a policy. +- Treat `fail_open` as an explicit choice to favor availability over enforcement. + +When a sandbox's effective configuration changes, its running supervisor validates the new service registry before installing it. If validation fails, the supervisor keeps the last-known-good registry and emits a configuration failure event. + +## Observe middleware + +OpenShell emits middleware activity through OCSF logging: + +- Each invocation records the policy-local configuration name, attached middleware name, decision, transformation state, and failure state. +- A denied invocation records a platform-owned reason built from the policy-local name and optional validated reason code. It does not record free-form reason text from the service. +- A bypass under `fail_open` emits a detection finding. +- A required stage that fails closed emits a high-severity detection finding. +- An HTTP-only host match on a WebSocket session emits an informational `binding_not_selected` coverage event. +- A binary message encountered by an active WebSocket stage emits an informational `unsupported_message_type` event with message type, sequence, and byte count. It is not an invocation or failure. +- Registry reload success and failure emit configuration state changes. + +Built-in findings include their type, label, and aggregate count. Operator-run findings use the registration name, a platform label, and the aggregate count. OpenShell does not log service-provided finding text or diagnostic metadata. + +A stage can return at most 32 findings. A 10-stage chain can retain and emit up to 320 findings. Exceeding the per-stage cap makes the response invalid and applies `on_error`. + +See [Logging](/observability/logging) for log access and [OCSF JSON Export](/observability/ocsf-json-export) for structured export. + +## Size gRPC messages + +The 4 MiB platform payload maximum does not include the rest of the protobuf envelope. OpenShell also bounds: + +| Component | Limit | +| --- | --- | +| Service configuration | 64 KiB. | +| Request context | 4 KiB. | +| Target | 32 KiB. | +| Request headers | 128 lines and 64 KiB encoded. | +| Discarded free-form reason | 4 KiB. | +| Validated reason code | 64 bytes. | +| Header mutations | 64 operations and 64 KiB encoded. | +| Findings | 32 entries of at most 4 KiB encoded each. | +| Metadata | 64 entries and 32 KiB total. | + +Configure middleware gRPC servers to accept at least 4 MiB plus 293 KiB for requests and responses so they can process every platform-valid envelope. + +## Security and deployment limits + +- A `fail_closed` selector cannot cover a `tls: skip` endpoint because OpenShell cannot inspect that traffic. An all-`fail_open` match may cover it; OpenShell bypasses middleware and emits a detection finding. +- Operator-run services use TLS `https://` when gateway JWT signing is enabled unless their registration sets `allow_insecure_transport`. Certificates must chain to the configured custom CA or platform roots, and the hostname must match. +- Extension and sandbox admission tokens use the same signing key. Audience and `typ` separate them, but the extension credential path cannot rotate or revoke independently. +- OpenShell does not track or revoke `jti`. Bearer tokens can be replayed until expiry unless the service adds proof of possession or request binding. +- mTLS client authentication, health checks, runtime registration, and overlapping signing-key rotation are not available. + +Protocol-specific runtime limits are documented in [HTTP requests](/extensibility/supervisor-middleware/http#payload-and-capacity-limits) and [WebSocket sessions](/extensibility/supervisor-middleware/websocket#payload-assembly-and-capacity-limits). diff --git a/docs/extensibility/supervisor-middleware/websocket.mdx b/docs/extensibility/supervisor-middleware/websocket.mdx new file mode 100644 index 0000000000..fdbde4ea19 --- /dev/null +++ b/docs/extensibility/supervisor-middleware/websocket.mdx @@ -0,0 +1,109 @@ +--- +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +title: "WebSocket Supervisor Middleware" +sidebar-title: "WebSocket Sessions" +description: "Understand WebSocket middleware preflight, session events, message inspection, limits, and close behavior." +keywords: "Supervisor Middleware, WebSocket, RFC 6455, Session Stream, Text Messages" +position: 4 +--- + +WebSocket middleware uses the `WEBSOCKET_MESSAGE/PRE_CREDENTIALS` binding to evaluate an upgrade preflight and complete client-to-upstream text messages. OpenShell opens one ordered `EvaluateWebSocketSession` stream for each selected stage. + +## Session flow + +```mermaid +sequenceDiagram + participant App as Sandbox process + participant Supervisor as Network supervisor + participant Stages as Selected middleware stages + participant Upstream as Upstream service + + App->>Supervisor: WebSocket upgrade request + Supervisor->>Supervisor: Admit HTTP upgrade and select bindings + Supervisor->>Stages: Open streams and send preflights concurrently + Stages-->>Supervisor: Inspect, skip, or deny + Note over Supervisor,Stages: SKIP ends that stage stream + alt Any stage denies + Supervisor-->>App: Reject upgrade + Supervisor->>Stages: Session end for each writable stream + else No stage denies + Supervisor->>Upstream: Forward upgrade + Upstream-->>Supervisor: 101 Switching Protocols + Supervisor->>Stages: Session start for inspecting stages + loop Each client text message + App->>Supervisor: Frames + Supervisor->>Supervisor: Reassemble and decompress + Supervisor->>Stages: Complete text message + Stages-->>Supervisor: Allow, deny, or replace + alt Message allowed + Supervisor->>Upstream: Re-frame and forward text + else Message denied or a stage fails closed + Supervisor-->>App: Close with code 1008 + end + end + Supervisor->>Stages: Best-effort session end + end +``` + +The supervisor first finds host-matched middleware configurations, then keeps only implementations that advertise the WebSocket binding. An HTTP-only attachment may still inspect the upgrade request through its HTTP binding, but it does not join the post-upgrade message chain. OpenShell emits `binding_not_selected` coverage for that attachment. + +## Session events + +OpenShell sends these events over each selected stage stream: + +1. A preflight before contacting the upstream. The stage returns `INSPECT`, voluntary `SKIP`, or authoritative `DENY`, plus optional bounded findings and metadata. Selected preflights run concurrently. Any `DENY` rejects the upgrade regardless of `on_error`. +2. A session-start notification after the upstream accepts the upgrade. It includes the negotiated subprotocol. +3. Complete client-to-upstream text messages in sequence order. OpenShell reassembles fragmented messages and decompresses negotiated `permessage-deflate` messages before evaluation. +4. A best-effort session-end notification while the stream remains writable. + +Preflight and message events require `WebSocketSessionEventResult` responses. Session start and end are notifications. OpenShell attempts one terminal event for every opened stream, including a stream opened for a preflight that later rejects the upgrade. It half-closes the request stream and briefly drains the response stream. Middleware services should finish their response stream after request EOF. + +The protobuf also reserves `PRE_RETURN` for future upstream-to-client inspection. A service that eventually advertises both phases receives two independent streams for one WebSocket session. + +## Message handling + +The protobuf represents each logical message with a `text` or `binary` variant. Text uses the protobuf `string` type, so invalid UTF-8 cannot enter the middleware contract. + +A result may omit its replacement to preserve the input or return a text replacement, including an empty string. OpenShell rejects a replacement that changes the message type. It re-frames an allowed replacement, re-compresses it when required, and forwards it. + +V1 does not deliver binary messages to middleware. Binary messages pass through under both `on_error` modes. OpenShell emits `unsupported_message_type` coverage for each active stage and advances the session-wide sequence number, so the next text event can contain a valid sequence gap. Control frames and upstream-to-client messages also remain uninspected. + + + +If your deployment requires inspection of every WebSocket message class or both directions, V1 cannot enforce that requirement. + + + +## Failure behavior + +A preflight `DENY` is a successful policy decision, not a middleware failure. OpenShell rejects the upgrade before contacting the upstream and sends `MIDDLEWARE_DENIAL` session-end notifications to streams opened by successful preflights. + +If a selected stage fails under `fail_open`, OpenShell disables it for the rest of the connection and continues the remaining chain. It emits both a bypass finding and a state-change finding. Under `fail_closed`, OpenShell rejects the upgrade or closes the connection. + +A message result may include the same validated `reason_code` used by HTTP results. It must be 1 through 64 bytes, start with a lowercase ASCII letter, and contain only lowercase ASCII letters, digits, and underscores. An invalid code is a middleware failure governed by `on_error`. + +## Payload, assembly, and capacity limits + +For a WebSocket binding, `max_payload_bytes` covers one complete client text message and its replacement. It does not cover the whole session or binary relay traffic. Exceeding a selected stage's effective limit follows that stage's `on_error`. + +The parsed-text platform maximum is 4 MiB. A text message may contain at most 4,096 fragments, must make input progress within 30 seconds, and must finish assembly within 2 minutes. Forwarding the completed message must finish within another 2 minutes. + +The supervisor allows at most 32 concurrent text assemblies and 64 additional callers waiting without buffered payload bytes. If both bounds are full, it closes the connection with code `1013` before reading the new payload. This process-wide assembly budget applies even when no middleware is selected and persists across policy reloads. + +Active message evaluations share the middleware budget with HTTP bodies. At most 32 evaluations run and 64 additional callers wait. Persistent middleware streams have a separate process-wide limit of 32 sessions. Session admission does not wait when that limit is full. OpenShell applies each selected configuration's `on_error` before it opens a stream. + +## Close codes + +| Code | Meaning in the parsed relay | +| --- | --- | +| `1002` | WebSocket protocol error. | +| `1007` | Invalid UTF-8 text. | +| `1008` | Middleware or policy denial. | +| `1009` | Parsed text exceeds the platform limit. | +| `1012` | Policy reload makes the pinned generation stale. | +| `1013` | Assembly capacity is full. | + +Raw binary frames retain the 16 MiB relay-safety bound. The middleware payload limit does not change that bound because middleware never receives binary messages. + +See [Configure middleware](/extensibility/supervisor-middleware/configure) for attachment and failure settings, and [Operate middleware](/extensibility/supervisor-middleware/operate) for coverage events and reload behavior. diff --git a/docs/reference/gateway-config.mdx b/docs/reference/gateway-config.mdx index 14cdb05ece..fd3d54b7da 100644 --- a/docs/reference/gateway-config.mdx +++ b/docs/reference/gateway-config.mdx @@ -299,13 +299,17 @@ 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 only to `EvaluateHttpRequest`, WebSocket preflight, and each WebSocket message. An accepted WebSocket stream has 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 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. + +Response middleware also has an 8 MiB aggregate retained-body limit per response session across stage buffers and pending output. This limit is separate from each service's `max_payload_bytes` and cannot be configured. A replacement that exceeds the available budget follows the stage's `on_error` policy. Fail-open disables that stage and forwards its original input; fail-closed returns 502 before commitment or stops delivery after commitment. Streaming inputs flush after a short coalescing window without waiting for an entire upstream transfer chunk. + +Whole-body response inspection has a fixed 120-second deadline shared by all whole-body stages in one response, separate from middleware RPC timeouts. 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. 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 f19951981b..357ebacf31 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.lock +++ b/examples/supervisor-middleware-content-guard/Cargo.lock @@ -231,6 +231,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" @@ -286,6 +297,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" @@ -302,6 +323,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" @@ -445,6 +487,7 @@ dependencies = [ "cfg-if", "libc", "r-efi", + "rand_core", ] [[package]] @@ -499,6 +542,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" @@ -710,6 +772,8 @@ checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", "hashbrown 0.17.1", + "serde", + "serde_core", ] [[package]] @@ -745,6 +809,55 @@ 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 = "jobserver" version = "0.1.35" @@ -761,6 +874,12 @@ version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "libm" +version = "0.2.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6d2cec3eae94f9f509c767b45932f1ada8350c4bdb85af2fcab4a3c14807981" + [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -874,6 +993,22 @@ dependencies = [ "libc", ] +[[package]] +name = "noyalib" +version = "0.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f075ef19fa3bcf8697c0ef96c37d5c435d339a40ab8081cae3aac3a4e7fee9a" +dependencies = [ + "hashbrown 0.17.1", + "indexmap", + "libm", + "memchr", + "rustc-hash", + "serde", + "serde_core", + "smallvec", +] + [[package]] name = "object" version = "0.37.3" @@ -888,6 +1023,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" @@ -936,12 +1075,26 @@ dependencies = [ "tower", ] +[[package]] +name = "openshell-policy" +version = "0.0.0" +dependencies = [ + "hickory-proto", + "miette", + "noyalib", + "openshell-core", + "prost-types", + "serde", + "serde_json", +] + [[package]] name = "openshell-supervisor-middleware-content-guard" version = "0.0.0" dependencies = [ "clap", "openshell-core", + "openshell-policy", "prost-types", "tokio", "tokio-stream", @@ -1032,6 +1185,12 @@ version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" +[[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" @@ -1212,6 +1371,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" @@ -1270,6 +1446,21 @@ version = "0.1.27" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b50b8869d9fc858ce7266cce0194bd74df58b9d0e3f6df3a9fc8eb470d95c09d" +[[package]] +name = "rustc-hash" +version = "2.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" + +[[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" @@ -1340,6 +1531,15 @@ dependencies = [ "untrusted", ] +[[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" @@ -1378,6 +1578,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" @@ -1437,6 +1643,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" @@ -1589,6 +1811,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" @@ -1849,6 +2086,16 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" +[[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" @@ -1864,6 +2111,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 135316979c..89e7b7a85e 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..fcdc37df56 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 configured-literal guard applies the same case-sensitive terms to UTF-8 HTTP request bodies, complete HTTP response bodies, and client WebSocket text messages. It is not a general PII detector. > [!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. +> 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 and response 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`, `mise`, and `uv` with Python 3 on the host before running the smoke script. Start Docker or Podman. The supervisor image build uses the repository's Linux cross-compilation toolchain, including `cargo-zigbuild` and Zig on macOS. Install the repository's mise tools before running it. ## 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,16 @@ 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 script defaults to Docker. Set `CONTENT_GUARD_SMOKE_DRIVER=podman` to build and run with Podman instead. + +On Linux and macOS, the script runs `mise run docker:build:supervisor` with the selected container engine to build a Linux supervisor from the current checkout. It configures that driver's `supervisor_image` with a unique local tag, so the response checks exercise the local runtime changes. macOS host binaries are never used inside the sandbox. The local image remains available after the smoke run. + +Cargo's configured target directory applies to the host binaries and the Linux supervisor build. For example: + +```shell +CARGO_TARGET_DIR=/tmp/content-guard-target ./examples/supervisor-middleware-content-guard/smoke.sh --test-suite +``` + ## 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 +64,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 +95,34 @@ curl -sS https://httpbin.org/anything \ The echoed JSON body contains `[FILTERED]` instead of the configured term. +## HTTP response behavior + +The smoke launcher starts the local fixture. To start it manually: + +```shell +uv run --no-project python examples/supervisor-middleware-content-guard/upstream.py +``` + +The policy permits `GET /clean` and `GET /sensitive` on +`http://host.openshell.internal:18081`. The first returns ordinary public text. +The second contains both configured terms. Redact mode returns +`contains [FILTERED] and [FILTERED]`. Deny mode returns typed `BlockDelivery` +with reason code `content_match`, which produces the canonical 403 response +before delivery. The smoke suite recreates the sandbox in deny mode and checks +both clean and matching responses through the external gRPC service. + +Every selected response requires `WHOLE_BODY_BYTES`. If that mode is unavailable, +the service returns a middleware failure and the policy's `on_error` decides +whether delivery fails open or closed. This includes encoded, partial, +no-transform, bodyless, and known oversized responses. Unknown-length bodies can +also exceed the runtime limit during collection. Invalid UTF-8 fails the same way. +The example policy uses `fail_closed`. + +Clean bodies pass unchanged. Matching spans are merged and replaced in the +complete body, so transport chunk boundaries do not affect matching. Trailers +are accepted without mutation. The guard does not decode compressed bodies, +normalize Unicode, scan response headers, retain stream units, or spool bodies. + ## 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 +146,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..da08607f27 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,18 @@ network_policies: path: /anything binaries: - path: /usr/bin/curl + guard-responses: + name: Guard responses + endpoints: + - host: host.openshell.internal + port: 18081 + protocol: rest + rules: + - allow: + method: GET + path: /clean + - allow: + method: GET + path: /sensitive + 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..f16ed28ca7 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 + Compute driver: docker (default) or podman. EOF } @@ -103,19 +105,30 @@ detect_service_host() { } SERVICE_HOST="$(detect_service_host)" +COMPUTE_DRIVER="${CONTENT_GUARD_SMOKE_DRIVER:-docker}" +case "$COMPUTE_DRIVER" in + docker | podman) ;; + *) + echo "CONTENT_GUARD_SMOKE_DRIVER must be docker or podman" >&2 + exit 1 + ;; +esac 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 fi -TMPDIR="$(mktemp -d)" -LOG_DIR="$TMPDIR/logs" -JWT_DIR="$TMPDIR/jwt" -GATEWAY_CONFIG="$TMPDIR/gateway.toml" +SMOKE_TMP_DIR="$(mktemp -d)" +LOG_DIR="$SMOKE_TMP_DIR/logs" +JWT_DIR="$SMOKE_TMP_DIR/jwt" +GATEWAY_CONFIG="$SMOKE_TMP_DIR/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" +SUPERVISOR_IMAGE="localhost/openshell-content-guard/supervisor:$RUN_ID" # Sandbox names are capped at 19 characters. Use a short prefix with # the PID for uniqueness; keep the full RUN_ID for gateway identity. SANDBOX_NAME="cg-$$-$RANDOM" @@ -141,8 +154,13 @@ 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" + rm -rf "$SMOKE_TMP_DIR" else echo "logs retained in $LOG_DIR" >&2 fi @@ -214,8 +232,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.$COMPUTE_DRIVER] +supervisor_image = "$SUPERVISOR_IMAGE" EOF } @@ -239,11 +261,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 +278,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,17 +355,42 @@ wait_for_middleware() { fail "content guard service is reachable at $SERVICE_HOST:$MIDDLEWARE_PORT" } +start_upstream() { + printf 'INFO starting content guard upstream at %s:18081\n' "$SERVICE_HOST" + uv run --no-project python "$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 "content guard upstream starts" + fi + if curl -fsS --max-time 1 "http://127.0.0.1:18081/clean" >/dev/null 2>&1; then + printf 'INFO content guard upstream is ready\n' + return + fi + sleep 1 + done + fail "content guard 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 & + --db-url "sqlite://$SMOKE_TMP_DIR/gateway.db" >"$GATEWAY_LOG" 2>&1 & GATEWAY_PID=$! } @@ -354,10 +418,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 +432,33 @@ request() { --data '{"note":"prototype-secret"}' } +response_request() { + local path="$1" + "${CLI[@]}" sandbox exec --name "$SANDBOX_NAME" --no-tty -- \ + curl -sS -i --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 checking response pass-through and redaction\n' + if ! response_request clean >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'ordinary public text' "$response_output"; then + fail "clean response passes unchanged" + fi + if ! response_request sensitive >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'contains [FILTERED] and [FILTERED]' "$response_output" || + grep -Fq 'prototype-secret' "$response_output"; then + fail "configured response terms are redacted" + fi + printf 'PASS response pass-through and redaction\n' if grep -Fq '[FILTERED]' "$guarded_output" && ! grep -Fq 'prototype-secret' "$guarded_output"; then printf 'PASS guarded request is filtered\n' else @@ -394,6 +477,24 @@ run_suite() { fail "unguarded request is unchanged" fi + # Recreate with the same terms in deny mode, through the external service. + "${CLI[@]}" sandbox delete "$SANDBOX_NAME" >>"$SETUP_LOG" 2>&1 + SANDBOX_CREATED=0 + sed '/replacement:/d; s/mode: redact/mode: deny/' "$EXAMPLE_DIR/policy.yaml" >"$SMOKE_TMP_DIR/deny.yaml" + SANDBOX_CREATED=1 + run_setup_step "creating deny sandbox" "${CLI[@]}" sandbox create --name "$SANDBOX_NAME" --policy "$SMOKE_TMP_DIR/deny.yaml" --no-tty --detach -- sleep infinity + if ! response_request sensitive >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'HTTP/1.1 403 Forbidden' "$response_output" || + ! grep -Fq 'content_match' "$response_output" || + grep -Fq 'prototype-secret' "$response_output"; then + fail "configured response term blocks delivery" + fi + if ! response_request clean >"$response_output" 2>>"$SETUP_LOG" || + ! grep -Fq 'ordinary public text' "$response_output"; then + fail "deny mode passes clean responses" + fi + printf 'PASS response denial\n' + "${CLI[@]}" sandbox delete "$SANDBOX_NAME" >>"$SETUP_LOG" 2>&1 SANDBOX_CREATED=0 echo "ALL PASS content guard smoke" @@ -439,15 +540,27 @@ require_command cargo require_command curl require_command jq require_command openssl +require_command uv +require_command mise 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 +# Always rebuild from this checkout and load into the selected runtime. Native +# macOS binaries cannot run in Linux sandboxes; Podman also needs an image. +# A unique tag prevents the driver from selecting an older published runtime. +run_setup_step "building Linux sandbox supervisor image" \ + env -u CI -u DOCKER_PLATFORM -u DOCKER_PUSH -u DOCKER_OUTPUT \ + CONTAINER_ENGINE="$COMPUTE_DRIVER" PREBUILT_AUTO_STAGE=1 \ + IMAGE_REGISTRY=localhost/openshell-content-guard IMAGE_TAG="$RUN_ID" \ + mise run docker:build:supervisor 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..f395b2a070 100644 --- a/examples/supervisor-middleware-content-guard/src/main.rs +++ b/examples/supervisor-middleware-content-guard/src/main.rs @@ -6,20 +6,30 @@ 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, - web_socket_session_event, web_socket_session_event_result, + Decision, Finding, HttpRequestEvaluation, HttpRequestResult, HttpResponseBlockDelivery, + HttpResponseBodyMode, HttpResponseBodyResult, HttpResponseBodyTransform, HttpResponseEvent, + HttpResponseEventResult, HttpResponsePreflightInspect, HttpResponsePreflightResult, + HttpResponseTrailersResult, MiddlewareBinding, MiddlewareManifest, + SupervisorMiddlewareOperation, SupervisorMiddlewarePhase, ValidateConfigRequest, + ValidateConfigResponse, WebSocketMessage, WebSocketMessageResult, WebSocketPreflightAction, + WebSocketPreflightDecision, WebSocketSessionEvent, WebSocketSessionEventResult, + 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 +248,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 +298,146 @@ impl SupervisorMiddleware for ContentGuard { } } +#[derive(Debug, Default)] +struct ResponseSessionState { + config: Option, + body_ended: bool, + trailers_seen: bool, +} +impl ResponseSessionState { + fn preflight( + &mut self, + preflight: openshell_core::proto::HttpResponsePreflight, + ) -> Result { + if self.config.is_some() { + return Err(Status::failed_precondition("duplicate preflight")); + } + let config = + GuardConfig::parse(preflight.config.as_ref()).map_err(Status::invalid_argument)?; + if !preflight + .permitted_body_modes + .contains(&(HttpResponseBodyMode::WholeBodyBytes as i32)) + { + return Err(Status::failed_precondition( + "content guard requires WHOLE_BODY_BYTES", + )); + } + self.config = Some(config); + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult { + action: Some(http_response_preflight_result::Action::Inspect( + HttpResponsePreflightInspect { + body_mode: HttpResponseBodyMode::WholeBodyBytes as i32, + header_mutations: vec![], + }, + )), + ..Default::default() + }, + )), + }) + } + fn body( + &mut self, + body: openshell_core::proto::HttpResponseBodyUnit, + ) -> Result { + let config = self + .config + .as_ref() + .ok_or_else(|| Status::failed_precondition("body before preflight"))?; + if self.body_ended || body.sequence != 1 || !body.end_of_stream { + return Err(Status::failed_precondition( + "expected one complete response body", + )); + } + let Some(http_response_body_unit::Payload::Data(data)) = body.payload else { + return Err(Status::invalid_argument("body data required")); + }; + let text = std::str::from_utf8(&data) + .map_err(|_| Status::invalid_argument("content guard requires a UTF-8 body"))?; + let result = inspect(config, text); + let action = if result.denied { + http_response_body_result::Action::BlockDelivery(HttpResponseBlockDelivery {}) + } else if let Some(replacement) = result.replacement { + http_response_body_result::Action::Transform(HttpResponseBodyTransform { + replacement: Some(http_response_body_transform::Replacement::Data( + replacement.into_bytes(), + )), + }) + } else { + http_response_body_result::Action::PassThrough( + openshell_core::proto::HttpResponseBodyPassThrough {}, + ) + }; + self.body_ended = true; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: body.sequence, + action: Some(action), + reason: result.reason, + reason_code: result.reason_code, + findings: result.findings, + metadata: result.metadata, + }, + )), + }) + } + fn trailers(&mut self) -> Result { + if !self.body_ended || self.trailers_seen { + return Err(Status::failed_precondition("expected trailers after body")); + } + self.trailers_seen = true; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult::default(), + )), + }) + } +} + +#[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}'")); @@ -289,11 +445,37 @@ fn validate_phase(phase: i32) -> Result<(), String> { Ok(()) } +#[derive(Default)] +struct GuardOutcome { + denied: bool, + replacement: Option, + reason: String, + reason_code: String, + findings: Vec, + metadata: HashMap, +} fn evaluate(config: &GuardConfig, body: &str) -> HttpRequestResult { + let result = inspect(config, body); + HttpRequestResult { + decision: if result.denied { + Decision::Deny + } else { + Decision::Allow + } as i32, + has_body: result.replacement.is_some(), + body: result.replacement.unwrap_or_default().into_bytes(), + reason: result.reason, + reason_code: result.reason_code, + findings: result.findings, + metadata: result.metadata, + ..Default::default() + } +} +fn inspect(config: &GuardConfig, body: &str) -> GuardOutcome { let (ranges, match_count, matched_term_count) = find_match_ranges(body, &config.terms); if match_count == 0 { - return allow_result(); + return GuardOutcome::default(); } let finding = Finding { @@ -315,27 +497,22 @@ fn evaluate(config: &GuardConfig, body: &str) -> HttpRequestResult { ), ]); - match config.mode { - Mode::Redact => HttpRequestResult { - decision: Decision::Allow as i32, - reason: String::new(), - body: redact_ranges(body, &ranges, &config.replacement).into_bytes(), - has_body: true, - header_mutations: Vec::new(), - findings: vec![finding], - metadata, - reason_code: String::new(), + GuardOutcome { + denied: config.mode == Mode::Deny, + replacement: (config.mode == Mode::Redact) + .then(|| redact_ranges(body, &ranges, &config.replacement)), + reason: if config.mode == Mode::Deny { + "payload matched configured content".into() + } else { + String::new() }, - Mode::Deny => HttpRequestResult { - decision: Decision::Deny as i32, - reason: "payload matched configured content".into(), - body: Vec::new(), - has_body: false, - header_mutations: Vec::new(), - findings: vec![finding], - metadata, - reason_code: "content_match".into(), + reason_code: if config.mode == Mode::Deny { + "content_match".into() + } else { + String::new() }, + findings: vec![finding], + metadata, } } @@ -356,23 +533,21 @@ fn evaluate_websocket_message( "WebSocket text message exceeds {MAX_PAYLOAD_BYTES} bytes" ))); } - let result = evaluate(config, payload); - let replacement = if result.has_body { - Some(web_socket_message_result::Replacement::Text( - String::from_utf8(result.body) - .expect("content guard replacements are constructed from UTF-8 text"), - )) - } else { - None - }; + let result = inspect(config, payload); Ok(WebSocketMessageResult { sequence: message.sequence, - decision: result.decision, - replacement, + decision: if result.denied { + Decision::Deny + } else { + Decision::Allow + } as i32, + replacement: result + .replacement + .map(web_socket_message_result::Replacement::Text), reason: result.reason, + reason_code: result.reason_code, findings: result.findings, metadata: result.metadata, - reason_code: result.reason_code, }) } @@ -451,25 +626,13 @@ fn redact_ranges(body: &str, ranges: &[Range], replacement: &str) -> Stri transformed } -fn allow_result() -> HttpRequestResult { - HttpRequestResult { - decision: Decision::Allow as i32, - reason: String::new(), - body: Vec::new(), - has_body: false, - header_mutations: Vec::new(), - findings: Vec::new(), - metadata: HashMap::new(), - reason_code: String::new(), - } -} - #[tokio::main] async fn main() -> Result<(), Box> { let cli = Cli::parse(); 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 +641,10 @@ async fn main() -> Result<(), Box> { #[cfg(test)] mod tests { use super::*; - use openshell_core::proto::{MiddlewareSessionEnd, WebSocketPreflight, WebSocketSessionStart}; + use openshell_core::proto::{ + HttpResponseBodyUnit, HttpResponsePreflight, MiddlewareSessionEnd, WebSocketPreflight, + WebSocketSessionStart, + }; use prost_types::{ListValue, Value}; use std::collections::BTreeMap; @@ -511,13 +677,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 +694,110 @@ 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(mode: &str) -> HttpResponsePreflight { + HttpResponsePreflight { + config: Some(config(mode, &["prototype-secret", "秘密"], None)), + permitted_body_modes: vec![HttpResponseBodyMode::WholeBodyBytes as i32], + ..Default::default() + } + } + #[test] + fn response_guard_passes_redacts_and_denies() { + for (mode, input, expected) in [ + ("redact", "clean", None), + ( + "redact", + "a prototype-secret 秘密", + Some("a [REDACTED] [REDACTED]"), + ), + ("deny", "prototype-secret", None), + ] { + let mut state = ResponseSessionState::default(); + state.preflight(response_preflight(mode)).unwrap(); + let unit = HttpResponseBodyUnit { + sequence: 1, + payload: Some(http_response_body_unit::Payload::Data( + input.as_bytes().to_vec(), + )), + end_of_stream: true, + }; + let result = state.body(unit.clone()).unwrap(); + assert!(state.body(unit).is_err()); + let Some(http_response_event_result::Result::BodyResult(result)) = result.result else { + panic!("body result") + }; + if mode == "deny" { + assert!(matches!( + result.action, + Some(http_response_body_result::Action::BlockDelivery(_)) + )); + assert_eq!(result.reason_code, "content_match"); + } else if let Some(expected) = expected { + let Some(http_response_body_result::Action::Transform(transform)) = result.action + else { + panic!("transform") + }; + assert_eq!( + transform.replacement, + Some(http_response_body_transform::Replacement::Data( + expected.as_bytes().to_vec() + )) + ); + } else { + assert!(matches!( + result.action, + Some(http_response_body_result::Action::PassThrough(_)) + )); + } + let trailers = state.trailers().unwrap(); + let Some(http_response_event_result::Result::TrailersResult(trailers)) = + trailers.result + else { + panic!("trailers") + }; + assert!(trailers.trailer_mutations.is_empty()); + assert!(state.trailers().is_err()); + } + } + #[test] + fn response_guard_rejects_unavailable_inspection_and_invalid_input() { + let mut preflight = response_preflight("redact"); + preflight.permitted_body_modes = vec![HttpResponseBodyMode::HeadersOnly as i32]; + assert!( + ResponseSessionState::default() + .preflight(preflight) + .is_err() + ); + for (sequence, end_of_stream, payload) in [ + (2, true, Some(vec![])), + (1, false, Some(vec![])), + (1, true, Some(vec![0xff])), + (1, true, None), + ] { + let mut state = ResponseSessionState::default(); + assert!(state.trailers().is_err()); + state.preflight(response_preflight("redact")).unwrap(); + assert!(state.preflight(response_preflight("redact")).is_err()); + assert!( + state + .body(HttpResponseBodyUnit { + sequence, + end_of_stream, + payload: payload.map(http_response_body_unit::Payload::Data) + }) + .is_err() + ); + } } #[tokio::test] @@ -737,4 +1007,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..85a85396cf --- /dev/null +++ b/examples/supervisor-middleware-content-guard/upstream.py @@ -0,0 +1,26 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + + +class Handler(BaseHTTPRequestHandler): + def do_GET(self): + bodies = { + "/clean": b"ordinary public text", + "/sensitive": b"contains prototype-secret and internal-only", + } + body = bodies.get(self.path, b"not found") + self.send_response(200 if self.path in bodies else 404) + self.send_header("Content-Type", "text/plain; charset=utf-8") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + +with ThreadingHTTPServer(("0.0.0.0", 18081), Handler) as server: + print("content guard upstream listening on 0.0.0.0:18081", flush=True) + try: + server.serve_forever() + except KeyboardInterrupt: + pass diff --git a/skills/debug-openshell-cluster/SKILL.md b/skills/debug-openshell-cluster/SKILL.md index 7d47b18065..ae19330520 100644 --- a/skills/debug-openshell-cluster/SKILL.md +++ b/skills/debug-openshell-cluster/SKILL.md @@ -113,7 +113,9 @@ 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 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 fixed, non-resetting 120-second wall-clock deadline shared across response reads and whole-body barriers; inspect the active stage's `on_error` 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. @@ -699,6 +701,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`; `whole_body_accumulation_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..2b84d54d87 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 a fixed 120-second accumulation deadline shared across the response chain. Expanding response transformations also share an 8 MiB retained-body budget per session; exhaustion follows the stage's `on_error` policy. 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 9d82e442bc..9b42632d96 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 a fixed 120-second accumulation deadline shared across the response chain. 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 diff --git a/tasks/scripts/stage-prebuilt-binaries.sh b/tasks/scripts/stage-prebuilt-binaries.sh index b3f75bbaba..8f40788ffa 100755 --- a/tasks/scripts/stage-prebuilt-binaries.sh +++ b/tasks/scripts/stage-prebuilt-binaries.sh @@ -171,6 +171,7 @@ build_component_for_arch() { local current_host_os local current_host_arch local binary_path + local cargo_output_dir local build_rustflags resolve_component "$component" @@ -257,7 +258,8 @@ build_component_for_arch() { CARGO_INCREMENTAL=0 mise x -- ${cargo_env[@]+"${cargo_env[@]}"} "${cargo_subcommand[@]}" "${args[@]}" ) - binary_path="${ROOT}/target/${target}/release/${binary}" + cargo_output_dir="$(cd "$ROOT" && mise x -- cargo metadata --format-version=1 --no-deps | jq -er '.target_directory')" + binary_path="${cargo_output_dir}/${target}/release/${binary}" if [[ "$component" == "gateway" ]]; then "$SCRIPT_DIR/verify-glibc-symbols.sh" 2.28 "$binary_path" elif [[ "$component" == "supervisor" ]]; then diff --git a/tasks/test.toml b/tasks/test.toml index bb1fa2e9ab..eb61d52257 100644 --- a/tasks/test.toml +++ b/tasks/test.toml @@ -75,6 +75,7 @@ run = [ # with test-only helpers enabled. "cargo test --workspace --exclude openshell-server", "cargo test -p openshell-server --features test-support", + "cargo test --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml", ] run_windows = "powershell -NoProfile -ExecutionPolicy Bypass -File tasks/scripts/windows-msvc.ps1 test-precommit native" hide = true @@ -267,4 +268,3 @@ description = "Run GPU e2e against a standalone gateway with the Docker compute env = { OPENSHELL_E2E_DOCKER_GPU = "1", OPENSHELL_E2E_DOCKER_TEST = "gpu", OPENSHELL_E2E_DOCKER_FEATURES = "e2e-docker-gpu" } depends = ["e2e:conformance:build"] run = "OPENSHELL_CONFORMANCE_BIN=\"${OPENSHELL_CONFORMANCE_BIN:-$PWD/target/debug/openshell-conformance}\" e2e/rust/e2e-docker.sh" -