diff --git a/.github/workflows/branch-checks.yml b/.github/workflows/branch-checks.yml index 7b0c77c8fd..a5e5904a3f 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 nextest run --config-file .config/nextest.toml --profile ci --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml - name: Verify telemetry can be compiled out run: | diff --git a/architecture/sandbox.md b/architecture/sandbox.md index 7be66e97e6..d3427480a7 100644 --- a/architecture/sandbox.md +++ b/architecture/sandbox.md @@ -187,6 +187,23 @@ middleware registry validates implementation-owned config. The generic registry and chain runner live in `openshell-supervisor-middleware`; first-party implementations live in `openshell-supervisor-middleware-builtins`. +The selected middleware chain can also inspect the final HTTP response before +it returns to the workload. Stages select header-only, whole-body, or streaming +inspection independently. The relay owns response framing when body bytes can +change. Preflight exposes upstream `Content-Length`, `Content-Encoding`, and +`Content-Range` as read-only metadata, while the relay emits final framing +separately from middleware-visible headers. Stage failures follow policy-local +`on_error`; explicit denials always block delivery. Once delivery has started, +blocking aborts the response. + +The network supervisor represents the destination-selected request and response +pair as one `HttpMiddlewareExchange`. It retains the full chain, runner, request +identity, and policy generation while request and response bindings are selected +independently. The HTTP response adapter owns wire parsing, downstream commit +state, generation fences, framing, and transport error classification. The +generic middleware crate owns stage selection, remote stream lifecycle, ordered +body processing, limits, and result validation. + The supervisor installs policy and middleware registry changes as one runtime generation and preserves the last-known-good generation if preparation fails. Policy-only updates reuse the connected registry, so an external middleware diff --git a/crates/openshell-supervisor-middleware/src/headers.rs b/crates/openshell-supervisor-middleware/src/headers.rs index 1dc37bff20..056ad32cf7 100644 --- a/crates/openshell-supervisor-middleware/src/headers.rs +++ b/crates/openshell-supervisor-middleware/src/headers.rs @@ -305,28 +305,39 @@ fn is_request_protected(name: &str) -> bool { || name.starts_with("x-openshell-credential") } -fn is_response_protected(name: &str) -> bool { +/// Return whether a response header can carry authentication material and must +/// never be exposed to middleware. +#[must_use] +pub fn is_response_credential_header(name: &str) -> bool { + let name = name.to_ascii_lowercase(); matches!( - name, + name.as_str(), "authentication-info" - | "connection" - | "content-encoding" - | "content-length" - | "content-range" - | "keep-alive" | "proxy-authenticate" | "proxy-authentication-info" | "proxy-authorization" - | "proxy-connection" | "set-cookie" - | "te" - | "trailer" - | "transfer-encoding" - | "upgrade" | "www-authenticate" ) || name.starts_with("x-openshell-credential") } +fn is_response_protected(name: &str) -> bool { + is_response_credential_header(name) + || matches!( + name, + "connection" + | "content-encoding" + | "content-length" + | "content-range" + | "keep-alive" + | "proxy-connection" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) +} + fn is_response_remove_only(name: &str) -> bool { matches!( name, @@ -642,6 +653,25 @@ mod tests { } } + #[test] + fn response_authority_keeps_visible_body_metadata_read_only() { + let existing = [ + header("content-length", "5"), + header("content-encoding", "gzip"), + header("content-range", "bytes 0-4/10"), + ]; + for name in ["Content-Length", "Content-Encoding", "Content-Range"] { + for mutation in [ + write(name, "replacement", ExistingHeaderAction::Overwrite), + remove(name), + ] { + let error = apply(HeaderAuthority::Response, &existing, &[], &[mutation]) + .expect_err("read-only response body metadata"); + assert!(matches!(error, HeaderMutationError::Protected { .. })); + } + } + } + #[test] fn response_authority_protects_credential_headers_from_writes_and_removals() { let existing = [header("set-cookie", "session=upstream")]; diff --git a/crates/openshell-supervisor-middleware/src/lib.rs b/crates/openshell-supervisor-middleware/src/lib.rs index 8755ddae61..a0136ef551 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, is_stale_http_response_integrity_header, +}; + 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 { @@ -836,6 +854,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), @@ -3691,7 +3708,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(), @@ -3705,13 +3722,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..bf7b61a454 --- /dev/null +++ b/crates/openshell-supervisor-middleware/src/response.rs @@ -0,0 +1,2615 @@ +// 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. + +mod preflight; +mod validation; + +#[cfg(test)] +use validation::permitted_body_modes; +use validation::{ + BodyAction, CurrentBodyAction, encoded_header_bytes, strip_stale_integrity, + validate_body_result, validate_trailers_result, +}; + +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; +const SESSION_END_TIMEOUT: Duration = Duration::from_millis(10); +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; + +/// Return whether a response metadata field becomes stale after body changes. +#[must_use] +pub fn is_stale_http_response_integrity_header(name: &str) -> bool { + matches!( + name.to_ascii_lowercase().as_str(), + "accept-ranges" + | "etag" + | "content-md5" + | "digest" + | "content-digest" + | "repr-digest" + | "signature" + | "signature-input" + ) +} + +#[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, +} + +impl HttpResponseStageTransport { + async fn end(self, reason: MiddlewareSessionEndReason) { + let _ = tokio::time::timeout(SESSION_END_TIMEOUT, self.end_inner(reason)).await; + } + + async fn end_inner(self, reason: MiddlewareSessionEndReason) { + if self.sender.send(session_end_event(reason)).await.is_err() { + return; + } + self.drain().await; + } + + async fn drain(self) { + let Self { + sender, + mut responses, + } = self; + // Keep the response stream alive while half-closing the request side. + // Dropping both handles together schedules an HTTP/2 CANCEL and can + // discard the terminal event before remote middleware receives it. + drop(sender); + while responses.next().await.is_some() {} + } +} + +#[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() { + transport.end(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(), + }) + } + } +} + +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; + } +} + +async fn handle_opened_preflight_failure( + entry: &DescribedChainEntry, + current_stage: &mut HttpResponseStage, + prior_stages: &mut [HttpResponseStage], + reason: &str, + invocations: &mut Vec, +) -> Option { + current_stage + .end(MiddlewareSessionEndReason::MiddlewareFailure) + .await; + let failure = collect_preflight_failure(entry, reason, invocations); + if failure.is_some() { + end_stages(prior_stages, MiddlewareSessionEndReason::MiddlewareFailure).await; + } + failure +} + +#[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, + } + + struct PreflightLifecycleService { + completion_tx: mpsc::UnboundedSender<(String, Vec)>, + } + + #[derive(Clone)] + struct RemoteResponseService { + session_end_tx: Option>, + } + + #[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); + let session_end_tx = self.session_end_tx.clone(); + 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(end)) => { + if let Some(sender) = &session_end_tx + && let Ok(reason) = MiddlewareSessionEndReason::try_from(end.reason) + { + let _ = sender.send(reason); + } + break; + } + 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))) + } + } + + #[tonic::async_trait] + impl InProcessMiddleware for PreflightLifecycleService { + async fn describe(&self) -> MiddlewareManifest { + response_manifest("test/preflight-lifecycle") + } + + 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 { + unreachable!() + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> Result { + let (sender, receiver) = mpsc::channel(4); + let completion_tx = self.completion_tx.clone(); + tokio::spawn(async move { + let Some(HttpResponseEvent { + event: Some(http_response_event::Event::Preflight(preflight)), + }) = requests.recv().await + else { + return; + }; + let config_value = |name: &str| { + preflight + .config + .as_ref() + .and_then(|config| config.fields.get(name)) + .and_then(|value| value.kind.as_ref()) + .and_then(|kind| match kind { + prost_types::value::Kind::StringValue(value) => Some(value.clone()), + _ => None, + }) + .unwrap_or_default() + }; + let label = config_value("label"); + let behavior = config_value("behavior"); + let inspect = |body_mode, header_mutations| HttpResponsePreflightResult { + action: Some(http_response_preflight_result::Action::Inspect( + HttpResponsePreflightInspect { + body_mode, + header_mutations, + }, + )), + ..Default::default() + }; + let result = match behavior.as_str() { + "stream" => http_response_event_result::Result::PreflightResult(inspect( + HttpResponseBodyMode::StreamBytes as i32, + Vec::new(), + )), + "wrong-envelope" => http_response_event_result::Result::BodyResult( + HttpResponseBodyResult::default(), + ), + "invalid-diagnostics" => { + let mut result = + inspect(HttpResponseBodyMode::HeadersOnly as i32, Vec::new()); + result.reason = "x".repeat(MAX_MIDDLEWARE_REASON_BYTES + 1); + http_response_event_result::Result::PreflightResult(result) + } + "unsupported-body-mode" => http_response_event_result::Result::PreflightResult( + inspect(i32::MAX, Vec::new()), + ), + "invalid-header-mutation" => { + http_response_event_result::Result::PreflightResult(inspect( + HttpResponseBodyMode::HeadersOnly as i32, + vec![write_header("content-length", "1")], + )) + } + "no-action" => http_response_event_result::Result::PreflightResult( + HttpResponsePreflightResult::default(), + ), + behavior => panic!("unknown lifecycle test behavior: {behavior}"), + }; + if sender + .send(Ok(HttpResponseEventResult { + result: Some(result), + })) + .await + .is_err() + { + return; + } + + let mut terminal_reasons = Vec::new(); + while let Some(event) = requests.recv().await { + if let Some(http_response_event::Event::SessionEnd(end)) = event.event + && let Ok(reason) = MiddlewareSessionEndReason::try_from(end.reason) + { + terminal_reasons.push(reason); + } + } + let _ = completion_tx.send((label, terminal_reasons)); + }); + 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 lifecycle_entry(name: &str, order: i32, behavior: &str, on_error: OnError) -> ChainEntry { + let string_value = |value: &str| prost_types::Value { + kind: Some(prost_types::value::Kind::StringValue(value.into())), + }; + ChainEntry { + name: name.into(), + implementation: "test/preflight-lifecycle".into(), + order, + config: prost_types::Struct { + fields: [ + ("label".into(), string_value(name)), + ("behavior".into(), string_value(behavior)), + ] + .into(), + }, + on_error, + } + } + + 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") + })); + } + + #[tokio::test] + async fn invalid_opened_preflight_stages_receive_one_failure_terminal_event() { + for behavior in [ + "wrong-envelope", + "invalid-diagnostics", + "unsupported-body-mode", + "invalid-header-mutation", + "no-action", + ] { + for on_error in [OnError::FailOpen, OnError::FailClosed] { + let (completion_tx, mut completion_rx) = mpsc::unbounded_channel(); + let runner = + ChainRunner::new(Arc::new(PreflightLifecycleService { completion_tx })); + let entries = [ + lifecycle_entry("prior", 0, "stream", OnError::FailClosed), + lifecycle_entry("invalid", 1, behavior, on_error), + ]; + let mut outcome = runner + .preflight_http_response(&entries, input(200)) + .await + .expect("invalid preflight response"); + assert_eq!(outcome.allowed, on_error == OnError::FailOpen); + if let Some(session) = outcome.session.take() { + session.end(MiddlewareSessionEndReason::Normal).await; + } + + let mut completions = BTreeMap::new(); + for _ in 0..2 { + let (label, reasons) = + tokio::time::timeout(Duration::from_secs(1), completion_rx.recv()) + .await + .expect("bounded terminal event delivery") + .expect("opened stage completion"); + assert!(completions.insert(label, reasons).is_none()); + } + assert_eq!( + completions.get("invalid").map(Vec::as_slice), + Some([MiddlewareSessionEndReason::MiddlewareFailure].as_slice()), + "invalid behavior: {behavior}, policy: {on_error:?}" + ); + let prior_reason = if on_error == OnError::FailOpen { + MiddlewareSessionEndReason::Normal + } else { + MiddlewareSessionEndReason::MiddlewareFailure + }; + assert_eq!( + completions.get("prior").map(Vec::as_slice), + Some([prior_reason].as_slice()), + "invalid behavior: {behavior}, policy: {on_error:?}" + ); + assert!(completion_rx.try_recv().is_err()); + } + } + } + + #[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 (session_end_tx, mut session_end_rx) = mpsc::unbounded_channel(); + let service = RemoteResponseService { + session_end_tx: Some(session_end_tx), + }; + let server = tonic::transport::Server::builder() + .add_service(SupervisorMiddlewareServer::new(service.clone())) + .add_service(HttpResponsePreReturnServer::new(service)) + .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()); + assert_eq!( + tokio::time::timeout(Duration::from_secs(1), session_end_rx.recv()) + .await + .expect("bounded session end delivery"), + Some(MiddlewareSessionEndReason::Normal) + ); + assert!(session_end_rx.try_recv().is_err()); + + 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-middleware/src/response/preflight.rs b/crates/openshell-supervisor-middleware/src/response/preflight.rs new file mode 100644 index 0000000000..e2e53e9bbe --- /dev/null +++ b/crates/openshell-supervisor-middleware/src/response/preflight.rs @@ -0,0 +1,427 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! HTTP response preflight and stage selection. + +use super::validation::{ + body_restriction, permitted_body_modes, strip_stale_integrity, validate_diagnostics, + validate_inspect, validate_preflight_input, +}; +use super::*; + +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?; + self.preflight_described_http_response(described, input) + .await + } + + /// Run preflight with a response-filtered chain that the caller already + /// described. This keeps response parsing and binding selection on one + /// snapshot without repeating remote capability discovery. + pub async fn preflight_described_http_response( + &self, + described: Vec, + input: HttpResponsePreflightInput, + ) -> miette::Result { + 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 mut current_stage = HttpResponseStage { + entry: entry.clone(), + transport: Some(HttpResponseStageTransport { sender, responses }), + mode: StageMode::HeadersOnly, + next_sequence: 1, + whole_body: Vec::new(), + }; + let Some(http_response_event_result::Result::PreflightResult(decision)) = + response.result + else { + if let Some(reason) = handle_opened_preflight_failure( + &entry, + &mut current_stage, + &mut stages, + "unexpected_response_result", + &mut invocations, + ) + .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) = handle_opened_preflight_failure( + &entry, + &mut current_stage, + &mut stages, + reason, + &mut invocations, + ) + .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, + }); + current_stage + .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) = handle_opened_preflight_failure( + &entry, + &mut current_stage, + &mut stages, + &reason, + &mut invocations, + ) + .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) = handle_opened_preflight_failure( + &entry, + &mut current_stage, + &mut stages, + &reason, + &mut invocations, + ) + .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, + }); + current_stage.mode = mode; + if mode == StageMode::HeadersOnly { + current_stage.end(MiddlewareSessionEndReason::Normal).await; + } else { + stages.push(current_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(current_stage); + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareDenial).await; + return Ok(blocked_preflight_outcome( + headers, + crate::MiddlewareDenial { + config_name: entry.entry.name.clone(), + reason_code, + }, + findings, + metadata, + invocations, + )); + } + None => { + if let Some(reason) = handle_opened_preflight_failure( + &entry, + &mut current_stage, + &mut stages, + "invalid_preflight_decision", + &mut invocations, + ) + .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, + }) + } +} diff --git a/crates/openshell-supervisor-middleware/src/response/validation.rs b/crates/openshell-supervisor-middleware/src/response/validation.rs new file mode 100644 index 0000000000..03febd4479 --- /dev/null +++ b/crates/openshell-supervisor-middleware/src/response/validation.rs @@ -0,0 +1,328 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! HTTP response middleware protocol and payload validation. + +use super::*; + +pub(super) enum BodyAction { + PassThrough, + Transform(Vec), + BlockDelivery, + SkipRemaining(CurrentBodyAction), +} + +pub(super) enum CurrentBodyAction { + PassThrough, + Transform(Vec), +} + +pub(super) struct BodyDecision { + pub(super) action: BodyAction, + pub(super) reason_code: String, + pub(super) findings: Vec, + pub(super) metadata: std::collections::HashMap, +} + +pub(super) struct TrailersDecision { + pub(super) headers: Vec, + pub(super) reason_code: String, + pub(super) findings: Vec, + pub(super) metadata: std::collections::HashMap, +} + +pub(super) 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, + }) +} + +pub(super) fn encoded_header_bytes(headers: &[HttpHeader]) -> usize { + headers.iter().fold(0usize, |total, header| { + total.saturating_add(header.encoded_len()) + }) +} + +pub(super) 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) +} + +pub(super) 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) +} + +pub(super) 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(()) +} + +pub(super) 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(()) +} + +pub(super) 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 +} + +pub(super) 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") + ) + }) +} + +pub(super) fn strip_stale_integrity(headers: &mut Vec) { + headers.retain(|header| !is_stale_http_response_integrity_header(&header.name)); +} diff --git a/crates/openshell-supervisor-network/src/l7/middleware.rs b/crates/openshell-supervisor-network/src/l7/middleware.rs index f2df288019..c469ab5233 100644 --- a/crates/openshell-supervisor-network/src/l7/middleware.rs +++ b/crates/openshell-supervisor-network/src/l7/middleware.rs @@ -26,6 +26,89 @@ pub enum MiddlewareApplyResult { AdmissionExhausted, } +/// One destination-selected middleware chain shared by an HTTP request and +/// its matching response. The request and response phases filter bindings +/// independently, so the full chain must remain available until relay ends. +#[derive(Clone)] +pub struct HttpMiddlewareExchange { + request_id: String, + chain: Vec, + runner: openshell_supervisor_middleware::ChainRunner, + generation_guard: PolicyGenerationGuard, +} + +impl HttpMiddlewareExchange { + pub fn new( + request_id: String, + chain: Vec, + runner: openshell_supervisor_middleware::ChainRunner, + generation_guard: PolicyGenerationGuard, + ) -> Self { + Self { + request_id, + chain, + runner, + generation_guard, + } + } + + pub async fn apply_request( + &self, + request: crate::l7::provider::L7Request, + client: &mut C, + ctx: &L7EvalContext, + scheme: &str, + transformed_body_policy: openshell_supervisor_middleware::TransformedBodyPolicy<'_>, + ) -> Result + where + C: AsyncRead + AsyncWrite + Unpin + Send, + { + apply_middleware_chain_for_scheme_with_request_id( + request, + client, + ctx, + scheme, + self.chain.clone(), + &self.runner, + &self.generation_guard, + transformed_body_policy, + &self.request_id, + ) + .await + } + + pub fn response_relay<'a>( + &'a self, + request: &crate::l7::provider::L7Request, + ctx: &'a L7EvalContext, + scheme: &str, + ) -> crate::l7::rest::HttpResponseMiddlewareRelay<'a> { + let sandbox = openshell_ocsf::ctx::ctx(); + crate::l7::rest::HttpResponseMiddlewareRelay { + chain: &self.chain, + runner: &self.runner, + request_context: openshell_core::proto::RequestContext { + request_id: self.request_id.clone(), + 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: super::relay::policy_safe_response_query(&request.query_params), + }, + policy_name: &ctx.policy_name, + generation_guard: Some(&self.generation_guard), + whole_body_timeout: super::rest::DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT, + } + } +} + /// How traffic a middleware chain can never inspect (h2c, non-HTTP TCP, /// protocols without an L7 relay) must be handled for a matching chain. /// @@ -197,7 +280,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 +506,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 +515,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 +526,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 +543,7 @@ pub async fn apply_middleware_chain_for_scheme, + request_id: &str, ) -> Result { if chain.is_empty() { return Ok(MiddlewareApplyResult::Allowed(req)); @@ -479,7 +568,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 +729,32 @@ pub(super) fn middleware_request_input( } } +#[cfg(test)] +#[allow(clippy::too_many_arguments)] +pub(super) fn middleware_request_input( + sandbox: &openshell_ocsf::EventContext, + 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 +1234,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 +1243,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 aa036dee94..4bdd036a45 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)] @@ -398,13 +398,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 { @@ -425,6 +432,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, + } +} + +pub(super) 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, @@ -857,13 +948,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, @@ -871,6 +964,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await; let req = match middleware_result? { @@ -975,6 +1069,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 { @@ -1588,11 +1691,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, @@ -1600,6 +1705,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + &request_id, ) .await; let req = match middleware_result? { @@ -1715,6 +1821,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 { @@ -2005,12 +2120,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, @@ -2018,6 +2135,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await? { @@ -2082,6 +2200,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 { @@ -2251,12 +2378,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, @@ -2264,6 +2393,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await? { @@ -2318,6 +2448,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 { @@ -2870,6 +3009,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)?; @@ -2877,19 +3018,24 @@ where return Ok(()); } let runner = engine.middleware_runner()?; + let exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + request_id.clone(), + chain, + runner, + generation_guard.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( - req, - client, - ctx, - chain, - &runner, - generation_guard, - openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, - ) - .await? - { + let result = exchange + .apply_request( + req, + client, + ctx, + "http", + openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + ) + .await?; + let request = match result { MiddlewareApplyResult::Allowed(request) => request, MiddlewareApplyResult::Denied { denial, .. } => { let denied_request = crate::l7::provider::L7Request { @@ -2926,7 +3072,9 @@ where .await?; return Ok(()); } - } + }; + response_selection = Some(exchange); + request } else { req }; @@ -2948,6 +3096,9 @@ 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(|exchange| exchange.response_relay(&req_with_auth, ctx, "http")); // Forward request with credential rewriting and relay the response. // relay_http_request_with_resolver handles both directions: it sends @@ -2963,6 +3114,7 @@ where ..Default::default() }, ctx, + response_middleware, ) .await? else { @@ -3034,6 +3186,7 @@ mod tests { ..Default::default() }, &L7EvalContext::default(), + None, ) .await .unwrap(); @@ -3307,6 +3460,7 @@ mod tests { ..options }, &ctx, + None, ) .await .expect("typed credential denial"); @@ -6230,7 +6384,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, @@ -6239,6 +6393,7 @@ network_policies: &runner, tunnel_engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + "test-request-id", ) .await .expect("apply middleware chain"); @@ -6366,6 +6521,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 d34a05bd1b..2c2e0b0326 100644 --- a/crates/openshell-supervisor-network/src/l7/rest.rs +++ b/crates/openshell-supervisor-network/src/l7/rest.rs @@ -7,12 +7,28 @@ //! policy, and relays allowed requests to upstream. Handles Content-Length //! and chunked transfer encoding for body framing. +mod http_response; + +pub(crate) use http_response::{ + DEFAULT_HTTP_RESPONSE_WHOLE_BODY_TIMEOUT, HttpResponseMiddlewareRelay, +}; +use http_response::{RelayResponseOptions, relay_response}; +#[cfg(test)] +use http_response::{ + http_response_middleware_fail_open_finding_event, http_response_middleware_invocation_events, + parse_connection_close, parse_status_code, response_is_event_stream, + strip_response_integrity_headers, +}; + use crate::l7::provider::{BodyLength, L7Provider, L7Request, RelayOutcome}; 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::{ SecretResolver, contains_reserved_credential_marker, contains_reserved_credential_marker_bytes, rewrite_http_header_block, @@ -20,7 +36,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 +65,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 ", @@ -800,6 +817,20 @@ 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 +} + +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, @@ -1154,6 +1185,7 @@ where websocket: websocket_response, client_requested_upgrade, }, + response_middleware, ) .await?; @@ -3109,230 +3141,6 @@ fn find_crlf(buf: &[u8], start: usize) -> Option { .map(|offset| start + offset) } -#[derive(Clone)] -struct RelayResponseOptions { - websocket_extensions: WebSocketExtensionMode, - client_requested_upgrade: bool, - websocket: Option, -} - -impl Default for RelayResponseOptions { - fn default() -> Self { - Self { - websocket_extensions: WebSocketExtensionMode::Preserve, - client_requested_upgrade: true, - websocket: None, - } - } -} - -async fn relay_response( - request_method: &str, - upstream: &mut U, - client: &mut C, - options: RelayResponseOptions, -) -> Result -where - U: AsyncRead + Unpin, - C: AsyncWrite + Unpin, -{ - let started_at = std::time::Instant::now(); - let mut buf = Vec::with_capacity(4096); - let mut tmp = [0u8; 1024]; - - // Read response headers - loop { - if buf.len() > MAX_HEADER_BYTES { - return Err(miette!("HTTP response headers exceed limit")); - } - - let n = upstream.read(&mut tmp).await.into_diagnostic()?; - if n == 0 { - // Upstream closed — forward whatever we have - if !buf.is_empty() { - client.write_all(&buf).await.into_diagnostic()?; - } - return Ok(RelayOutcome::Consumed); - } - buf.extend_from_slice(&tmp[..n]); - - if buf.windows(4).any(|w| w == b"\r\n\r\n") { - break; - } - } - - let header_end = buf.windows(4).position(|w| w == b"\r\n\r\n").unwrap() + 4; - - // Parse response framing - let header_str = String::from_utf8_lossy(&buf[..header_end]); - let status_code = parse_status_code(&header_str).unwrap_or(200); - let server_wants_close = parse_connection_close(&header_str); - let event_stream = response_is_event_stream(&header_str); - let body_length = parse_body_length(&header_str)?; - - debug!( - status_code, - ?body_length, - server_wants_close, - request_method, - overflow_bytes = buf.len() - header_end, - "relay_response framing" - ); - - // 101 Switching Protocols: the connection has been upgraded (e.g. to - // WebSocket). Forward the 101 headers to the client and signal the - // caller to switch to raw bidirectional TCP relay. Any bytes read - // from upstream beyond the headers are overflow that belong to the - // upgraded protocol and must be forwarded before switching. - if status_code == 101 { - if !options.client_requested_upgrade { - return Ok(RelayOutcome::Consumed); - } - let (websocket_permessage_deflate, websocket_subprotocol) = validate_websocket_response( - &header_str, - options.websocket_extensions, - options.websocket.as_ref(), - )?; - client - .write_all(&buf[..header_end]) - .await - .into_diagnostic()?; - client.flush().await.into_diagnostic()?; - let overflow = buf[header_end..].to_vec(); - debug!( - request_method, - overflow_bytes = overflow.len(), - "101 Switching Protocols — signaling protocol upgrade" - ); - return Ok(RelayOutcome::Upgraded { - overflow, - websocket_permessage_deflate, - websocket_subprotocol, - }); - } - - // Bodiless responses (HEAD, 1xx, 204, 304): forward headers only, skip body - if is_bodiless_response(request_method, status_code) { - client - .write_all(&buf[..header_end]) - .await - .into_diagnostic()?; - client.flush().await.into_diagnostic()?; - return if server_wants_close { - Ok(RelayOutcome::Consumed) - } else { - Ok(RelayOutcome::Reusable) - }; - } - - // No explicit framing (no Content-Length, no Transfer-Encoding). - // Per RFC 7230 §3.3.3 the body is delimited by connection close. - if matches!(body_length, BodyLength::None) { - if server_wants_close || event_stream { - // Server indicated it will close, or this is a streaming response - // such as SSE where the body is intentionally delimited by EOF. - let before_end = &buf[..header_end - 2]; - client.write_all(before_end).await.into_diagnostic()?; - if server_wants_close { - client - .write_all(b"Connection: close\r\n\r\n") - .await - .into_diagnostic()?; - } else { - client.write_all(b"\r\n").await.into_diagnostic()?; - } - let overflow = &buf[header_end..]; - if !overflow.is_empty() { - client.write_all(overflow).await.into_diagnostic()?; - client.flush().await.into_diagnostic()?; - } - 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); - } - // No Connection: close — an HTTP/1.1 keep-alive server that omits - // framing headers has an empty body. Forward headers and continue - // the relay loop instead of blocking on relay_until_eof. - debug!("BodyLength::None without Connection: close — treating body as empty"); - client - .write_all(&buf[..header_end]) - .await - .into_diagnostic()?; - client.flush().await.into_diagnostic()?; - return Ok(RelayOutcome::Reusable); - } - - // Forward response headers + any overflow body bytes - client.write_all(&buf).await.into_diagnostic()?; - let overflow_len = (buf.len() - header_end) as u64; - - // Forward remaining response body - match body_length { - BodyLength::ContentLength(len) => { - let remaining = len.saturating_sub(overflow_len); - if remaining > 0 { - relay_fixed(upstream, client, remaining, None).await?; - } - } - BodyLength::Chunked => { - relay_chunked(upstream, client, &buf[header_end..], None).await?; - } - BodyLength::None => unreachable!(), - } - client.flush().await.into_diagnostic()?; - debug!( - request_method, - elapsed_ms = started_at.elapsed().as_millis(), - "relay_response complete (explicit framing)" - ); - - // When body framing is explicit (Content-Length / Chunked), always report - // the connection as reusable so the relay loop continues. If the server - // sent `Connection: close`, the *next* upstream write will fail and the - // loop will exit via the normal error path. Exiting early here would - // tear down the CONNECT tunnel before the client can detect the close, - // causing ~30 s retry delays in clients like `gh`. - 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"); - } - } - 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, @@ -3630,17 +3438,287 @@ mod tests { 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, + body_gate: Option, + captured_preflight_headers: Option>>>, + } + + #[derive(Clone)] + struct ResponseBodyGate { + entered: Arc, + release: Arc, + } + + #[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 body_gate = self.body_gate.clone(); + let captured_preflight_headers = self.captured_preflight_headers.clone(); + 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 let Some(captured) = &captured_preflight_headers { + *captured.lock().expect("preflight capture lock") = + preflight.headers.clone(); + } + 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) => { + if let Some(gate) = &body_gate { + gate.entered.notify_one(); + gate.release.notified().await; + } + 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, @@ -5511,104 +5589,1308 @@ mod tests { } #[tokio::test] - async fn relay_response_no_framing_with_connection_close_reads_until_eof() { - // Response with Connection: close but no Content-Length/TE: body is - // delimited by connection close — relay_until_eof should forward it. - let response = b"HTTP/1.1 200 OK\r\nConnection: close\r\nServer: test\r\n\r\nhello world"; - - let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); - let (mut client_read, mut client_write) = tokio::io::duplex(4096); - - tokio::spawn(async move { - upstream_write.write_all(response).await.unwrap(); - upstream_write.shutdown().await.unwrap(); - }); + 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, + body_gate: None, + captured_preflight_headers: None, + })); + 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); + } + } + } - let result = tokio::time::timeout( - std::time::Duration::from_secs(2), - relay_response( - "GET", - &mut upstream_read, - &mut client_write, - RelayResponseOptions::default(), - ), + #[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"X-Ordinary: value\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 - .expect("relay_response should not deadlock"); - - let outcome = result.expect("relay_response should succeed"); - assert!( - matches!(outcome, RelayOutcome::Consumed), - "connection consumed by read-until-EOF" - ); - - client_write.shutdown().await.unwrap(); - let mut received = Vec::new(); - client_read.read_to_end(&mut received).await.unwrap(); - let received_str = String::from_utf8_lossy(&received); - assert!( - received_str.contains("Connection: close"), - "should preserve Connection: close" - ); - assert!( - received_str.contains("hello world"), - "body should be forwarded" - ); + .unwrap(); + assert!(matches!(outcome, RelayOutcome::Consumed)); + assert!(delivered.starts_with(b"HTTP/1.1 502 ")); } #[tokio::test] - async fn relay_response_no_framing_event_stream_reads_until_eof() { - let response = - b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\nevent: message\ndata: {}\r\n\r\n"; - - let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); - let (mut client_read, mut client_write) = tokio::io::duplex(4096); - - tokio::spawn(async move { - upstream_write.write_all(response).await.unwrap(); - upstream_write.shutdown().await.unwrap(); - }); - - let result = tokio::time::timeout( - std::time::Duration::from_secs(2), - relay_response( - "GET", - &mut upstream_read, - &mut client_write, - RelayResponseOptions::default(), - ), + async fn response_middleware_omits_credentials_and_preserves_downstream_fields() { + let captured = Arc::new(std::sync::Mutex::new(Vec::new())); + let runner = + openshell_supervisor_middleware::ChainRunner::new(Arc::new(ResponseRelayService { + script: ResponseRelayScript::HeadersOnly, + request_only: false, + body_gate: None, + captured_preflight_headers: Some(Arc::clone(&captured)), + })); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "response".into(), + implementation: "test/response-relay".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: openshell_supervisor_middleware::OnError::FailClosed, + }]; + let response = b"HTTP/1.1 200 OK\r\n\ + Content-Length: 5\r\n\ + Cache-Control: public\r\n\ + Set-Cookie: session=secret; HttpOnly\r\n\ + WWW-Authenticate: Bearer realm=private\r\n\ + Authentication-Info: nextnonce=secret\r\n\ + Proxy-Authenticate: Basic realm=proxy\r\n\ + Proxy-Authentication-Info: nextnonce=proxy-secret\r\n\ + Proxy-Authorization: Basic proxy-secret\r\n\ + X-OpenShell-Credential-Token: injected-secret\r\n\r\nhello"; + 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 - .expect("relay_response should not deadlock"); - - let outcome = result.expect("relay_response should succeed"); - assert!( - matches!(outcome, RelayOutcome::Consumed), - "event stream is consumed by read-until-EOF" - ); + .expect("response relay"); - client_write.shutdown().await.unwrap(); - let mut received = Vec::new(); - client_read.read_to_end(&mut received).await.unwrap(); - let received_str = String::from_utf8_lossy(&received); - assert!(received_str.contains("Content-Type: text/event-stream")); - assert!(received_str.contains("event: message")); + assert!(matches!(outcome, RelayOutcome::Reusable)); + let observed = captured.lock().expect("preflight capture lock"); + assert_eq!( + observed + .iter() + .map(|header| header.name.as_str()) + .collect::>(), + vec!["content-length", "cache-control"] + ); + drop(observed); + + let delivered = String::from_utf8(delivered).expect("UTF-8 response"); + for credential_field in [ + "Set-Cookie: session=secret; HttpOnly\r\n", + "WWW-Authenticate: Bearer realm=private\r\n", + "Authentication-Info: nextnonce=secret\r\n", + "Proxy-Authenticate: Basic realm=proxy\r\n", + "Proxy-Authentication-Info: nextnonce=proxy-secret\r\n", + "Proxy-Authorization: Basic proxy-secret\r\n", + "X-OpenShell-Credential-Token: injected-secret\r\n", + ] { + assert!(delivered.contains(credential_field), "{delivered}"); + } + assert!(delivered.contains("cache-control: private\r\n")); + assert!(delivered.ends_with("\r\n\r\nhello")); } #[tokio::test] - async fn relay_response_no_framing_event_stream_survives_idle_gap() { - let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); - let (mut client_read, mut client_write) = tokio::io::duplex(4096); - - let upstream_task = tokio::spawn(async move { - upstream_write - .write_all( - b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream; charset=utf-8\r\n\r\n", + 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 - .unwrap(); - upstream_write - .write_all(b"event: first\ndata: {}\r\n\r\n") + }); + 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()); + } + } + + fn response_middleware_fixture( + script: ResponseRelayScript, + ) -> ( + openshell_supervisor_middleware::ChainRunner, + Vec, + ) { + response_middleware_fixture_with_error( + script, + openshell_supervisor_middleware::OnError::FailClosed, + ) + } + + fn response_middleware_fixture_with_error( + script: ResponseRelayScript, + on_error: openshell_supervisor_middleware::OnError, + ) -> ( + openshell_supervisor_middleware::ChainRunner, + Vec, + ) { + let runner = + openshell_supervisor_middleware::ChainRunner::new(Arc::new(ResponseRelayService { + script, + request_only: false, + body_gate: None, + captured_preflight_headers: None, + })); + 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 + } + + 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 + } + + 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 response_middleware_headers_only_mutates_head_and_preserves_framing() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\nCache-Control: public\r\n\r\nhello", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("cache-control: private\r\n"), + "{delivered}" + ); + assert!(delivered.contains("Content-Length: 5\r\n"), "{delivered}"); + assert!( + !delivered.to_ascii_lowercase().contains("transfer-encoding"), + "{delivered}" + ); + assert!(delivered.ends_with("\r\n\r\nhello"), "{delivered}"); + } + + #[tokio::test] + async fn response_middleware_headers_only_preserves_chunked_and_close_delimited_bodies() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n2;ext=yes\r\nhe\r\n3\r\nllo\r\n0\r\n\r\n", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("Transfer-Encoding: chunked\r\n"), + "{delivered}" + ); + assert!( + delivered.ends_with("2;ext=yes\r\nhe\r\n3\r\nllo\r\n0\r\n\r\n"), + "{delivered}" + ); + + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\nhello", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + assert!( + String::from_utf8(delivered) + .unwrap() + .ends_with("\r\n\r\nhello") + ); + } + + #[tokio::test] + 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!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.contains("Content-Length: 11\r\n"), "{delivered}"); + assert!( + !delivered.to_ascii_lowercase().contains("etag:"), + "{delivered}" + ); + assert!(delivered.ends_with("\r\n\r\nwhole:hello"), "{delivered}"); + } + + #[tokio::test] + async fn 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!(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 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); + } + + #[tokio::test] + async fn response_middleware_whole_body_block_returns_403_before_commit() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::BlockWholeBody, + ) + .await; + 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 response_middleware_stream_block_aborts_after_commit_without_error_bytes() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + ResponseRelayScript::BlockStream, + ) + .await; + assert!(outcome.is_err()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.1 200 OK\r\n"), "{delivered}"); + assert!(!delivered.contains("middleware_denied"), "{delivered}"); + assert!( + !delivered.contains("response_delivery_failed"), + "{delivered}" + ); + } + + async fn run_response_relay_across_policy_reload( + script: ResponseRelayScript, + ) -> (Result, Vec) { + let entered = Arc::new(tokio::sync::Notify::new()); + let release = Arc::new(tokio::sync::Notify::new()); + let runner = + openshell_supervisor_middleware::ChainRunner::new(Arc::new(ResponseRelayService { + script, + request_only: false, + body_gate: Some(ResponseBodyGate { + entered: Arc::clone(&entered), + release: Arc::clone(&release), + }), + captured_preflight_headers: None, + })); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "response".into(), + implementation: "test/response-relay".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: openshell_supervisor_middleware::OnError::FailClosed, + }]; + let engine = OpaEngine::from_strings(TEST_POLICY, "network_policies: {}\n").unwrap(); + let guard = engine + .generation_guard(engine.current_generation()) + .expect("initial generation guard"); + let mut middleware = response_middleware_context(&runner, &chain, "GET"); + middleware.generation_guard = Some(&guard); + let mut upstream = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello".as_slice(); + let mut delivered = Vec::new(); + let relay = relay_response( + "GET", + &mut upstream, + &mut delivered, + RelayResponseOptions::default(), + Some(middleware), + ); + let reload = async { + entered.notified().await; + engine + .reload(TEST_POLICY, "network_policies: {}\n") + .expect("policy reload"); + release.notify_one(); + }; + let (outcome, ()) = tokio::join!(relay, reload); + (outcome, delivered) + } + + #[tokio::test] + async fn response_middleware_rechecks_generation_after_stream_exchange() { + let (outcome, delivered) = + run_response_relay_across_policy_reload(ResponseRelayScript::Stream).await; + + let error = outcome.expect_err("stale stream output must not be delivered"); + assert!(error.to_string().contains("policy generation is stale")); + assert!(delivered.ends_with(b"\r\n\r\n")); + assert!(!delivered.windows(5).any(|window| window == b"HELLO")); + } + + #[tokio::test] + async fn response_middleware_rechecks_generation_after_whole_body_finish() { + let (outcome, delivered) = + run_response_relay_across_policy_reload(ResponseRelayScript::WholeBody).await; + + let error = outcome.expect_err("stale whole-body output must not be delivered"); + assert!(error.to_string().contains("policy generation is stale")); + assert!(delivered.is_empty()); + } + + #[tokio::test] + 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; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + assert!( + String::from_utf8(delivered) + .unwrap() + .ends_with("\r\n\r\nhello") + ); + + 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; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + assert!( + String::from_utf8(delivered) + .unwrap() + .contains("response_delivery_failed") + ); + } + + #[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}"); + } + + #[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, + ); + 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() + + 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 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; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.contains("Trailer: x-upstream\r\n"), "{delivered}"); + assert!(delivered.contains("5\r\nHELLO\r\n"), "{delivered}"); + assert!(delivered.contains("x-upstream: kept\r\n"), "{delivered}"); + assert!(!delivered.contains("digest:"), "{delivered}"); + assert!(!delivered.contains("ext=yes"), "{delivered}"); + } + + #[tokio::test] + async fn response_middleware_rejects_malformed_response_fields_before_commit() { + for response in [ + b"HTTP/1.1 200 OK\r\nBad Name: value\r\nContent-Length: 0\r\n\r\n".as_slice(), + b"HTTP/1.1 200 OK\r\nTrailer: content-length\r\nTransfer-Encoding: chunked\r\n\r\n0\r\n\r\n".as_slice(), + b"HTTP/1.1 200 OK\r\nConnection: x-private\r\nTrailer: x-private\r\nTransfer-Encoding: chunked\r\n\r\n0\r\n\r\n".as_slice(), + ] { + let (outcome, delivered) = run_response_middleware_relay( + response, + "GET", + ResponseRelayScript::WholeBody, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.starts_with("HTTP/1.1 502 Bad Gateway\r\n"), + "{delivered}" + ); + } + } + + #[tokio::test] + async fn response_middleware_rejects_malformed_or_protected_upstream_trailers_atomically() { + for response in [ + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\nBad Name: value\r\n\r\n".as_slice(), + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n0\r\nContent-Length: 7\r\n\r\n".as_slice(), + ] { + let (outcome, delivered) = run_response_middleware_relay( + response, + "GET", + ResponseRelayScript::WholeBody, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Consumed)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.starts_with("HTTP/1.1 502 Bad Gateway\r\n"), + "{delivered}" + ); + assert!(!delivered.contains("whole:hello"), "{delivered}"); + } + } + + #[tokio::test] + async fn response_middleware_never_uses_chunked_framing_for_http_10() { + for (script, expected_body) in [ + (ResponseRelayScript::HeadersOnly, "hello"), + (ResponseRelayScript::Stream, "HELLO"), + (ResponseRelayScript::WholeBodyWithTrailer, "whole:hello"), + ] { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.0 200 OK\r\nContent-Length: 5\r\n\r\nhello", + "GET", + script, + ) + .await; + assert!(outcome.is_ok()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.0 200 OK\r\n"), "{delivered}"); + assert!( + !delivered.to_ascii_lowercase().contains("transfer-encoding"), + "{delivered}" + ); + assert!( + !delivered.to_ascii_lowercase().contains("trailer:"), + "{delivered}" + ); + assert!(delivered.ends_with(expected_body), "{delivered}"); + } + } + + #[tokio::test] + async fn response_middleware_preserves_baseline_connection_outcomes() { + let (outcome, _) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\nConnection: close\r\n\r\nhello", + "GET", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(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 response_middleware_forwards_interim_head_before_final_preflight() { + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 100 Continue\r\nX-Interim: yes\r\n\r\nHTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok", + "GET", + ResponseRelayScript::WholeBody, + ) + .await; + assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); + let delivered = String::from_utf8(delivered).unwrap(); + assert!(delivered.starts_with("HTTP/1.1 100 Continue\r\nX-Interim: yes\r\n\r\n")); + assert!(delivered.ends_with("whole:ok"), "{delivered}"); + } + + #[tokio::test] + async fn response_middleware_handles_bodyless_responses_without_body_events() { + for response in [ + b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\n\r\n".as_slice(), + b"HTTP/1.1 304 Not Modified\r\nContent-Length: 5\r\n\r\n".as_slice(), + ] { + let (outcome, delivered) = + run_response_middleware_relay(response, "GET", ResponseRelayScript::HeadersOnly) + .await; + assert!(outcome.is_ok()); + let delivered = String::from_utf8(delivered).unwrap(); + assert!( + delivered.contains("cache-control: private\r\n"), + "{delivered}" + ); + assert!(delivered.ends_with("\r\n\r\n"), "{delivered}"); + } + + let (outcome, delivered) = run_response_middleware_relay( + b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\n", + "HEAD", + ResponseRelayScript::HeadersOnly, + ) + .await; + assert!(outcome.is_ok()); + let split = delivered + .windows(4) + .position(|window| window == b"\r\n\r\n") + .unwrap() + + 4; + assert_eq!(&delivered[split..], b""); + } + + #[tokio::test] + async fn 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!(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 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; + 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("\"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 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; + 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""); + } + + #[tokio::test] + 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::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; + 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") + ); + } + } + } + } + + #[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")); + } + } + + #[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}"); + } + } + + #[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()); + } + + #[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()); + } + + #[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 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 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] + async fn relay_response_no_framing_with_connection_close_reads_until_eof() { + // Response with Connection: close but no Content-Length/TE: body is + // delimited by connection close — relay_until_eof should forward it. + let response = b"HTTP/1.1 200 OK\r\nConnection: close\r\nServer: test\r\n\r\nhello world"; + + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); + let (mut client_read, mut client_write) = tokio::io::duplex(4096); + + tokio::spawn(async move { + upstream_write.write_all(response).await.unwrap(); + upstream_write.shutdown().await.unwrap(); + }); + + let result = tokio::time::timeout( + std::time::Duration::from_secs(2), + relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + None, + ), + ) + .await + .expect("relay_response should not deadlock"); + + let outcome = result.expect("relay_response should succeed"); + assert!( + matches!(outcome, RelayOutcome::Consumed), + "connection consumed by read-until-EOF" + ); + + client_write.shutdown().await.unwrap(); + let mut received = Vec::new(); + client_read.read_to_end(&mut received).await.unwrap(); + let received_str = String::from_utf8_lossy(&received); + assert!( + received_str.contains("Connection: close"), + "should preserve Connection: close" + ); + assert!( + received_str.contains("hello world"), + "body should be forwarded" + ); + } + + #[tokio::test] + async fn relay_response_no_framing_event_stream_reads_until_eof() { + let response = + b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\nevent: message\ndata: {}\r\n\r\n"; + + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); + let (mut client_read, mut client_write) = tokio::io::duplex(4096); + + tokio::spawn(async move { + upstream_write.write_all(response).await.unwrap(); + upstream_write.shutdown().await.unwrap(); + }); + + let result = tokio::time::timeout( + std::time::Duration::from_secs(2), + relay_response( + "GET", + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + None, + ), + ) + .await + .expect("relay_response should not deadlock"); + + let outcome = result.expect("relay_response should succeed"); + assert!( + matches!(outcome, RelayOutcome::Consumed), + "event stream is consumed by read-until-EOF" + ); + + client_write.shutdown().await.unwrap(); + let mut received = Vec::new(); + client_read.read_to_end(&mut received).await.unwrap(); + let received_str = String::from_utf8_lossy(&received); + assert!(received_str.contains("Content-Type: text/event-stream")); + assert!(received_str.contains("event: message")); + } + + #[tokio::test] + async fn relay_response_no_framing_event_stream_survives_idle_gap() { + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(4096); + let (mut client_read, mut client_write) = tokio::io::duplex(4096); + + let upstream_task = tokio::spawn(async move { + upstream_write + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream; charset=utf-8\r\n\r\n", + ) + .await + .unwrap(); + upstream_write + .write_all(b"event: first\ndata: {}\r\n\r\n") .await .unwrap(); upstream_write.flush().await.unwrap(); @@ -5626,6 +6908,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5670,6 +6953,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5711,6 +6995,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5749,6 +7034,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5789,6 +7075,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5833,6 +7120,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5876,6 +7164,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5912,6 +7201,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5958,6 +7248,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -6005,6 +7296,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -8091,4 +9383,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/l7/rest/http_response.rs b/crates/openshell-supervisor-network/src/l7/rest/http_response.rs new file mode 100644 index 0000000000..2188eb3e39 --- /dev/null +++ b/crates/openshell-supervisor-network/src/l7/rest/http_response.rs @@ -0,0 +1,2095 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! HTTP response relay and pre-return middleware integration. + +use super::*; + +/// Default wall-clock bound shared by whole-body stages in one response. +pub 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 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, +} + +#[derive(Clone)] +pub(super) struct RelayResponseOptions { + pub(super) websocket_extensions: WebSocketExtensionMode, + pub(super) client_requested_upgrade: bool, + pub(super) websocket: Option, +} + +impl Default for RelayResponseOptions { + fn default() -> Self { + Self { + websocket_extensions: WebSocketExtensionMode::Preserve, + client_requested_upgrade: true, + websocket: None, + } + } +} + +pub(super) async fn relay_response( + request_method: &str, + upstream: &mut U, + client: &mut C, + options: RelayResponseOptions, + response_middleware: Option>, +) -> Result +where + U: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + let started_at = std::time::Instant::now(); + let mut buf = Vec::with_capacity(4096); + let mut tmp = [0u8; 1024]; + + // 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")); + } + + let n = upstream.read(&mut tmp).await.into_diagnostic()?; + if n == 0 { + // Upstream closed — forward whatever we have + if !buf.is_empty() { + client.write_all(&buf).await.into_diagnostic()?; + } + return Ok(RelayOutcome::Consumed); + } + 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; + } + } + + let header_end = buf.windows(4).position(|w| w == b"\r\n\r\n").unwrap() + 4; + + // Parse response framing + let header_str = String::from_utf8_lossy(&buf[..header_end]); + let status_code = parse_status_code(&header_str).unwrap_or(200); + let server_wants_close = parse_connection_close(&header_str); + let event_stream = response_is_event_stream(&header_str); + let body_length = parse_body_length(&header_str)?; + + debug!( + status_code, + ?body_length, + server_wants_close, + request_method, + overflow_bytes = buf.len() - header_end, + "relay_response framing" + ); + + // 101 Switching Protocols: the connection has been upgraded (e.g. to + // WebSocket). Forward the 101 headers to the client and signal the + // caller to switch to raw bidirectional TCP relay. Any bytes read + // from upstream beyond the headers are overflow that belong to the + // upgraded protocol and must be forwarded before switching. + if status_code == 101 { + if !options.client_requested_upgrade { + return Ok(RelayOutcome::Consumed); + } + let (websocket_permessage_deflate, websocket_subprotocol) = validate_websocket_response( + &header_str, + options.websocket_extensions, + options.websocket.as_ref(), + )?; + client + .write_all(&buf[..header_end]) + .await + .into_diagnostic()?; + client.flush().await.into_diagnostic()?; + let overflow = buf[header_end..].to_vec(); + debug!( + request_method, + overflow_bytes = overflow.len(), + "101 Switching Protocols — signaling protocol upgrade" + ); + return Ok(RelayOutcome::Upgraded { + overflow, + websocket_permessage_deflate, + websocket_subprotocol, + }); + } + + 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 + .write_all(&buf[..header_end]) + .await + .into_diagnostic()?; + client.flush().await.into_diagnostic()?; + return if server_wants_close { + Ok(RelayOutcome::Consumed) + } else { + Ok(RelayOutcome::Reusable) + }; + } + + // No explicit framing (no Content-Length, no Transfer-Encoding). + // Per RFC 7230 §3.3.3 the body is delimited by connection close. + if matches!(body_length, BodyLength::None) { + if server_wants_close || event_stream { + // Server indicated it will close, or this is a streaming response + // such as SSE where the body is intentionally delimited by EOF. + let before_end = &buf[..header_end - 2]; + client.write_all(before_end).await.into_diagnostic()?; + if server_wants_close { + client + .write_all(b"Connection: close\r\n\r\n") + .await + .into_diagnostic()?; + } else { + client.write_all(b"\r\n").await.into_diagnostic()?; + } + let overflow = &buf[header_end..]; + if !overflow.is_empty() { + client.write_all(overflow).await.into_diagnostic()?; + client.flush().await.into_diagnostic()?; + } + 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); + } + // No Connection: close — an HTTP/1.1 keep-alive server that omits + // framing headers has an empty body. Forward headers and continue + // the relay loop instead of blocking on relay_until_eof. + debug!("BodyLength::None without Connection: close — treating body as empty"); + client + .write_all(&buf[..header_end]) + .await + .into_diagnostic()?; + client.flush().await.into_diagnostic()?; + return Ok(RelayOutcome::Reusable); + } + + // Forward response headers + any overflow body bytes + client.write_all(&buf).await.into_diagnostic()?; + let overflow_len = (buf.len() - header_end) as u64; + + // Forward remaining response body + match body_length { + BodyLength::ContentLength(len) => { + let remaining = len.saturating_sub(overflow_len); + if remaining > 0 { + relay_fixed(upstream, client, remaining, None).await?; + } + } + BodyLength::Chunked => { + relay_chunked(upstream, client, &buf[header_end..], None).await?; + } + BodyLength::None => unreachable!(), + } + client.flush().await.into_diagnostic()?; + debug!( + request_method, + elapsed_ms = started_at.elapsed().as_millis(), + "relay_response complete (explicit framing)" + ); + + // When body framing is explicit (Content-Length / Chunked), always report + // the connection as reusable so the relay loop continues. If the server + // sent `Connection: close`, the *next* upstream write will fail and the + // loop will exit via the normal error path. Exiting early here would + // tear down the CONNECT tunnel before the client can detect the close, + // causing ~30 s retry delays in clients like `gh`. + Ok(RelayOutcome::Reusable) +} + +#[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), + }; + if let Some(guard) = middleware.generation_guard { + guard.ensure_current()?; + } + 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 preserved_credential_headers = parsed.preserved_credential_headers; + 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_described_http_response(described, input) + .await + } else { + Ok(middleware + .runner + .http_response_input_unrepresentable(&described)) + }; + let mut 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)); + } + }; + if let Some(guard) = middleware.generation_guard + && let Err(error) = guard.ensure_current() + { + if let Some(session) = preflight.session.take() { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) + .await; + } + return Err(error); + } + 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)); + } + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &preflight.invocations, + ); + + 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, + &preserved_credential_headers, + &upstream_declared_trailers, + &buffered[header_end..], + status_code, + body_length, + server_wants_close, + event_stream, + ) + .await?; + return Ok(Some(outcome)); + }; + + 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)); + } + }; + if let Some(guard) = middleware.generation_guard { + guard.ensure_current()?; + } + 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, + &preserved_credential_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 + })); + } + + let whole_body = session.requires_whole_body(); + if whole_body { + session.start_whole_body_deadline(middleware.whole_body_timeout); + } + 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, + &preserved_credential_headers, + if chunked_output { + ResponseFraming::Chunked + } else { + ResponseFraming::Preserve(BodyLength::None) + }, + server_wants_close || close_delimited_output, + downstream_trailers, + ); + let mut committed = !whole_body; + if committed { + if let Err(error) = client.write_all(&streaming_head).await { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect) + .await; + return Err(error).into_diagnostic(); + } + if let Err(error) = client.flush().await { + session + .end(openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect) + .await; + return Err(error).into_diagnostic(); + } + } + + let 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.downcast_ref::().is_some() { + openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect + } else if middleware_stop.is_some_and(|stop| stop.failure.denial.is_some()) { + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareDenial + } else if middleware_stop.is_some() { + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareFailure + } else { + openshell_core::proto::MiddlewareSessionEndReason::UpstreamDisconnect + }; + session.end(end_reason).await; + if committed { + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + true, + ); + return Err(error); + } + debug!(error = %error, "HTTP response processing failed before commitment"); + if let Some(denial) = middleware_stop.and_then(|stop| stop.failure.denial.as_ref()) { + send_response_middleware_denial( + client, + request_method, + middleware.policy_name, + &middleware.target, + denial, + ) + .await?; + } else { + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + } + return Ok(Some(RelayOutcome::Consumed)); + } + }; + + 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 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)); + } + }; + if let Some(guard) = middleware.generation_guard { + guard.ensure_current()?; + } + 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 whole_body && !committed { + let mut headers = preflight.headers; + if finish.strip_stale_integrity_headers { + strip_response_integrity_headers(&mut headers); + } + 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, + &preserved_credential_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_downstream_response_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_downstream_response_chunk(client, unit).await?; + } else { + client.write_all(unit).await.into_diagnostic()?; + } + } + if chunked_output { + write_response_trailers(client, &finish.trailers).await?; + } + } + 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 + }, + )) +} + +#[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], + preserved_credential_headers: &[String], + declared_trailers: &[String], + overflow: &[u8], + status_code: u16, + body_length: BodyLength, + server_wants_close: bool, + event_stream: bool, +) -> Result +where + U: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + let head = serialize_response_head( + status_line, + headers, + preserved_credential_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?; + } + } + 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) +} + +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); + } + } +} + +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); + } +} + +pub(super) fn http_response_middleware_invocation_events( + policy_name: &str, + target: &HttpRequestTarget, + status_code: u16, + invocations: &[openshell_supervisor_middleware::HttpResponseInvocation], +) -> Vec { + invocations + .iter() + .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() +} + +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(super) 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(), + ) +} + +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); +} + +#[derive(Debug)] +struct ParsedResponseHead { + representable: bool, + headers: Vec, + preserved_credential_headers: Vec, + connection_nominated: Vec, + declared_trailers: Vec, +} + +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); + } + } + } + } + for trailer in &declared_trailers { + if is_protected_response_field(trailer) || nominated.contains(trailer) { + return Err(miette!("HTTP response declares a protected trailer field")); + } + } + let mut headers = Vec::new(); + let mut preserved_credential_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 openshell_supervisor_middleware::headers::is_response_credential_header(&name) { + // Middleware must not observe credential-bearing response fields. + // Keep the original line separately so downstream delivery retains + // its exact name, whitespace, and value bytes. + preserved_credential_headers.push(line.to_string()); + continue; + } + if nominated.contains(&name) || is_hidden_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() }, + preserved_credential_headers: if representable { + preserved_credential_headers + } else { + Vec::new() + }, + connection_nominated, + declared_trailers, + }) +} + +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(()) +} + +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(()) +} + +fn is_protected_response_field(name: &str) -> bool { + name.eq_ignore_ascii_case("content-length") + || is_hidden_response_field(name) + || openshell_supervisor_middleware::headers::is_response_credential_header(name) +} + +fn is_hidden_response_field(name: &str) -> bool { + matches!( + name.to_ascii_lowercase().as_str(), + "connection" + | "keep-alive" + | "proxy-authenticate" + | "proxy-authorization" + | "proxy-connection" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) +} + +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")) +} + +#[derive(Clone, Copy)] +enum ResponseFraming { + Preserve(BodyLength), + ContentLength(u64), + Chunked, +} + +fn serialize_response_head( + status_line: &str, + headers: &[HttpHeader], + preserved_credential_headers: &[String], + framing: ResponseFraming, + connection_close: bool, + trailer_names: &[String], +) -> Vec { + let mut output = format!("{status_line}\r\n"); + for header in headers { + // Content-Length is read-only middleware metadata. The relay emits the + // final framing exactly once from `framing` below. + if header.name.eq_ignore_ascii_case("content-length") { + continue; + } + output.push_str(&header.name); + output.push_str(": "); + output.push_str(&header.value); + output.push_str("\r\n"); + } + for header in preserved_credential_headers { + output.push_str(header); + 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() +} + +pub(super) fn strip_response_integrity_headers(headers: &mut Vec) { + headers.retain(|header| { + !openshell_supervisor_middleware::is_stale_http_response_integrity_header(&header.name) + }); +} + +struct BufferedResponseReader<'a, R> { + upstream: &'a mut R, + buffered: &'a [u8], + position: usize, + exact_buffer: Vec, + exact_target: Option, + line_buffer: Vec, +} + +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(), + } + } + + 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)) + } + + 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)) + } + + 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)); + } + } + } +} + +#[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, + generation_guard, + 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], + generation_guard: Option<&'a PolicyGenerationGuard>, + 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 {} + +#[derive(Debug)] +struct ResponseDownstreamWrite(std::io::Error); + +impl fmt::Display for ResponseDownstreamWrite { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(formatter, "HTTP response client write failed: {}", self.0) + } +} + +impl std::error::Error for ResponseDownstreamWrite { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + Some(&self.0) + } +} + +impl miette::Diagnostic for ResponseDownstreamWrite {} + +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; + if let Some(guard) = framing.generation_guard { + guard.ensure_current()?; + } + 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; + if let Some(guard) = framing.generation_guard { + guard.ensure_current()?; + } + 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(ResponseDownstreamWrite)?; + client.flush().await.map_err(ResponseDownstreamWrite)?; + } + if *framing.committed { + for unit in output { + if framing.chunked { + write_downstream_response_chunk(client, &unit).await?; + } else { + client + .write_all(&unit) + .await + .map_err(ResponseDownstreamWrite)?; + } + } + client.flush().await.map_err(ResponseDownstreamWrite)?; + } + Ok(()) +} + +async fn write_downstream_response_chunk( + client: &mut C, + payload: &[u8], +) -> Result<()> { + if payload.is_empty() { + return Ok(()); + } + client + .write_all(format!("{:X}\r\n", payload.len()).as_bytes()) + .await + .map_err(ResponseDownstreamWrite)?; + client + .write_all(payload) + .await + .map_err(ResponseDownstreamWrite)?; + client + .write_all(b"\r\n") + .await + .map_err(ResponseDownstreamWrite)?; + 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`. +pub(super) 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`. +pub(super) 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 +} + +pub(super) 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") + }) +} + +#[cfg(test)] +mod tests { + use super::{ResponseFraming, parse_response_head_for_middleware, serialize_response_head}; + use crate::l7::provider::BodyLength; + use openshell_core::proto::HttpHeader; + + #[test] + fn response_middleware_preflight_keeps_read_only_content_length() { + let parsed = parse_response_head_for_middleware( + b"HTTP/1.1 206 Partial Content\r\n\ + Content-Length: 5\r\n\ + Content-Encoding: gzip\r\n\ + Content-Range: bytes 0-4/10\r\n\ + Connection: x-hop\r\n\ + X-Hop: omitted\r\n\r\n", + ) + .expect("parse response head"); + + assert_eq!( + parsed.headers, + vec![ + HttpHeader { + name: "content-length".into(), + value: "5".into(), + }, + HttpHeader { + name: "content-encoding".into(), + value: "gzip".into(), + }, + HttpHeader { + name: "content-range".into(), + value: "bytes 0-4/10".into(), + }, + ] + ); + + let nominated = parse_response_head_for_middleware( + b"HTTP/1.1 200 OK\r\nConnection: Content-Length\r\nContent-Length: 5\r\n\r\n", + ) + .expect("parse nominated response head"); + assert!(nominated.headers.is_empty()); + } + + #[test] + fn response_middleware_serializes_only_relay_owned_framing() { + let headers = vec![ + HttpHeader { + name: "content-length".into(), + value: "999".into(), + }, + HttpHeader { + name: "Content-Length".into(), + value: "998".into(), + }, + HttpHeader { + name: "content-type".into(), + value: "text/plain".into(), + }, + ]; + + for (framing, expected) in [ + ( + ResponseFraming::Preserve(BodyLength::ContentLength(5)), + Some("Content-Length: 5\r\n"), + ), + ( + ResponseFraming::ContentLength(9), + Some("Content-Length: 9\r\n"), + ), + (ResponseFraming::Chunked, None), + (ResponseFraming::Preserve(BodyLength::None), None), + ] { + let serialized = String::from_utf8(serialize_response_head( + "HTTP/1.1 200 OK", + &headers, + &[], + framing, + false, + &[], + )) + .expect("serialized response head"); + + assert_eq!( + serialized + .lines() + .filter(|line| line.to_ascii_lowercase().starts_with("content-length:")) + .count(), + usize::from(expected.is_some()) + ); + if let Some(expected) = expected { + assert!(serialized.contains(expected)); + } + if matches!(framing, ResponseFraming::Chunked) { + assert!(serialized.contains("Transfer-Encoding: chunked\r\n")); + } + } + } +} diff --git a/crates/openshell-supervisor-network/src/proxy.rs b/crates/openshell-supervisor-network/src/proxy.rs index e3a481941c..d76ee7e6eb 100644 --- a/crates/openshell-supervisor-network/src/proxy.rs +++ b/crates/openshell-supervisor-network/src/proxy.rs @@ -1155,8 +1155,7 @@ struct ForwardL7Reevaluation<'a> { struct ForwardMiddlewarePipeline<'a> { ctx: &'a crate::l7::relay::L7EvalContext, scheme: &'a str, - runner: &'a openshell_supervisor_middleware::ChainRunner, - generation_guard: &'a PolicyGenerationGuard, + exchange: &'a crate::l7::middleware::HttpMiddlewareExchange, l7_reevaluation: Option>, } @@ -1169,7 +1168,6 @@ impl ForwardMiddlewarePipeline<'_> { &self, request: crate::l7::provider::L7Request, client: &mut C, - chain: Vec, ) -> Result where C: TokioAsyncRead + TokioAsyncWrite + Unpin + Send, @@ -1188,17 +1186,15 @@ impl ForwardMiddlewarePipeline<'_> { None => openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, }; - crate::l7::middleware::apply_middleware_chain_for_scheme( - request, - client, - self.ctx, - self.scheme, - chain, - self.runner, - self.generation_guard, - transformed_body_policy, - ) - .await + self.exchange + .apply_request( + request, + client, + self.ctx, + self.scheme, + transformed_body_policy, + ) + .await } } @@ -1648,7 +1644,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[..], @@ -1665,7 +1661,7 @@ async fn handle_tcp_connection( dynamic_credentials, denial_tx.as_ref(), activity_tx.as_ref(), - ) + )) .await; } @@ -3986,6 +3982,13 @@ 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, + exchange: &'a crate::l7::middleware::HttpMiddlewareExchange, } async fn relay_rewritten_forward_request( @@ -4006,16 +4009,22 @@ 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| { + middleware + .exchange + .response_relay(&req, middleware.ctx, middleware.scheme) + }); + + crate::l7::rest::relay_http_request_with_response_middleware_guarded( &req, client, upstream, @@ -4033,6 +4042,7 @@ where host: options.host, port: options.port, }, + response_middleware, ) .await } @@ -4974,8 +4984,11 @@ 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()); - if !chain.is_empty() { + let response_selection = if chain.is_empty() { + None + } else { let middleware_runner = opa_engine.middleware_runner()?; let request = crate::l7::rest::request_from_buffered_http( method, @@ -4991,14 +5004,19 @@ async fn handle_forward_proxy( }), _ => None, }; + let middleware_exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + request_id.clone(), + chain, + middleware_runner, + forward_generation_guard.clone(), + ); let pipeline = ForwardMiddlewarePipeline { ctx: &l7_ctx, scheme: &scheme, - runner: &middleware_runner, - generation_guard: &forward_generation_guard, + exchange: &middleware_exchange, l7_reevaluation, }; - forward_request_bytes = match pipeline.apply(request, client, chain).await? { + forward_request_bytes = match pipeline.apply(request, client).await? { crate::l7::middleware::MiddlewareApplyResult::Allowed(request) => request.raw_header, crate::l7::middleware::MiddlewareApplyResult::Denied { denial, .. } => { emit_activity_simple(activity_tx, true, "middleware"); @@ -5016,7 +5034,8 @@ async fn handle_forward_proxy( return Ok(()); } }; - } + Some(middleware_exchange) + }; let mut middleware_session = if let Some(chain) = websocket_chain.as_deref() { let request = crate::l7::rest::request_from_buffered_http( method, @@ -5338,6 +5357,13 @@ async fn handle_forward_proxy( signing_region, host: &host_lc, port, + response_middleware: response_selection.as_ref().map(|exchange| { + ForwardResponseMiddleware { + ctx: &l7_ctx, + scheme: &scheme, + exchange, + } + }), }, ) .await; @@ -5766,6 +5792,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 { @@ -5981,7 +6119,7 @@ network_policies: }); let (mut proxy_connection, _) = proxy_listener.accept().await.unwrap(); - tokio::time::timeout( + Box::pin(tokio::time::timeout( std::time::Duration::from_secs(30), handle_forward_proxy( "POST", @@ -6001,7 +6139,7 @@ network_policies: None, None, ), - ) + )) .await .expect("MCP forwarding should complete") .expect("handle valid MCP request"); @@ -6097,7 +6235,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(), @@ -6114,7 +6252,7 @@ network_policies: None, None, None, - ), + )), ) .await .expect("denied preflight must complete without an upstream response") @@ -6230,7 +6368,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(), @@ -6247,7 +6385,7 @@ network_policies: None, None, None, - ) + )) .await }); let scenario = tokio::time::timeout(std::time::Duration::from_mins(1), async { @@ -7131,28 +7269,33 @@ network_policies: .next() .expect("built-in middleware service"), ); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "redactor".into(), + implementation: openshell_supervisor_middleware_builtins::BUILTIN_REGEX.into(), + order: 0, + config: prost_types::Struct::default(), + on_error: openshell_supervisor_middleware::OnError::FailClosed, + }]; + let exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + "test-request-id".into(), + chain, + runner, + tunnel_engine.generation_guard().clone(), + ); let pipeline = ForwardMiddlewarePipeline { ctx: &ctx, scheme: "http", - runner: &runner, - generation_guard: tunnel_engine.generation_guard(), + exchange: &exchange, l7_reevaluation: Some(ForwardL7Reevaluation { config: &config, engine: &tunnel_engine, request_info: &request_info, }), }; - let chain = vec![openshell_supervisor_middleware::ChainEntry { - name: "redactor".into(), - implementation: openshell_supervisor_middleware_builtins::BUILTIN_REGEX.into(), - order: 0, - config: prost_types::Struct::default(), - on_error: openshell_supervisor_middleware::OnError::FailClosed, - }]; let (_app, mut client) = tokio::io::duplex(8192); let outcome = pipeline - .apply(request, &mut client, chain) + .apply(request, &mut client) .await .expect("forward middleware pipeline"); @@ -7229,13 +7372,6 @@ network_policies: canonicalize_forward_host_header(raw, "api.example.test").unwrap(), ) .unwrap(); - let pipeline = ForwardMiddlewarePipeline { - ctx: &ctx, - scheme: "http", - runner: &runner, - generation_guard: &guard, - l7_reevaluation: None, - }; let chain = vec![openshell_supervisor_middleware::ChainEntry { name: "blocker".into(), implementation: "test/blocking-forward".into(), @@ -7243,13 +7379,25 @@ network_policies: config: prost_types::Struct::default(), on_error: openshell_supervisor_middleware::OnError::FailClosed, }]; + let exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + "test-request-id".into(), + chain, + runner, + guard.clone(), + ); + let pipeline = ForwardMiddlewarePipeline { + ctx: &ctx, + scheme: "http", + exchange: &exchange, + l7_reevaluation: None, + }; let (_app, mut client) = tokio::io::duplex(8192); let revoke = async { entered.notified().await; 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), revoke); let request = match outcome.expect("middleware pipeline") { crate::l7::middleware::MiddlewareApplyResult::Allowed(request) => request, crate::l7::middleware::MiddlewareApplyResult::Denied { .. } => { @@ -7362,6 +7510,181 @@ 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 exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + "correlated-request-id".into(), + chain, + runner, + guard.clone(), + ); + 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 { + body_classifier: None, + 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", + exchange: &exchange, + }), + }, + ) + .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 exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + "correlated-request-id".into(), + chain, + runner, + guard.clone(), + ); + 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 { + body_classifier: None, + 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", + exchange: &exchange, + }), + }, + ) + .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, @@ -7459,6 +7782,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await?; @@ -7724,6 +8048,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await?; @@ -9844,13 +10169,6 @@ network_policies: .expect("built-in middleware service"), ); let guard = forward_test_guard(); - let pipeline = ForwardMiddlewarePipeline { - ctx: &ctx, - scheme: "http", - runner: &runner, - generation_guard: &guard, - l7_reevaluation: None, - }; let chain = vec![openshell_supervisor_middleware::ChainEntry { name: "redactor".into(), implementation: openshell_supervisor_middleware_builtins::BUILTIN_REGEX.into(), @@ -9858,10 +10176,22 @@ network_policies: config: prost_types::Struct::default(), on_error: openshell_supervisor_middleware::OnError::FailClosed, }]; + let exchange = crate::l7::middleware::HttpMiddlewareExchange::new( + "test-request-id".into(), + chain, + runner, + guard, + ); + let pipeline = ForwardMiddlewarePipeline { + ctx: &ctx, + scheme: "http", + exchange: &exchange, + l7_reevaluation: None, + }; let (_app, mut client) = tokio::io::duplex(8192); let allowed = pipeline - .apply(request, &mut client, chain) + .apply(request, &mut client) .await .expect("middleware pipeline"); let crate::l7::middleware::MiddlewareApplyResult::Allowed(request) = allowed else { @@ -10231,6 +10561,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await @@ -10313,6 +10644,7 @@ network_policies: signing_region: "us-west-2", host: "api.example.com", port: 80, + response_middleware: None, }, ) .await @@ -10402,6 +10734,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await; @@ -10453,6 +10786,7 @@ network_policies: signing_region: "", host: "", port: 0, + response_middleware: None, }, ) .await; diff --git a/docs/extensibility/supervisor-middleware.mdx b/docs/extensibility/supervisor-middleware.mdx index 9f3c723b16..13935436c4 100644 --- a/docs/extensibility/supervisor-middleware.mdx +++ b/docs/extensibility/supervisor-middleware.mdx @@ -22,6 +22,8 @@ For each inspected HTTP request, the supervisor: 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. +Response middleware advertising `HTTP_RESPONSE/PRE_RETURN` inspects the final upstream response before OpenShell returns it to the workload. Its preflight includes upstream `Content-Length`, `Content-Encoding`, and `Content-Range` fields as read-only metadata, except when `Connection` nominates the field as hop-by-hop. Middleware cannot write or remove these fields. OpenShell computes downstream framing separately and emits the final `Content-Length` or `Transfer-Encoding` exactly once after middleware processing. + 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`. @@ -144,6 +146,8 @@ See [Policy Schema](/reference/policy-schema#network-middleware) for the complet | `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. | +A valid upstream response can exceed the 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. + 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. @@ -220,10 +224,16 @@ Middleware activity is emitted through OpenShell's OCSF logging: See [Logging](/observability/logging) for log access and [OCSF JSON Export](/observability/ocsf-json-export) for structured export. +## Runnable example + +The [content guard example](https://github.com/NVIDIA/OpenShell/tree/main/examples/supervisor-middleware-content-guard) matches configured literal terms in UTF-8 request bodies, complete response bodies, and client WebSocket text messages. It supports redaction and denial. Responses require whole-body inspection; unavailable inspection or invalid UTF-8 invokes the configured `on_error` policy. It is not a general PII detector. + +The example includes a policy, local fixture, and smoke launcher. + ## 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`. +- The typed operation and phase pairs are `HTTP_REQUEST/PRE_CREDENTIALS`, `HTTP_RESPONSE/PRE_RETURN`, 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. diff --git a/docs/reference/gateway-config.mdx b/docs/reference/gateway-config.mdx index e2d48087cd..34154c0f5d 100644 --- a/docs/reference/gateway-config.mdx +++ b/docs/reference/gateway-config.mdx @@ -355,13 +355,13 @@ 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 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. 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 96aa9e3c34..e37c8a5317 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 @@ -213,8 +231,12 @@ gateway_id = "$RUN_ID" [[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 } @@ -238,11 +260,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 @@ -253,8 +277,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 } @@ -315,17 +354,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=(--compute-driver "$COMPUTE_DRIVER") + fi printf 'INFO starting gateway\n' - env -u OPENSHELL_COMPUTE_DRIVER "$GATEWAY_BIN" \ + env -u OPENSHELL_DRIVERS -u OPENSHELL_COMPUTE_DRIVER "$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=$! } @@ -353,10 +417,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() { @@ -367,14 +431,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 @@ -393,6 +476,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" @@ -438,15 +539,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 360bf28ee4..76c7f84fdf 100644 --- a/skills/debug-openshell-cluster/SKILL.md +++ b/skills/debug-openshell-cluster/SKILL.md @@ -130,18 +130,9 @@ journalctl -u openshell-gateway --no-pager --lines=200 The gateway calls each interceptor's `Describe` RPC and validates its manifest at startup. Check for unreachable endpoints, invalid RPC/phase bindings, strict `allowlist` or `exact` mismatches, and `post_commit` bindings that resolve to `fail_closed`. If gateway JWT signing is enabled, authenticated network interceptors require HTTPS and a valid bearer token; check the private CA path, endpoint hostname, expected audience, issuer, `kid`, and interceptor logs for token rejection. `allow_insecure_transport = true` explicitly preserves unauthenticated plaintext behavior. If `provider_profile_sources` names an interceptor, that interceptor must advertise provider-profile capability and return a valid, duplicate-free catalog. A selected interceptor-only source is authoritative; include `builtin` or `user` sources explicitly when composition is intended. -For operator-run supervisor middleware, inspect `[[openshell.supervisor.middleware]]`, service reachability, and both gateway and supervisor logs: - -```bash -rg -n 'supervisor|middleware|grpc_endpoint|tls_ca_cert_path|audience|allow_insecure_transport|max_payload_bytes|timeout|gateway_jwt' /etc/openshell/gateway.toml -journalctl -u --no-pager --lines=200 -journalctl -u openshell-gateway --no-pager --lines=200 -openshell logs --tail --source sandbox -``` - -The middleware service must start before the gateway and be reachable from both the gateway and sandbox supervisors. Gateway startup fails if `Describe` is unavailable, a manifest exposes duplicate operation/phase bindings, the registration claims the reserved `openshell/` namespace, or payload and timeout limits are invalid. Supported V1 bindings are `HTTP_REQUEST/PRE_CREDENTIALS` and `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. When gateway JWT signing is disabled, supervisors preserve the legacy unauthenticated connector and do not request extension credentials. When signing is enabled, credential acquisition and verification failures are fail closed: check HTTPS trust and hostname validation, audience and issuer agreement, the token `kid`, gateway `RefreshSandboxToken` errors, and middleware logs. Changing a registration requires a gateway restart. A policy update can also fail before persistence if the selected implementation rejects its `network_middlewares` config. - -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. +If the deployment uses supervisor middleware, follow the +[supervisor middleware troubleshooting reference](references/supervisor-middleware.md) +for startup, authentication, policy validation, and HTTP or WebSocket failures. For network policy validation failures, first distinguish a gateway mutation rejection from a supervisor runtime rejection. Direct policy updates, @@ -745,17 +736,10 @@ configuration — check that the gateway spawned the driver binary you expect | CLI TLS error | Local mTLS bundle does not match server cert/CA | Check `~/.config/openshell/gateways//mtls/` | | Edge or OIDC gateway returns `Unauthenticated` | Stored login expired, audience/scopes mismatch, or gateway auth configuration changed | `openshell gateway info`, `openshell gateway login `, gateway auth logs | | Gateway fails before serving health after enabling an interceptor | Interceptor endpoint unavailable or manifest/binding validation failed | Gateway and interceptor logs; interceptor socket; `binding_policy`, phases, and failure policy | -| Authenticated interceptor or middleware rejects gateway calls | Private CA or hostname mismatch, expected audience or issuer mismatch, stale/unknown `kid`, or malformed extension token | `tls_ca_cert_path`, registration `audience`, service verifier config and logs; fetch well-known metadata only through the already-trusted gateway TLS endpoint | +| Authenticated interceptor rejects gateway calls | Private CA or hostname mismatch, expected audience or issuer mismatch, stale/unknown `kid`, or malformed extension token | `tls_ca_cert_path`, registration `audience`, service verifier config and logs; fetch well-known metadata only through the already-trusted gateway TLS endpoint | | Provider profiles disappear after enabling an interceptor catalog | `provider_profile_sources` selected only an authoritative interceptor or returned invalid/duplicate IDs | Inspect source list and interceptor `Describe`/catalog logs; include `builtin` and `user` when intended | -| Gateway fails after registering supervisor middleware | Service unavailable, invalid manifest, duplicate binding, reserved name, or invalid payload/timeout limit | Middleware service and gateway logs; `[[openshell.supervisor.middleware]]`; `Describe` response | -| Policy update rejects `network_middlewares` | Unknown middleware name, implementation-owned config invalid, duplicate order, broad/invalid host selector, or fail-closed coverage of `tls: skip` | Policy error, gateway logs, middleware `ValidateConfig`, selector and order fields | | 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` | -| 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 | -| Supervisor repeatedly fails to install middleware after enabling gateway JWT signing | Extension credential minting, distribution, or authenticated service connection failed; last-known-good registry remains active | Gateway `RefreshSandboxToken` logs, sandbox configuration events, service token-verification logs, registration TLS/audience settings | | Custom compute driver is unavailable | Driver process/socket missing, inaccessible, or selected name does not match its endpoint/config key | Socket ownership/mode, driver service logs, gateway `GetCapabilities` logs | | Sandbox remains `Stopping` or `Starting` | Driver stop/start failed, retained resource is missing, or a fresh supervisor has not connected | Gateway and driver logs; `docker inspect`, `podman inspect`, Agent Sandbox status/PVC, or VM state marker and launcher process | | Image pull failure | Gateway or sandbox image cannot be pulled | Runtime events and image pull credentials | diff --git a/skills/debug-openshell-cluster/references/supervisor-middleware.md b/skills/debug-openshell-cluster/references/supervisor-middleware.md new file mode 100644 index 0000000000..a1d7524662 --- /dev/null +++ b/skills/debug-openshell-cluster/references/supervisor-middleware.md @@ -0,0 +1,53 @@ +# Supervisor middleware troubleshooting + +Use this reference when the deployment registers supervisor middleware or a +sandbox policy attaches it through `network_middlewares`. Start with gateway +reachability and compute-platform checks in the [main skill](../SKILL.md). +Use installed CLI help for command syntax and the published +[gateway configuration reference](https://docs.nvidia.com/openshell/latest/reference/gateway-config.md) +for registration settings. + +## Collect diagnostics + +For operator-run supervisor middleware, inspect `[[openshell.supervisor.middleware]]`, service reachability, and both gateway and supervisor logs: + +```shell +rg -n 'supervisor|middleware|grpc_endpoint|tls_ca_cert_path|audience|allow_insecure_transport|max_payload_bytes|timeout|gateway_jwt' /etc/openshell/gateway.toml +journalctl -u --no-pager --lines=200 +journalctl -u openshell-gateway --no-pager --lines=200 +openshell logs --tail --source sandbox +``` + +## Startup and authentication + +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. + +## HTTP response failures + +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. + +## Request and WebSocket failures + +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. + +## Common failure patterns + +| Symptom | Likely cause | Check | +|---|---|---| +| Authenticated middleware rejects gateway calls | Private CA or hostname mismatch, expected audience or issuer mismatch, stale/unknown `kid`, or malformed extension token | `tls_ca_cert_path`, registration `audience`, service verifier config and logs; fetch well-known metadata only through the already-trusted gateway TLS endpoint | +| Gateway fails after registering supervisor middleware | Service unavailable, invalid manifest, duplicate binding, reserved name, or invalid payload/timeout limit | Middleware service and gateway logs; `[[openshell.supervisor.middleware]]`; `Describe` response | +| Policy update rejects `network_middlewares` | Unknown middleware name, implementation-owned config invalid, duplicate order, broad/invalid host selector, or fail-closed coverage of `tls: skip` | Policy error, gateway logs, middleware `ValidateConfig`, selector and order fields | +| 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 | +| Supervisor repeatedly fails to install middleware after enabling gateway JWT signing | Extension credential minting, distribution, or authenticated service connection failed; last-known-good registry remains active | Gateway `RefreshSandboxToken` logs, sandbox configuration events, service token-verification logs, registration TLS/audience settings | diff --git a/skills/generate-sandbox-policy/SKILL.md b/skills/generate-sandbox-policy/SKILL.md index 73c0863df7..39ab4cadef 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,14 +209,13 @@ 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`. -- 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`. +- Confirm that the implementation advertises the requested binding: `HTTP_REQUEST/PRE_CREDENTIALS`, `HTTP_RESPONSE/PRE_RETURN`, or `WEBSOCKET_MESSAGE/PRE_CREDENTIALS`. A host match alone does not enable inspection. +- WebSocket middleware inspects client text messages only, over both `ws://` and `wss://`. Binary and upstream-to-client messages pass without inspection, even with `fail_closed`. +- `on_error` controls selected-stage failures. Explicit denials always block traffic. A failed WebSocket stage with `fail_open` can remain bypassed for the rest of the connection. - Default `on_error` to `fail_closed`. Use `fail_open` only when bypassing the stage preserves the user's stated security requirement. - Assign unique `order` values across the complete policy. Lower values run first, and at most 10 configs may be selected. - Match the narrowest destination hosts possible with `endpoints.include`; use `exclude` when a broad selector has trusted exceptions. @@ -380,6 +379,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` - [ ] 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 d986a45cc1..cf713b07d1 100644 --- a/skills/openshell-cli/SKILL.md +++ b/skills/openshell-cli/SKILL.md @@ -498,11 +498,15 @@ 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`. +Middleware can inspect HTTP requests, HTTP responses, or client WebSocket text +messages when the implementation advertises the matching binding. The built-in +`openshell/regex` supports request bodies and client WebSocket text messages. +Use the `generate-sandbox-policy` skill to choose attachments and failure policy, +and `debug-openshell-cluster` to investigate middleware failures. ### Step 5: Push the updated policy diff --git a/tasks/test.toml b/tasks/test.toml index 8c41e70025..727979544c 100644 --- a/tasks/test.toml +++ b/tasks/test.toml @@ -85,6 +85,7 @@ run = [ # with test-only helpers enabled. "cargo test --workspace --exclude openshell-server", "cargo test -p openshell-server --features test-support", + "cargo nextest run --config-file .config/nextest.toml --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