From 132cdc4384cac364aad6bd2dcb7d5a936f01817c Mon Sep 17 00:00:00 2001 From: Piotr Mlocek Date: Mon, 31 Aug 2026 23:01:37 -0700 Subject: [PATCH 1/6] feat(middleware): add HTTP response session runtime Signed-off-by: Piotr Mlocek --- .../src/headers.rs | 6 +- .../src/lib.rs | 17 + .../src/remote.rs | 8 + .../src/response.rs | 2282 +++++++++++++++++ 4 files changed, 2310 insertions(+), 3 deletions(-) create mode 100644 crates/openshell-supervisor-middleware/src/response.rs diff --git a/crates/openshell-supervisor-middleware/src/headers.rs b/crates/openshell-supervisor-middleware/src/headers.rs index 1dc37bff20..5ee9fca2ee 100644 --- a/crates/openshell-supervisor-middleware/src/headers.rs +++ b/crates/openshell-supervisor-middleware/src/headers.rs @@ -138,7 +138,7 @@ pub fn apply( for mutation in mutations { match mutation.operation.as_ref() { Some(header_mutation::Operation::Write(write)) => { - let name = validate_name(&write.name)?; + let name = normalize_name(&write.name)?; validate_authority(authority, MutationKind::Write, &write.name, &name)?; if authority == HeaderAuthority::ResponseTrailers && !existing_headers @@ -193,7 +193,7 @@ pub fn apply( } } Some(header_mutation::Operation::Remove(remove)) => { - let name = validate_name(&remove.name)?; + let name = normalize_name(&remove.name)?; validate_authority(authority, MutationKind::Remove, &remove.name, &name)?; if is_connection_nominated(connection_nominated_headers, &name) { return Err(HeaderMutationError::HopByHop { @@ -217,7 +217,7 @@ fn enforce_size_limit(mutation_bytes: usize) -> Result<(), HeaderMutationError> Ok(()) } -fn validate_name(name: &str) -> Result { +pub fn normalize_name(name: &str) -> Result { let lower = name.to_ascii_lowercase(); if lower.is_empty() || !lower.bytes().all(is_name_token_byte) { return Err(HeaderMutationError::InvalidName { diff --git a/crates/openshell-supervisor-middleware/src/lib.rs b/crates/openshell-supervisor-middleware/src/lib.rs index 1065c18faf..f6a4e720fc 100644 --- a/crates/openshell-supervisor-middleware/src/lib.rs +++ b/crates/openshell-supervisor-middleware/src/lib.rs @@ -5,8 +5,15 @@ pub mod headers; mod remote; +mod response; mod websocket; +pub use response::{ + HttpResponseFinish, HttpResponseInvocation, HttpResponseInvocationOutcome, + HttpResponseMiddlewareFailure, HttpResponsePreflightInput, HttpResponsePreflightOutcome, + HttpResponseSession, MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES, +}; + pub use websocket::{ WebSocketCoverage, WebSocketCoverageState, WebSocketInvocation, WebSocketInvocationOutcome, WebSocketMessageAdmission, WebSocketMessageOutcome, WebSocketMessageType, @@ -626,6 +633,16 @@ impl MiddlewareDispatch { Self::Grpc(service) => service.open_websocket_session(receiver).await, } } + + async fn open_http_response_pre_return( + &self, + receiver: tokio::sync::mpsc::Receiver, + ) -> std::result::Result { + match self { + Self::InProcess(service) => service.open_http_response_pre_return(receiver).await, + Self::Grpc(service) => service.open_http_response_pre_return(receiver).await, + } + } } struct MiddlewareServiceState { 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..a2fe8a7e8e --- /dev/null +++ b/crates/openshell-supervisor-middleware/src/response.rs @@ -0,0 +1,2282 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! HTTP response pre-return middleware chain execution. + +use std::collections::{BTreeMap, BTreeSet}; +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, HeaderMutation, HttpHeader, HttpRequestTarget, HttpResponseBodyEnd, + HttpResponseBodyMode, HttpResponseBodyPassThrough, HttpResponseBodyUnit, HttpResponseEvent, + HttpResponseEventResult, HttpResponsePreflight, HttpResponseSessionEnd, + HttpResponseSessionEndReason, HttpResponseTrailers, RemoveHeader, RequestContext, + header_mutation, http_response_body_result, http_response_body_transform, + http_response_body_unit, http_response_event, http_response_event_result, + http_response_preflight_decision, +}; + +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_PREFLIGHT_TIMEOUT, MAX_MIDDLEWARE_REASON_BYTES, + MAX_MIDDLEWARE_REASON_CODE_BYTES, MAX_MIDDLEWARE_TARGET_BYTES, MiddlewareDiagnosticPolicy, + MiddlewareSessionAdmission, MiddlewareSessionPermit, NamespacedFinding, OnError, headers, + is_stable_reason_code, +}; + +const STREAM_CHANNEL_CAPACITY: usize = 4; +pub const MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES: usize = 64 * 1024; + +#[derive(Debug, Clone)] +pub struct HttpResponsePreflightInput { + pub context: RequestContext, + pub target: HttpRequestTarget, + pub status_code: u16, + /// Parsed upstream Content-Length when present and valid. + pub declared_body_length: Option, + /// Sanitized, lowercased final response headers in wire order. + pub headers: Vec, + /// Lowercased names nominated by the original response's `Connection` + /// fields. Their values are not exposed to middleware. + pub connection_nominated_headers: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum HttpResponseInvocationOutcome { + Skip, + HeadersOnly, + WholeBody, + Stream, + PassThrough, + Transform, + 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 struct HttpResponsePreflightOutcome { + pub allowed: bool, + pub reason: String, + pub headers: Vec, + pub declared_trailer_names: 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, +} + +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 {} + +#[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, +} + +struct HttpResponseStageTransport { + sender: mpsc::Sender, + responses: super::HttpResponseResultStream, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum StageMode { + HeadersOnly, + WholeBody, + Stream, +} + +struct HttpResponseStage { + entry: DescribedChainEntry, + transport: Option, + mode: StageMode, + next_sequence: u64, + declared_trailer_names: BTreeSet, + 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: HttpResponseSessionEndReason) { + if let Some(transport) = self.transport.take() { + let _ = tokio::time::timeout( + Duration::from_millis(10), + transport.sender.send(session_end_event(reason)), + ) + .await; + } + } +} + +pub struct HttpResponseSession { + runner: ChainRunner, + stages: Vec, + connection_nominated_headers: Vec, + findings: Vec, + metadata: BTreeMap>, + invocations: Vec, + session_admission: Option, + body_transformed: bool, + defer_output_until_finish: bool, + deferred_output: Vec>, +} + +impl HttpResponseSession { + #[must_use] + pub fn requires_whole_body(&self) -> bool { + self.stages + .iter() + .any(|stage| stage.is_active() && stage.mode == StageMode::WholeBody) + } + + #[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 + .min(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(), + }); + } + let _work = self + .runner + .reserve_middleware_work_admission() + .await + .map_err(|error| HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {error}"), + })?; + let deadline = Instant::now() + MAX_MIDDLEWARE_CHAIN_TIMEOUT; + let output = self.process_units_from(0, vec![data], deadline).await?; + if self.defer_output_until_finish { + self.deferred_output.extend(output); + Ok(Vec::new()) + } else { + Ok(output) + } + } + + /// Finalize every body stage, process normalized trailers, and end streams. + pub async fn finish( + mut self, + trailers: Vec, + ) -> Result { + let _work = self + .runner + .reserve_middleware_work_admission() + .await + .map_err(|error| HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {error}"), + })?; + 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(HttpResponseSessionEndReason::MiddlewareFailure) + .await; + return Err(failure); + } + }; + if !stage_output.is_empty() { + let output = self + .process_units_from(index + 1, stage_output, deadline) + .await?; + released.extend(output); + } + } + + let trailers = match self.process_trailers(trailers, deadline).await { + Ok(trailers) => trailers, + Err(failure) => { + self.end_all(HttpResponseSessionEndReason::MiddlewareFailure) + .await; + return Err(failure); + } + }; + self.end_all(HttpResponseSessionEndReason::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: HttpResponseSessionEndReason) { + self.end_all(reason).await; + } + + async fn process_units_from( + &mut self, + start: usize, + mut units: Vec>, + deadline: Instant, + ) -> Result>, HttpResponseMiddlewareFailure> { + for index in start..self.stages.len() { + let mut next = Vec::new(); + for unit in units { + let chunk_limit = if self.stages[index].mode == StageMode::Stream { + self.stages[index] + .entry + .max_payload_bytes + .min(MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) + } else { + unit.len().max(1) + }; + if unit.is_empty() { + next.extend(self.process_stage_unit(index, unit, deadline).await?); + } else { + for chunk in unit.chunks(chunk_limit) { + next.extend( + self.process_stage_unit(index, chunk.to_vec(), deadline) + .await?, + ); + } + } + } + units = next; + if units.is_empty() + && self.stages[index + 1..] + .iter() + .all(|stage| stage.mode != StageMode::WholeBody) + { + break; + } + } + Ok(units) + } + + async fn process_stage_unit( + &mut self, + index: usize, + data: Vec, + deadline: Instant, + ) -> Result>, HttpResponseMiddlewareFailure> { + let 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 input_size = data.len(); + let event = body_event(sequence, data.clone()); + let result = match exchange(stage, event, deadline).await { + Ok(result) => result, + Err(reason) => { + return self + .handle_stage_failure(index, &reason, Some(sequence), data) + .await; + } + }; + match validate_body_result(result, sequence, stage.entry.max_payload_bytes) { + Ok(BodyDecision::PassThrough(findings, metadata)) => { + collect_diagnostics( + stage, + findings, + metadata, + &mut self.findings, + &mut self.metadata, + ); + self.invocations.push(body_invocation( + stage, + HttpResponseInvocationOutcome::PassThrough, + sequence, + input_size, + input_size, + )); + Ok(vec![data]) + } + Ok(BodyDecision::Transform(replacement, findings, metadata)) => { + self.body_transformed = true; + collect_diagnostics( + stage, + findings, + metadata, + &mut self.findings, + &mut self.metadata, + ); + self.invocations.push(body_invocation( + stage, + HttpResponseInvocationOutcome::Transform, + sequence, + input_size, + replacement.len(), + )); + if replacement.is_empty() { + Ok(Vec::new()) + } else { + Ok(vec![replacement]) + } + } + Err(reason) => { + self.handle_stage_failure(index, reason, Some(sequence), data) + .await + } + } + } + + async fn finish_stage( + &mut self, + index: usize, + deadline: Instant, + ) -> Result>, HttpResponseMiddlewareFailure> { + let stage = &mut self.stages[index]; + if !stage.is_body_active() { + return Ok(Vec::new()); + } + let mut output = Vec::new(); + if stage.mode == StageMode::WholeBody { + let data = std::mem::take(&mut stage.whole_body); + let sequence = 1; + stage.next_sequence = 2; + let input_size = data.len(); + let result = match exchange(stage, body_event(sequence, data.clone()), deadline).await { + Ok(result) => result, + Err(reason) => { + return self + .handle_stage_failure(index, &reason, Some(sequence), data) + .await; + } + }; + match validate_body_result(result, sequence, stage.entry.max_payload_bytes) { + Ok(BodyDecision::PassThrough(findings, metadata)) => { + collect_diagnostics( + stage, + findings, + metadata, + &mut self.findings, + &mut self.metadata, + ); + self.invocations.push(body_invocation( + stage, + HttpResponseInvocationOutcome::PassThrough, + sequence, + input_size, + input_size, + )); + output.push(data); + } + Ok(BodyDecision::Transform(replacement, findings, metadata)) => { + self.body_transformed = true; + collect_diagnostics( + stage, + findings, + metadata, + &mut self.findings, + &mut self.metadata, + ); + self.invocations.push(body_invocation( + stage, + HttpResponseInvocationOutcome::Transform, + sequence, + input_size, + replacement.len(), + )); + if !replacement.is_empty() { + output.push(replacement); + } + } + Err(reason) => { + return self + .handle_stage_failure(index, reason, Some(sequence), data) + .await; + } + } + } + + let final_sequence = stage.next_sequence.saturating_sub(1); + let body_end_sent = if let Some(transport) = stage.transport.as_ref() { + send_without_result( + &transport.sender, + HttpResponseEvent { + event: Some(http_response_event::Event::BodyEnd(HttpResponseBodyEnd { + final_sequence, + })), + }, + deadline, + stage.entry.timeout, + ) + .await + .is_ok() + } else { + false + }; + if !body_end_sent { + self.handle_stage_failure(index, "middleware_stream_closed", None, Vec::new()) + .await?; + } + Ok(output) + } + + async fn process_trailers( + &mut self, + mut trailers: Vec, + deadline: Instant, + ) -> Result, HttpResponseMiddlewareFailure> { + if self.body_transformed { + strip_stale_integrity(&mut trailers); + } + 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) => { + let original = Vec::new(); + self.handle_stage_failure(index, &reason, None, original) + .await?; + continue; + } + }; + let Some(http_response_event_result::Result::TrailersResult(trailer_result)) = + result.result + else { + self.handle_stage_failure(index, "unexpected_response_result", None, Vec::new()) + .await?; + continue; + }; + if let Err(reason) = validate_diagnostics( + &trailer_result.reason, + "", + &trailer_result.findings, + &trailer_result.metadata, + ) { + self.handle_stage_failure(index, reason, None, Vec::new()) + .await?; + continue; + } + if let Err(reason) = + validate_trailer_mutations(&self.stages[index], &trailer_result.trailer_mutations) + { + self.handle_stage_failure(index, reason, None, Vec::new()) + .await?; + continue; + } + match headers::apply( + headers::HeaderAuthority::ResponseTrailers, + &trailers, + &self.connection_nominated_headers, + &trailer_result.trailer_mutations, + ) { + Ok(updated) => trailers = updated, + Err(error) => { + let reason = self.stages[index].entry.service.as_ref().map_or_else( + || error.to_string(), + |service| { + service + .diagnostic_policy + .header_mutation_error_reason(&error) + }, + ); + self.handle_stage_failure(index, &reason, None, Vec::new()) + .await?; + continue; + } + } + let findings = trailer_result.findings; + let metadata = trailer_result.metadata; + collect_diagnostics( + &self.stages[index], + findings, + metadata, + &mut self.findings, + &mut self.metadata, + ); + } + Ok(trailers) + } + + async fn handle_stage_failure( + &mut self, + index: usize, + reason: &str, + sequence: Option, + original: Vec, + ) -> Result>, HttpResponseMiddlewareFailure> { + let stage = &mut self.stages[index]; + let outcome = if stage.entry.on_error() == OnError::FailOpen { + 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, + }); + stage + .end(HttpResponseSessionEndReason::MiddlewareFailure) + .await; + if stage.entry.on_error() == OnError::FailOpen { + if original.is_empty() { + Ok(Vec::new()) + } else { + Ok(vec![original]) + } + } else { + Err(HttpResponseMiddlewareFailure { + reason: format!("middleware_failed: {reason}"), + }) + } + } + + async fn end_all(&mut self, reason: HttpResponseSessionEndReason) { + for stage in &mut self.stages { + stage.end(reason).await; + } + } +} + +impl ChainRunner { + pub async fn preflight_http_response( + &self, + entries: &[ChainEntry], + input: HttpResponsePreflightInput, + ) -> miette::Result { + validate_preflight_input(&input)?; + let described = self.describe_http_response_chain(entries).await?; + if described.is_empty() { + return Ok(empty_preflight_outcome(input.headers)); + } + let session_admission = match self.try_reserve_middleware_session() { + MiddlewareSessionAdmission::Admitted(admission) => admission, + MiddlewareSessionAdmission::AtCapacity => { + return Ok(response_session_capacity_exhausted( + described, + input.headers, + )); + } + }; + let _work = self.reserve_middleware_work_admission().await?; + let original_restriction = body_restriction(&input); + let mut headers = input.headers.clone(); + let mut stages = Vec::new(); + let mut findings = Vec::new(); + let mut metadata = BTreeMap::new(); + let mut invocations = Vec::new(); + let mut declared_trailer_names = BTreeSet::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, HttpResponseSessionEndReason::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, + }; + let timeout = entry.timeout.min(MAX_MIDDLEWARE_PREFLIGHT_TIMEOUT); + let opened = tokio::time::timeout(timeout, async { + let mut responses = service + .service + .open_http_response_pre_return(receiver) + .await?; + sender + .send(HttpResponseEvent { + event: Some(http_response_event::Event::Preflight(preflight)), + }) + .await + .map_err(|_| tonic::Status::unavailable("middleware request stream closed"))?; + let response = responses.next().await.ok_or_else(|| { + tonic::Status::unavailable("middleware result stream closed") + })??; + Ok::<_, tonic::Status>((responses, response)) + }) + .await; + let (responses, response) = match opened { + Ok(Ok(opened)) => opened, + Ok(Err(error)) => { + let reason = if error.code() == tonic::Code::DeadlineExceeded { + "middleware_timeout".to_string() + } else { + service.diagnostic_policy.error_reason(&error) + }; + if let Some(reason) = + collect_preflight_failure(&entry, &reason, &mut invocations) + { + end_stages(&mut stages, HttpResponseSessionEndReason::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, HttpResponseSessionEndReason::MiddlewareFailure) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + }; + let Some(http_response_event_result::Result::PreflightDecision(decision)) = + response.result + else { + if let Some(reason) = collect_preflight_failure( + &entry, + "unexpected_response_result", + &mut invocations, + ) { + end_stages(&mut stages, HttpResponseSessionEndReason::MiddlewareFailure).await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + }; + match decision.decision { + Some(http_response_preflight_decision::Decision::Skip(skip)) => { + let invalid = validate_diagnostics( + &skip.reason, + &skip.reason_code, + &skip.findings, + &skip.metadata, + ); + if let Err(reason) = invalid { + if let Some(reason) = + collect_preflight_failure(&entry, reason, &mut invocations) + { + end_stages( + &mut stages, + HttpResponseSessionEndReason::MiddlewareFailure, + ) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + collect_preflight_diagnostics( + &entry, + skip.findings, + skip.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: (!skip.reason_code.is_empty()).then_some(skip.reason_code), + }); + let mut skipped = HttpResponseStage { + entry, + transport: Some(HttpResponseStageTransport { sender, responses }), + mode: StageMode::HeadersOnly, + next_sequence: 1, + declared_trailer_names: BTreeSet::new(), + whole_body: Vec::new(), + }; + skipped + .end(HttpResponseSessionEndReason::StageSkipped) + .await; + } + Some(http_response_preflight_decision::Decision::Inspect(inspect)) => { + let mode = match validate_inspect( + &entry, + &inspect, + original_restriction.as_deref(), + input.declared_body_length, + &input.connection_nominated_headers, + ) { + Ok(mode) => mode, + Err(reason) => { + if let Some(reason) = + collect_preflight_failure(&entry, &reason, &mut invocations) + { + end_stages( + &mut stages, + HttpResponseSessionEndReason::MiddlewareFailure, + ) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + }; + let updated = match headers::apply( + headers::HeaderAuthority::Response, + &headers, + &input.connection_nominated_headers, + &inspect.header_mutations, + ) { + Ok(updated) => updated, + Err(error) => { + let reason = service + .diagnostic_policy + .header_mutation_error_reason(&error); + if let Some(reason) = + collect_preflight_failure(&entry, &reason, &mut invocations) + { + end_stages( + &mut stages, + HttpResponseSessionEndReason::MiddlewareFailure, + ) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + }; + headers = updated; + if mode == StageMode::Stream { + strip_stale_integrity(&mut headers); + } + let declared = normalize_declared_trailers( + &inspect.declared_trailer_names, + &input.connection_nominated_headers, + ) + .expect("inspect validation normalized trailer declarations"); + declared_trailer_names.extend(declared.iter().cloned()); + collect_preflight_diagnostics( + &entry, + inspect.findings, + inspect.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: None, + }); + stages.push(HttpResponseStage { + entry, + transport: Some(HttpResponseStageTransport { sender, responses }), + mode, + next_sequence: 1, + declared_trailer_names: declared, + whole_body: Vec::new(), + }); + } + None => { + if let Some(reason) = collect_preflight_failure( + &entry, + "invalid_preflight_decision", + &mut invocations, + ) { + end_stages(&mut stages, HttpResponseSessionEndReason::MiddlewareFailure) + .await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + } + } + } + + if stages.is_empty() { + drop(session_admission); + return Ok(HttpResponsePreflightOutcome { + allowed: true, + reason: String::new(), + headers, + declared_trailer_names: declared_trailer_names.into_iter().collect(), + 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(), + headers, + declared_trailer_names: declared_trailer_names.into_iter().collect(), + session: Some(HttpResponseSession { + runner: self.clone(), + stages, + connection_nominated_headers: input.connection_nominated_headers, + findings: Vec::new(), + metadata: BTreeMap::new(), + invocations: Vec::new(), + session_admission: Some(session_admission), + body_transformed: false, + defer_output_until_finish, + deferred_output: Vec::new(), + }), + findings, + metadata, + invocations, + session_capacity_exhausted: false, + }) + } +} + +enum BodyDecision { + PassThrough(Vec, std::collections::HashMap), + Transform( + Vec, + Vec, + std::collections::HashMap, + ), +} + +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.findings, &body.metadata)?; + match body.decision { + Some(http_response_body_result::Decision::PassThrough(HttpResponseBodyPassThrough {})) => { + Ok(BodyDecision::PassThrough(body.findings, body.metadata)) + } + Some(http_response_body_result::Decision::Transform(transform)) => { + let Some(http_response_body_transform::Replacement::Data(replacement)) = + transform.replacement + else { + return Err("response_body_replacement_missing"); + }; + if replacement.len() > max_payload_bytes { + return Err("response_body_replacement_over_capacity"); + } + Ok(BodyDecision::Transform( + replacement, + body.findings, + body.metadata, + )) + } + None => Err("invalid_response_body_decision"), + } +} + +fn validate_inspect( + entry: &DescribedChainEntry, + inspect: &openshell_core::proto::HttpResponsePreflightInspect, + body_restriction: Option<&str>, + declared_body_length: Option, + connection_nominated_headers: &[String], +) -> Result { + validate_diagnostics(&inspect.reason, "", &inspect.findings, &inspect.metadata) + .map_err(str::to_string)?; + 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 mode != StageMode::HeadersOnly + && let Some(restriction) = body_restriction + { + return Err(restriction.to_string()); + } + if mode == StageMode::WholeBody + && declared_body_length.is_some_and(|length| length > entry.max_payload_bytes as u64) + { + return Err("whole_body_over_capacity".into()); + } + if mode == StageMode::HeadersOnly && !inspect.declared_trailer_names.is_empty() { + return Err("response_trailer_declaration_without_body".into()); + } + if inspect + .header_mutations + .len() + .saturating_add(inspect.declared_trailer_names.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()) + }); + let declared_bytes = inspect + .declared_trailer_names + .iter() + .fold(0usize, |total, name| total.saturating_add(name.len())); + if encoded_mutations.saturating_add(declared_bytes) > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES + { + return Err("header_mutation_bytes_over_capacity".into()); + } + normalize_declared_trailers( + &inspect.declared_trailer_names, + connection_nominated_headers, + )?; + if entry.max_payload_bytes == 0 && mode != StageMode::HeadersOnly { + return Err("response_payload_limit_invalid".into()); + } + Ok(mode) +} + +fn normalize_declared_trailers( + names: &[String], + connection_nominated_headers: &[String], +) -> Result, String> { + let mut normalized = BTreeSet::new(); + for name in names { + let name = headers::normalize_name(name).map_err(|error| error.to_string())?; + let validation = HeaderMutation { + operation: Some(header_mutation::Operation::Remove(RemoveHeader { + name: name.clone(), + })), + }; + headers::apply( + headers::HeaderAuthority::ResponseTrailers, + &[], + connection_nominated_headers, + &[validation], + ) + .map_err(|error| error.to_string())?; + if !normalized.insert(name) { + return Err("response_trailer_declaration_duplicate".into()); + } + } + Ok(normalized) +} + +fn validate_trailer_mutations( + stage: &HttpResponseStage, + mutations: &[HeaderMutation], +) -> Result<(), &'static str> { + if mutations.len() > headers::MAX_HEADER_MUTATIONS { + return Err("header_mutation_count_over_capacity"); + } + if mutations.iter().fold(0usize, |total, mutation| { + total.saturating_add(mutation.encoded_len()) + }) > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES + { + return Err("header_mutation_bytes_over_capacity"); + } + for mutation in mutations { + if let Some(header_mutation::Operation::Write(write)) = mutation.operation.as_ref() { + let name = + headers::normalize_name(&write.name).map_err(|_| "header_mutation_invalid_name")?; + if !stage.declared_trailer_names.contains(&name) { + return Err("response_trailer_name_not_declared"); + } + } + } + Ok(()) +} + +fn validate_preflight_input(input: &HttpResponsePreflightInput) -> miette::Result<()> { + if input.context.encoded_len() > MAX_MIDDLEWARE_CONTEXT_BYTES { + return Err(miette::miette!("response context exceeds platform limit")); + } + if input.target.encoded_len() > MAX_MIDDLEWARE_TARGET_BYTES { + return Err(miette::miette!("response target exceeds platform limit")); + } + if input.headers.len() > MAX_MIDDLEWARE_HEADERS { + return Err(miette::miette!( + "response header count exceeds platform limit" + )); + } + if input.headers.iter().fold(0usize, |total, header| { + total.saturating_add(header.encoded_len()) + }) > MAX_MIDDLEWARE_HEADER_BYTES + { + return Err(miette::miette!("response headers exceed platform limit")); + } + Ok(()) +} + +fn validate_diagnostics( + reason: &str, + reason_code: &str, + findings: &[Finding], + metadata: &std::collections::HashMap, +) -> Result<(), &'static str> { + if reason.len() > MAX_MIDDLEWARE_REASON_BYTES { + return Err("response_reason_over_capacity"); + } + if !reason_code.is_empty() + && (reason_code.len() > MAX_MIDDLEWARE_REASON_CODE_BYTES + || !is_stable_reason_code(reason_code)) + { + return Err("response_reason_code_invalid"); + } + if findings.len() > MAX_MIDDLEWARE_FINDINGS_PER_STAGE { + return Err("response_findings_over_capacity"); + } + if findings + .iter() + .any(|finding| finding.encoded_len() > MAX_MIDDLEWARE_FINDING_BYTES) + { + return Err("response_finding_over_capacity"); + } + if metadata.len() > MAX_MIDDLEWARE_METADATA_ENTRIES { + return Err("response_metadata_count_over_capacity"); + } + if metadata.iter().fold(0usize, |total, (key, value)| { + total.saturating_add(key.len()).saturating_add(value.len()) + }) > MAX_MIDDLEWARE_METADATA_BYTES + { + return Err("response_metadata_bytes_over_capacity"); + } + Ok(()) +} + +fn body_restriction(input: &HttpResponsePreflightInput) -> Option { + if input.target.method.eq_ignore_ascii_case("HEAD") + || input.status_code == 204 + || input.status_code == 304 + { + return Some("bodyless_response".into()); + } + if input.status_code == 206 + || input + .headers + .iter() + .any(|header| header.name.eq_ignore_ascii_case("content-range")) + || input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("content-type") + && header + .value + .split(';') + .next() + .is_some_and(|value| value.trim().eq_ignore_ascii_case("multipart/byteranges")) + }) + { + return Some("unsupported_partial_response".into()); + } + if input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("cache-control") + && header.value.split(',').any(|directive| { + directive + .split('=') + .next() + .is_some_and(|name| name.trim().eq_ignore_ascii_case("no-transform")) + }) + }) { + return Some("response_no_transform".into()); + } + if input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("content-encoding") + && header + .value + .split(',') + .any(|coding| !coding.trim().eq_ignore_ascii_case("identity")) + }) { + return Some("unsupported_content_encoding".into()); + } + None +} + +fn strip_stale_integrity(headers: &mut Vec) { + headers.retain(|header| { + !matches!( + header.name.to_ascii_lowercase().as_str(), + "accept-ranges" + | "etag" + | "content-md5" + | "digest" + | "content-digest" + | "repr-digest" + | "signature" + | "signature-input" + ) + }); +} + +async fn exchange( + stage: &mut HttpResponseStage, + event: HttpResponseEvent, + chain_deadline: Instant, +) -> Result { + let remaining = chain_deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + return Err("middleware_chain_timeout".into()); + } + let timeout = stage.entry.timeout.min(remaining); + let Some(transport) = stage.transport.as_mut() else { + return Err("middleware_stream_closed".into()); + }; + match tokio::time::timeout(timeout, async { + transport + .sender + .send(event) + .await + .map_err(|_| tonic::Status::unavailable("middleware request stream closed"))?; + transport + .responses + .next() + .await + .ok_or_else(|| tonic::Status::unavailable("middleware result stream closed"))? + }) + .await + { + Ok(Ok(result)) => Ok(result), + Ok(Err(error)) => { + let policy = stage + .entry + .service + .as_ref() + .map_or(MiddlewareDiagnosticPolicy::Preserve, |service| { + service.diagnostic_policy + }); + Err(policy.error_reason(&error)) + } + Err(_) => Err("middleware_timeout".into()), + } +} + +async fn send_without_result( + sender: &mpsc::Sender, + event: HttpResponseEvent, + chain_deadline: Instant, + stage_timeout: Duration, +) -> Result<(), ()> { + let remaining = chain_deadline.saturating_duration_since(Instant::now()); + let timeout = stage_timeout.min(remaining); + tokio::time::timeout(timeout, sender.send(event)) + .await + .map_err(|_| ())? + .map_err(|_| ()) +} + +fn body_event(sequence: u64, data: Vec) -> HttpResponseEvent { + HttpResponseEvent { + event: Some(http_response_event::Event::Body(HttpResponseBodyUnit { + sequence, + payload: Some(http_response_body_unit::Payload::Data(data)), + })), + } +} + +fn session_end_event(reason: HttpResponseSessionEndReason) -> HttpResponseEvent { + HttpResponseEvent { + event: Some(http_response_event::Event::SessionEnd( + HttpResponseSessionEnd { + reason: reason as i32, + }, + )), + } +} + +fn body_invocation( + stage: &HttpResponseStage, + outcome: HttpResponseInvocationOutcome, + sequence: u64, + input_size: usize, + output_size: usize, +) -> 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: 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, + declared_trailer_names: BTreeSet::new(), + 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, + }); + fail_closed.then(|| format!("middleware_failed: {reason}")) +} + +fn empty_preflight_outcome(headers: Vec) -> HttpResponsePreflightOutcome { + HttpResponsePreflightOutcome { + allowed: true, + reason: String::new(), + headers, + declared_trailer_names: Vec::new(), + 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, + headers, + declared_trailer_names: Vec::new(), + session: None, + findings, + metadata, + invocations, + session_capacity_exhausted: false, + } +} + +fn response_session_capacity_exhausted( + entries: Vec, + headers: Vec, +) -> HttpResponsePreflightOutcome { + let mut invocations = Vec::new(); + let fail_closed = entries.iter().any(|entry| { + collect_preflight_failure( + entry, + "middleware_session_capacity_exhausted", + &mut invocations, + ) + .is_some() + }); + HttpResponsePreflightOutcome { + allowed: !fail_closed, + reason: if fail_closed { + "middleware_failed: middleware_session_capacity_exhausted".into() + } else { + String::new() + }, + headers, + declared_trailer_names: Vec::new(), + session: None, + findings: Vec::new(), + metadata: BTreeMap::new(), + invocations, + session_capacity_exhausted: true, + } +} + +async fn end_stages(stages: &mut [HttpResponseStage], reason: HttpResponseSessionEndReason) { + for stage in stages { + stage.end(reason).await; + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use openshell_core::middleware::{HttpRequestView, InProcessMiddleware}; + use openshell_core::proto::{ + Decision, ExistingHeaderAction, HttpRequestResult, HttpResponseBodyResult, + HttpResponseBodyTransform, HttpResponsePreflightDecision, HttpResponsePreflightInspect, + HttpResponsePreflightSkip, HttpResponseTrailersResult, MiddlewareBinding, + MiddlewareManifest, WriteHeader, http_response_preflight_decision, + }; + use tokio_stream::wrappers::ReceiverStream; + use tokio_stream::wrappers::TcpListenerStream; + + use super::*; + + #[derive(Clone, Copy)] + enum Script { + HeadersOnly, + Stream, + WholeBody, + InvalidSequence, + Configured, + HangBody, + LargeStream, + Skip, + InvalidSkipReason, + UndeclaredTrailer, + } + + struct ResponseService { + script: Script, + } + + #[derive(Clone)] + struct RemoteResponseService; + + #[tonic::async_trait] + impl openshell_core::proto::middleware::v1::supervisor_middleware_server::SupervisorMiddleware + for RemoteResponseService + { + type EvaluateWebSocketSessionStream = super::super::WebSocketResponseStream; + + async fn describe( + &self, + _request: tonic::Request<()>, + ) -> Result, tonic::Status> { + Ok(tonic::Response::new(response_manifest( + "test/remote-response", + ))) + } + + async fn validate_config( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> + { + Ok(tonic::Response::new( + openshell_core::proto::ValidateConfigResponse { + valid: true, + reason: String::new(), + }, + )) + } + + async fn evaluate_http_request( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Ok(tonic::Response::new(HttpRequestResult { + decision: Decision::Allow as i32, + ..Default::default() + })) + } + + async fn evaluate_web_socket_session( + &self, + _request: tonic::Request< + tonic::Streaming, + >, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented("HTTP response-only service")) + } + } + + #[tonic::async_trait] + impl openshell_core::proto::middleware::v1::http_response_pre_return_server::HttpResponsePreReturn + for RemoteResponseService + { + type EvaluateStream = super::super::HttpResponseResultStream; + + async fn evaluate( + &self, + request: tonic::Request>, + ) -> Result, tonic::Status> { + let mut requests = request.into_inner(); + let (sender, receiver) = mpsc::channel(4); + tokio::spawn(async move { + while let Some(Ok(event)) = requests.next().await { + match event.event { + Some(http_response_event::Event::Preflight(_)) => { + let result = HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightDecision( + HttpResponsePreflightDecision { + decision: Some( + http_response_preflight_decision::Decision::Inspect( + HttpResponsePreflightInspect { + body_mode: + HttpResponseBodyMode::HeadersOnly as i32, + header_mutations: vec![write_header( + "cache-control", + "remote", + )], + ..Default::default() + }, + ), + ), + }, + ), + ), + }; + if sender.send(Ok(result)).await.is_err() { + break; + } + } + Some(http_response_event::Event::SessionEnd(_)) | None => break, + _ => {} + } + } + }); + Ok(tonic::Response::new(Box::pin(ReceiverStream::new(receiver)))) + } + } + + #[tonic::async_trait] + impl InProcessMiddleware for ResponseService { + async fn describe(&self) -> MiddlewareManifest { + MiddlewareManifest { + name: "test/response".into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: openshell_core::proto::SupervisorMiddlewareOperation::HttpResponse + as i32, + phase: openshell_core::proto::SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: if matches!(self.script, Script::LargeStream) { + 128 * 1024 + } else { + 4096 + }, + timeout: if matches!(self.script, Script::HangBody) { + "10ms".into() + } else { + String::new() + }, + }], + expected_audience: String::new(), + } + } + + async fn validate_config( + &self, + _middleware_name: &str, + _config: &prost_types::Struct, + ) -> miette::Result<()> { + Ok(()) + } + + async fn evaluate_http_request( + &self, + _request: HttpRequestView<'_>, + ) -> miette::Result { + Ok(HttpRequestResult { + decision: Decision::Allow as i32, + ..Default::default() + }) + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> Result { + let (sender, receiver) = mpsc::channel(4); + let script = self.script; + tokio::spawn(async move { + let mut selected_script = script; + while let Some(event) = requests.recv().await { + let Some(event) = event.event else { + break; + }; + let result = match event { + http_response_event::Event::Preflight(preflight) => { + if matches!(script, Script::Configured) { + selected_script = match preflight + .config + .as_ref() + .and_then(|config| config.fields.get("mode")) + .and_then(|value| value.kind.as_ref()) + { + Some(prost_types::value::Kind::StringValue(mode)) + if mode == "whole" => + { + Script::WholeBody + } + Some(prost_types::value::Kind::StringValue(mode)) + if mode == "stream" => + { + Script::Stream + } + _ => Script::HeadersOnly, + }; + } + if matches!(selected_script, Script::Skip | Script::InvalidSkipReason) { + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightDecision( + HttpResponsePreflightDecision { + decision: Some( + http_response_preflight_decision::Decision::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, declared_trailer_names) = + match selected_script { + Script::HeadersOnly => ( + HttpResponseBodyMode::HeadersOnly, + vec![write_header("cache-control", "private")], + Vec::new(), + ), + Script::Stream | Script::InvalidSequence => ( + HttpResponseBodyMode::StreamBytes, + Vec::new(), + vec!["digest".into()], + ), + Script::HangBody + | Script::LargeStream + | Script::UndeclaredTrailer => ( + HttpResponseBodyMode::StreamBytes, + Vec::new(), + Vec::new(), + ), + Script::WholeBody => ( + HttpResponseBodyMode::WholeBodyBytes, + Vec::new(), + Vec::new(), + ), + Script::Configured + | Script::Skip + | Script::InvalidSkipReason => unreachable!(), + }; + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightDecision( + HttpResponsePreflightDecision { + decision: Some( + http_response_preflight_decision::Decision::Inspect( + HttpResponsePreflightInspect { + body_mode: body_mode as i32, + header_mutations, + declared_trailer_names, + ..Default::default() + }, + ), + ), + }, + ), + ), + } + } + } + http_response_event::Event::Body(body) => { + if matches!(selected_script, Script::HangBody) { + continue; + } + let Some(http_response_body_unit::Payload::Data(data)) = body.payload + else { + break; + }; + let replacement = match selected_script { + Script::Stream + | Script::InvalidSequence + | Script::LargeStream + | Script::UndeclaredTrailer => data.to_ascii_uppercase(), + Script::WholeBody => [b"whole:".as_slice(), &data].concat(), + Script::HeadersOnly + | Script::Configured + | Script::HangBody + | Script::Skip + | Script::InvalidSkipReason => break, + }; + HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: if matches!( + selected_script, + Script::InvalidSequence + ) { + body.sequence + 1 + } else { + body.sequence + }, + decision: Some( + http_response_body_result::Decision::Transform( + HttpResponseBodyTransform { + replacement: Some( + http_response_body_transform::Replacement::Data( + replacement, + ), + ), + }, + ), + ), + ..Default::default() + }, + )), + } + } + http_response_event::Event::Trailers(_) => HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult { + trailer_mutations: if matches!( + selected_script, + Script::Stream | Script::UndeclaredTrailer + ) { + vec![write_header("digest", "sha-256=:test:")] + } else { + Vec::new() + }, + ..Default::default() + }, + )), + }, + http_response_event::Event::BodyEnd(_) => continue, + http_response_event::Event::SessionEnd(_) => break, + }; + if sender.send(Ok(result)).await.is_err() { + break; + } + } + }); + Ok(Box::pin(ReceiverStream::new(receiver))) + } + } + + fn write_header(name: &str, value: &str) -> HeaderMutation { + HeaderMutation { + operation: Some(header_mutation::Operation::Write(WriteHeader { + name: name.into(), + value: value.into(), + on_existing: ExistingHeaderAction::Overwrite as i32, + })), + } + } + + fn response_manifest(name: &str) -> MiddlewareManifest { + MiddlewareManifest { + name: name.into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: openshell_core::proto::SupervisorMiddlewareOperation::HttpResponse + as i32, + phase: openshell_core::proto::SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: 4096, + timeout: String::new(), + }], + expected_audience: String::new(), + } + } + + fn entry(on_error: OnError) -> ChainEntry { + ChainEntry { + name: "response".into(), + implementation: "test/response".into(), + order: 0, + config: prost_types::Struct::default(), + on_error, + } + } + + fn configured_entry(name: &str, order: i32, mode: &str) -> ChainEntry { + ChainEntry { + name: name.into(), + implementation: "test/response".into(), + order, + config: prost_types::Struct { + fields: [( + "mode".into(), + prost_types::Value { + kind: Some(prost_types::value::Kind::StringValue(mode.into())), + }, + )] + .into(), + }, + on_error: OnError::FailClosed, + } + } + + fn input(status_code: u16) -> HttpResponsePreflightInput { + HttpResponsePreflightInput { + context: RequestContext { + request_id: "req-1".into(), + sandbox_id: "sandbox-1".into(), + ..Default::default() + }, + target: HttpRequestTarget { + scheme: "https".into(), + host: "example.com".into(), + port: 443, + method: "GET".into(), + path: "/data".into(), + query: String::new(), + }, + status_code, + declared_body_length: None, + headers: vec![HttpHeader { + name: "content-type".into(), + value: "text/plain".into(), + }], + connection_nominated_headers: Vec::new(), + } + } + + #[tokio::test] + async fn headers_only_preflight_applies_end_to_end_mutation() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::HeadersOnly, + })); + let outcome = runner + .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) + .await + .expect("response preflight"); + + assert!(outcome.allowed); + assert_eq!( + outcome + .headers + .iter() + .find(|header| header.name == "cache-control") + .map(|header| header.value.as_str()), + Some("private") + ); + outcome + .session + .expect("headers-only session") + .finish(Vec::new()) + .await + .expect("finish headers-only session"); + } + + #[tokio::test] + async fn stream_mode_transforms_lockstep_units_and_declared_trailer() { + 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"); + assert_eq!(outcome.declared_trailer_names, vec!["digest"]); + 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 finish = session.finish(Vec::new()).await.expect("finish stream"); + assert!(finish.body_units.is_empty()); + assert_eq!( + finish.trailers, + vec![HttpHeader { + name: "digest".into(), + value: "sha-256=:test:".into(), + }] + ); + } + + #[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!(pushed.unwrap().is_empty()); + let finish = session.finish(Vec::new()).await.expect("fail-open finish"); + assert_eq!(finish.body_units, vec![original]); + } + } + } + + #[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 session = outcome.session.take().expect("large stream session"); + assert_eq!( + session.stream_unit_limit(), + MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES + ); + 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 undeclared_response_trailer_obeys_on_error() { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::UndeclaredTrailer, + })); + let mut outcome = runner + .preflight_http_response(&[entry(on_error)], 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 finish = session.finish(Vec::new()).await; + assert_eq!(finish.is_ok(), allowed); + if let Ok(finish) = finish { + assert!(finish.trailers.is_empty()); + } + } + } + + #[tokio::test] + async fn invalid_sequence_obeys_fail_open_and_fail_closed() { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::InvalidSequence, + })); + let mut outcome = runner + .preflight_http_response(&[entry(on_error)], input(200)) + .await + .expect("response preflight"); + let mut session = outcome.session.take().expect("stream session"); + let result = session.push_body(b"unchanged".to_vec()).await; + assert_eq!(result.is_ok(), allowed); + if let Ok(units) = result { + assert_eq!(units, vec![b"unchanged".to_vec()]); + } + } + } + + #[tokio::test] + async fn body_inspection_restrictions_obey_fail_open_and_fail_closed() { + let mut cases = Vec::new(); + cases.push(input(206)); + for (name, value) in [ + ("content-range", "bytes 0-3/10"), + ("content-type", "multipart/byteranges; boundary=test"), + ("cache-control", "private, no-transform"), + ("content-encoding", "gzip"), + ] { + let mut candidate = input(200); + candidate.headers.push(HttpHeader { + name: name.into(), + value: value.into(), + }); + cases.push(candidate); + } + for status in [204, 304] { + cases.push(input(status)); + } + let mut head = input(200); + head.target.method = "HEAD".into(); + cases.push(head); + + for candidate in cases { + for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::Stream, + })); + let outcome = runner + .preflight_http_response(&[entry(on_error)], candidate.clone()) + .await + .expect("restricted response preflight"); + assert_eq!(outcome.allowed, allowed); + assert!(outcome.session.is_none()); + } + } + } + + #[tokio::test] + async fn remote_service_executes_through_http_response_pre_return_rpc() { + use openshell_core::proto::middleware::v1::http_response_pre_return_server::HttpResponsePreReturnServer; + use openshell_core::proto::middleware::v1::supervisor_middleware_server::SupervisorMiddlewareServer; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("bind response middleware"); + let address = listener.local_addr().expect("response middleware address"); + let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel(); + let server = tonic::transport::Server::builder() + .add_service(SupervisorMiddlewareServer::new(RemoteResponseService)) + .add_service(HttpResponsePreReturnServer::new(RemoteResponseService)) + .serve_with_incoming_shutdown(TcpListenerStream::new(listener), async { + let _ = shutdown_rx.await; + }); + let server_task = tokio::spawn(server); + let registry = super::super::MiddlewareRegistry::connect_services( + Vec::new(), + vec![openshell_core::proto::SupervisorMiddlewareService { + name: "remote-response".into(), + grpc_endpoint: format!("http://{address}"), + max_payload_bytes: 4096, + allow_insecure_transport: true, + ..Default::default() + }], + ) + .await + .expect("connect remote response middleware"); + let runner = ChainRunner::from_registry(registry); + let outcome = runner + .preflight_http_response( + &[ChainEntry { + name: "response".into(), + implementation: "remote-response".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: OnError::FailClosed, + }], + input(200), + ) + .await + .expect("remote response preflight"); + + assert!(outcome.allowed); + assert_eq!( + outcome + .headers + .iter() + .find(|header| header.name == "cache-control") + .map(|header| header.value.as_str()), + Some("remote") + ); + outcome + .session + .expect("remote response session") + .finish(Vec::new()) + .await + .expect("finish remote response session"); + + let _ = shutdown_tx.send(()); + server_task + .await + .expect("join response middleware server") + .expect("serve response middleware"); + } +} From 649ba78e93b473e541ee7345ad589aa9eac6e5a0 Mon Sep 17 00:00:00 2001 From: Piotr Mlocek Date: Mon, 31 Aug 2026 19:44:39 -0700 Subject: [PATCH 2/6] feat(network): enforce response middleware on HTTP relay Signed-off-by: Piotr Mlocek --- .../src/l7/middleware.rs | 77 +- .../src/l7/relay.rs | 160 +- .../src/l7/rest.rs | 1476 ++++++++++++++++- .../openshell-supervisor-network/src/proxy.rs | 12 +- 4 files changed, 1699 insertions(+), 26 deletions(-) diff --git a/crates/openshell-supervisor-network/src/l7/middleware.rs b/crates/openshell-supervisor-network/src/l7/middleware.rs index 6305653f6a..c50797693c 100644 --- a/crates/openshell-supervisor-network/src/l7/middleware.rs +++ b/crates/openshell-supervisor-network/src/l7/middleware.rs @@ -423,7 +423,8 @@ pub(super) fn middleware_chain_body_limit( .max() } -pub async fn apply_middleware_chain( +#[allow(clippy::too_many_arguments)] +pub async fn apply_middleware_chain_with_request_id( req: crate::l7::provider::L7Request, client: &mut C, ctx: &L7EvalContext, @@ -431,8 +432,9 @@ pub async fn apply_middleware_chain( runner: &openshell_supervisor_middleware::ChainRunner, generation_guard: &PolicyGenerationGuard, transformed_body_policy: openshell_supervisor_middleware::TransformedBodyPolicy<'_>, + request_id: &str, ) -> Result { - apply_middleware_chain_for_scheme( + apply_middleware_chain_for_scheme_with_request_id( req, client, ctx, @@ -441,6 +443,7 @@ pub async fn apply_middleware_chain( runner, generation_guard, transformed_body_policy, + request_id, ) .await } @@ -455,6 +458,35 @@ pub async fn apply_middleware_chain_for_scheme, +) -> Result { + let request_id = uuid::Uuid::new_v4().to_string(); + apply_middleware_chain_for_scheme_with_request_id( + req, + client, + ctx, + scheme, + chain, + runner, + generation_guard, + transformed_body_policy, + &request_id, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +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, + scheme: &str, + chain: Vec, + runner: &openshell_supervisor_middleware::ChainRunner, + generation_guard: &PolicyGenerationGuard, + transformed_body_policy: openshell_supervisor_middleware::TransformedBodyPolicy<'_>, + request_id: &str, ) -> Result { if chain.is_empty() { return Ok(MiddlewareApplyResult::Allowed(req)); @@ -479,7 +511,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 +672,32 @@ pub(super) fn middleware_request_input( } } +#[cfg(test)] +#[allow(clippy::too_many_arguments)] +pub(super) fn middleware_request_input( + sandbox: &openshell_ocsf::SandboxContext, + scheme: &str, + req: &crate::l7::provider::L7Request, + ctx: &L7EvalContext, + headers: Vec<(String, String)>, + connection_nominated_headers: Vec, + query: String, + body: Vec, +) -> openshell_supervisor_middleware::HttpRequestInput { + let request_id = uuid::Uuid::new_v4().to_string(); + middleware_request_input_with_id( + sandbox, + scheme, + req, + ctx, + headers, + connection_nominated_headers, + query, + body, + &request_id, + ) +} + pub(super) fn raw_query_from_request_headers(headers: &[u8]) -> Result { let header_str = std::str::from_utf8(headers).map_err(|_| miette!("HTTP headers contain invalid UTF-8"))?; @@ -1116,7 +1177,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 +1186,13 @@ mod tests { Vec::new(), String::new(), Vec::new(), + "exchange-123", ); assert_eq!(input.sandbox_name, "nightly-build"); assert_eq!(input.sandbox_id, "sbx-123"); assert_eq!(input.workspace, "wrks-default"); + assert_eq!(input.request_id, "exchange-123"); } #[tokio::test] diff --git a/crates/openshell-supervisor-network/src/l7/relay.rs b/crates/openshell-supervisor-network/src/l7/relay.rs index 9e4949b191..a4bc8f00bb 100644 --- a/crates/openshell-supervisor-network/src/l7/relay.rs +++ b/crates/openshell-supervisor-network/src/l7/relay.rs @@ -8,13 +8,14 @@ //! and either forwards or denies the request. use crate::l7::middleware::{ - MiddlewareApplyResult, UninspectableTrafficGate, apply_middleware_chain, - emit_middleware_uninspectable, middleware_network_input, uninspectable_traffic_gate, + MiddlewareApplyResult, UninspectableTrafficGate, apply_middleware_chain_with_request_id, + emit_middleware_uninspectable, middleware_network_input, raw_query_from_request_headers, + uninspectable_traffic_gate, }; #[cfg(test)] use crate::l7::middleware::{ middleware_chain_body_limit, middleware_events, middleware_request_input, - raw_query_from_request_headers, resolve_unbuffered_body, + resolve_unbuffered_body, }; use crate::l7::provider::{L7Provider, RelayOutcome}; use crate::l7::rest::WebSocketExtensionMode; @@ -288,13 +289,20 @@ async fn relay_http_request_with_credential_rejection( upstream: &mut U, options: crate::l7::rest::RelayRequestOptions<'_>, ctx: &L7EvalContext, + response_middleware: Option>, ) -> Result> where C: AsyncRead + AsyncWrite + Unpin, U: AsyncRead + AsyncWrite + Unpin, { - match crate::l7::rest::relay_http_request_with_options_guarded( - request, client, upstream, options, + match Box::pin( + crate::l7::rest::relay_http_request_with_response_middleware_guarded( + request, + client, + upstream, + options, + response_middleware, + ), ) .await { @@ -310,6 +318,39 @@ where } } +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: raw_query_from_request_headers(&request.raw_header).unwrap_or_default(), + }, + policy_name: &ctx.policy_name, + generation_guard, + } +} + #[derive(Default)] pub(crate) struct UpgradeRelayOptions<'a> { pub(crate) websocket_request: bool, @@ -736,13 +777,15 @@ where if allowed || (config.enforcement == EnforcementMode::Audit && !force_deny) { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); let websocket_chain = websocket_request.then(|| chain.clone()); // Route selection resolved `config` per request, so re-check the // body against that protocol's policy after every transforming // stage (a no-op for REST and websocket, whose policy inputs the // chain cannot mutate). let validate = transformed_body_validator(config, &engine, ctx, &request_info); - let middleware_result = apply_middleware_chain( + let middleware_result = apply_middleware_chain_with_request_id( req, client, ctx, @@ -750,6 +793,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await; let req = match middleware_result? { @@ -848,6 +892,15 @@ where port: ctx.port, }, ctx, + Some(http_response_middleware_relay( + &req, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await; let outcome_result = match outcome_result { @@ -1461,11 +1514,13 @@ where if allowed || config.enforcement == EnforcementMode::Audit { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); let websocket_chain = websocket_request.then(|| chain.clone()); // REST and websocket-upgrade policy evaluates only the method, // path, and query, which a middleware result cannot mutate, so no // per-stage body re-check is needed. - let middleware_result = apply_middleware_chain( + let middleware_result = apply_middleware_chain_with_request_id( req, client, ctx, @@ -1473,6 +1528,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + &request_id, ) .await; let req = match middleware_result? { @@ -1587,6 +1643,15 @@ where port: ctx.port, }, ctx, + Some(http_response_middleware_relay( + &req_with_auth, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await; let outcome_result = match outcome_result { @@ -1875,12 +1940,14 @@ where if allowed || (config.enforcement == EnforcementMode::Audit && !force_deny) { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); // Policy admitted the original body above; re-check the body // against the same body-aware policy after every transforming // stage so a middleware cannot smuggle a denied operation to the // upstream or the next stage. let validate = transformed_body_validator(config, engine, ctx, &request_info); - let req = match apply_middleware_chain( + let req = match apply_middleware_chain_with_request_id( req, client, ctx, @@ -1888,6 +1955,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await? { @@ -1946,6 +2014,15 @@ where ..Default::default() }, ctx, + Some(http_response_middleware_relay( + &req, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await? else { @@ -2115,12 +2192,14 @@ where if allowed || (config.enforcement == EnforcementMode::Audit && !force_deny) { let chain = engine.query_middleware_chain(&middleware_network_input(ctx))?; + let response_chain = chain.clone(); + let request_id = uuid::Uuid::new_v4().to_string(); // Policy admitted the original body above; re-check the body // against the same body-aware policy after every transforming // stage so a middleware cannot smuggle a denied operation to the // upstream or the next stage. let validate = transformed_body_validator(config, engine, ctx, &request_info); - let req = match apply_middleware_chain( + let req = match apply_middleware_chain_with_request_id( req, client, ctx, @@ -2128,6 +2207,7 @@ where engine.middleware_runner(), engine.generation_guard(), openshell_supervisor_middleware::TransformedBodyPolicy::Reevaluate(&validate), + &request_id, ) .await? { @@ -2181,6 +2261,15 @@ where ..Default::default() }, ctx, + Some(http_response_middleware_relay( + &req, + ctx, + "https", + &request_id, + &response_chain, + engine.middleware_runner(), + Some(engine.generation_guard()), + )), ) .await? else { @@ -2733,6 +2822,8 @@ where ocsf_emit!(event); } + let request_id = uuid::Uuid::new_v4().to_string(); + let mut response_selection = None; let req = if let Some(engine) = middleware_engine { let input = middleware_network_input(ctx); let (chain, generation) = engine.query_middleware_chain_with_generation(&input)?; @@ -2740,9 +2831,10 @@ where return Ok(()); } let runner = engine.middleware_runner()?; + response_selection = Some((chain.clone(), runner.clone())); // The passthrough path enforces no L7 policy, so there is no // body-aware decision to re-check after a transformation. - match apply_middleware_chain( + match apply_middleware_chain_with_request_id( req, client, ctx, @@ -2750,6 +2842,7 @@ where &runner, generation_guard, openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant, + &request_id, ) .await? { @@ -2811,6 +2904,17 @@ where let scoped_ctx = scoped_context_for_request(ctx, &req_with_auth); let ctx = scoped_ctx.as_ref().unwrap_or(ctx); let resolver = ctx.secret_resolver.as_deref(); + let response_middleware = response_selection.as_ref().map(|(chain, runner)| { + http_response_middleware_relay( + &req_with_auth, + ctx, + "http", + &request_id, + chain, + runner, + Some(generation_guard), + ) + }); // Forward request with credential rewriting and relay the response. // relay_http_request_with_resolver handles both directions: it sends @@ -2826,6 +2930,7 @@ where ..Default::default() }, ctx, + response_middleware, ) .await? else { @@ -3132,6 +3237,7 @@ mod tests { ..options }, &ctx, + None, ) .await .expect("typed credential denial"); @@ -6191,6 +6297,40 @@ 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::new(), + raw_header: b"GET /v1/data?cursor=next 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"); + assert_eq!(response.target.scheme, "https"); + } + #[test] fn middleware_ocsf_events_are_audit_safe() { use openshell_supervisor_middleware::{ diff --git a/crates/openshell-supervisor-network/src/l7/rest.rs b/crates/openshell-supervisor-network/src/l7/rest.rs index 93315a671a..49c222bdba 100644 --- a/crates/openshell-supervisor-network/src/l7/rest.rs +++ b/crates/openshell-supervisor-network/src/l7/rest.rs @@ -12,7 +12,10 @@ use crate::opa::PolicyGenerationGuard; use aws_sigv4::http_request::SignableBody; use base64::Engine as _; use miette::{IntoDiagnostic, Result, miette}; -use openshell_core::proto::{ExistingHeaderAction, HeaderMutation, header_mutation}; +use openshell_core::proto::{ + ExistingHeaderAction, HeaderMutation, HttpHeader, HttpRequestTarget, RequestContext, + header_mutation, +}; use openshell_core::secrets::{ CREDENTIAL_MARKER_SCAN_TAIL_BYTES, SecretResolver, contains_reserved_credential_marker, contains_reserved_credential_marker_bytes, rewrite_http_header_block, @@ -20,7 +23,7 @@ use openshell_core::secrets::{ use openshell_ocsf::ctx::ctx as ocsf_ctx; use sha1::{Digest, Sha1}; use std::collections::{HashMap, HashSet}; -use std::fmt; +use std::fmt::{self, Write as _}; use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use tracing::debug; @@ -49,6 +52,7 @@ async fn max_middleware_body_bytes() -> usize { chain[0].max_payload_bytes() } const RELAY_BUF_SIZE: usize = 8192; +const RESPONSE_UNIT_COALESCE_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(2); const HTTP_METHOD_PREFIXES: &[&[u8]] = &[ b"GET ", b"HEAD ", @@ -798,6 +802,30 @@ pub(crate) async fn relay_http_request_with_options_guarded( upstream: &mut U, options: RelayRequestOptions<'_>, ) -> Result +where + C: AsyncRead + AsyncWrite + Unpin, + U: AsyncRead + AsyncWrite + Unpin, +{ + relay_http_request_with_response_middleware_guarded(req, client, upstream, options, None).await +} + +/// Context retained from request evaluation for the matching response hook. +pub(crate) struct HttpResponseMiddlewareRelay<'a> { + pub(crate) chain: &'a [openshell_supervisor_middleware::ChainEntry], + pub(crate) runner: &'a openshell_supervisor_middleware::ChainRunner, + pub(crate) request_context: RequestContext, + pub(crate) target: HttpRequestTarget, + pub(crate) policy_name: &'a str, + pub(crate) generation_guard: Option<&'a PolicyGenerationGuard>, +} + +pub(crate) async fn relay_http_request_with_response_middleware_guarded( + req: &L7Request, + client: &mut C, + upstream: &mut U, + options: RelayRequestOptions<'_>, + response_middleware: Option>, +) -> Result where C: AsyncRead + AsyncWrite + Unpin, U: AsyncRead + AsyncWrite + Unpin, @@ -1152,6 +1180,7 @@ where websocket: websocket_response, client_requested_upgrade, }, + response_middleware, ) .await?; @@ -3124,6 +3153,7 @@ async fn relay_response( upstream: &mut U, client: &mut C, options: RelayResponseOptions, + response_middleware: Option>, ) -> Result where U: AsyncRead + Unpin, @@ -3133,7 +3163,8 @@ where let mut buf = Vec::with_capacity(4096); let mut tmp = [0u8; 1024]; - // Read response headers + // Read response headers. Forward interim responses unchanged, but retain + // the final response head until response middleware preflight completes. loop { if buf.len() > MAX_HEADER_BYTES { return Err(miette!("HTTP response headers exceed limit")); @@ -3149,6 +3180,21 @@ where } buf.extend_from_slice(&tmp[..n]); + while let Some(position) = buf.windows(4).position(|w| w == b"\r\n\r\n") { + let header_end = position + 4; + let header_str = String::from_utf8_lossy(&buf[..header_end]); + let status_code = parse_status_code(&header_str).unwrap_or(200); + if (100..200).contains(&status_code) && status_code != 101 { + client + .write_all(&buf[..header_end]) + .await + .into_diagnostic()?; + client.flush().await.into_diagnostic()?; + buf.drain(..header_end); + continue; + } + break; + } if buf.windows(4).any(|w| w == b"\r\n\r\n") { break; } @@ -3204,6 +3250,24 @@ where }); } + if let Some(response_middleware) = response_middleware + && let Some(outcome) = Box::pin(relay_response_through_middleware( + request_method, + upstream, + client, + response_middleware, + &buf, + header_end, + status_code, + body_length, + server_wants_close, + event_stream, + )) + .await? + { + return Ok(outcome); + } + // Bodiless responses (HEAD, 1xx, 204, 304): forward headers only, skip body if is_bodiless_response(request_method, status_code) { client @@ -3292,6 +3356,943 @@ where 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]; + let parsed = match parse_response_head_for_middleware(header_bytes) { + Ok(parsed) => parsed, + Err(error) => { + debug!(error = %error, "HTTP response head normalization failed"); + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); + } + }; + let original_headers = parsed.headers.clone(); + let upstream_declared_trailers = parsed.declared_trailers.clone(); + let input = openshell_supervisor_middleware::HttpResponsePreflightInput { + context: middleware.request_context, + target: middleware.target.clone(), + status_code, + declared_body_length: match body_length { + BodyLength::ContentLength(length) => Some(length), + BodyLength::Chunked | BodyLength::None => None, + }, + headers: parsed.headers, + connection_nominated_headers: parsed.connection_nominated, + }; + let preflight = match middleware + .runner + .preflight_http_response(middleware.chain, input) + .await + { + Ok(preflight) => preflight, + Err(error) => { + debug!(error = %error, "HTTP response middleware preflight failed"); + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); + } + }; + if !preflight.allowed { + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &preflight.invocations, + ); + 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 { + // No response binding selected (or every selected stage skipped), so + // retain the existing byte-for-byte relay behavior. + debug_assert_eq!(preflight.headers, original_headers); + return Ok(None); + }; + + let status_line = response_status_line(header_bytes)?; + 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"); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); + } + }; + let mut headers = preflight.headers; + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &finish.invocations, + ); + if finish.strip_stale_integrity_headers { + strip_response_integrity_headers(&mut headers); + } + let head = serialize_response_head( + &status_line, + &headers, + ResponseFraming::Preserve(body_length), + server_wants_close, + &[], + ); + client.write_all(&head).await.into_diagnostic()?; + client.flush().await.into_diagnostic()?; + return Ok(Some(if server_wants_close { + RelayOutcome::Consumed + } else { + RelayOutcome::Reusable + })); + } + + let whole_body = session.requires_whole_body(); + let unit_limit = session.stream_unit_limit().max(1); + let committed = if whole_body { + false + } else { + let mut declared_trailers = upstream_declared_trailers; + for name in &preflight.declared_trailer_names { + if !declared_trailers.contains(name) { + declared_trailers.push(name.clone()); + } + } + let head = serialize_response_head( + &status_line, + &preflight.headers, + ResponseFraming::Chunked, + server_wants_close, + &declared_trailers, + ); + client.write_all(&head).await.into_diagnostic()?; + client.flush().await.into_diagnostic()?; + true + }; + + 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, + committed, + unit_limit, + middleware.generation_guard, + ) + .await; + let trailers = match body_result { + Ok(trailers) => trailers, + Err(error) => { + let end_reason = if middleware + .generation_guard + .is_some_and(PolicyGenerationGuard::is_stale) + { + openshell_core::proto::HttpResponseSessionEndReason::PolicyReload + } else if error + .to_string() + .starts_with("HTTP response middleware failure:") + { + openshell_core::proto::HttpResponseSessionEndReason::MiddlewareFailure + } else { + openshell_core::proto::HttpResponseSessionEndReason::UpstreamError + }; + 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"); + 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::HttpResponseSessionEndReason::PolicyReload) + .await; + return Err(error); + } + + let finish = match session.finish(trailers).await { + Ok(finish) => finish, + Err(error) => { + 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"); + emit_http_response_middleware_failure( + middleware.policy_name, + &middleware.target, + status_code, + false, + ); + send_response_delivery_failure( + client, + request_method, + middleware.policy_name, + &middleware.target, + ) + .await?; + return Ok(Some(RelayOutcome::Consumed)); + } + }; + emit_http_response_middleware_invocations( + middleware.policy_name, + &middleware.target, + status_code, + &finish.invocations, + ); + + if whole_body { + 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() { + ResponseFraming::ContentLength(output_length as u64) + } else { + ResponseFraming::Chunked + }; + let trailer_names: Vec = finish + .trailers + .iter() + .map(|header| header.name.clone()) + .collect(); + let head = serialize_response_head( + &status_line, + &headers, + framing, + server_wants_close, + &trailer_names, + ); + client.write_all(&head).await.into_diagnostic()?; + if matches!(framing, ResponseFraming::Chunked) { + for unit in &finish.body_units { + write_chunk(client, unit).await?; + } + write_response_trailers(client, &finish.trailers).await?; + } else { + for unit in &finish.body_units { + client.write_all(unit).await.into_diagnostic()?; + } + } + } else { + for unit in &finish.body_units { + write_chunk(client, unit).await?; + } + write_response_trailers(client, &finish.trailers).await?; + } + client.flush().await.into_diagnostic()?; + Ok(Some( + if server_wants_close || matches!(body_length, BodyLength::None) { + RelayOutcome::Consumed + } else { + 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); + } +} + +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; + openshell_ocsf::HttpActivityBuilder::new(ocsf_ctx()) + .activity(openshell_ocsf::ActivityId::Other) + .action(if failed { + openshell_ocsf::ActionId::Other + } else { + openshell_ocsf::ActionId::Allowed + }) + .disposition(if failed { + openshell_ocsf::DispositionId::Error + } else { + openshell_ocsf::DispositionId::Allowed + }) + .severity(if failed { + openshell_ocsf::SeverityId::Medium + } else { + openshell_ocsf::SeverityId::Informational + }) + .status(if failed { + 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 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 { + headers: Vec, + connection_nominated: Vec, + declared_trailers: Vec, +} + +fn parse_response_head_for_middleware(header_bytes: &[u8]) -> Result { + let header = std::str::from_utf8(header_bytes) + .map_err(|_| miette!("HTTP response headers contain invalid UTF-8"))?; + if parse_status_code(header).is_none() { + return Err(miette!("HTTP response status line is malformed")); + } + 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; + }; + 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); + } + } + } + } + let mut headers = Vec::new(); + for line in header.split("\r\n").skip(1).filter(|line| !line.is_empty()) { + let (name, value) = line + .split_once(':') + .ok_or_else(|| miette!("Malformed HTTP response header field"))?; + let name = name.to_ascii_lowercase(); + if nominated.contains(&name) + || matches!( + name.as_str(), + "connection" + | "content-length" + | "keep-alive" + | "proxy-authenticate" + | "proxy-authorization" + | "proxy-connection" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) + { + continue; + } + headers.push(HttpHeader { + name, + value: value.trim().to_string(), + }); + } + let mut connection_nominated: Vec<_> = nominated.into_iter().collect(); + connection_nominated.sort(); + Ok(ParsedResponseHead { + headers, + connection_nominated, + declared_trailers, + }) +} + +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], + framing: ResponseFraming, + connection_close: bool, + trailer_names: &[String], +) -> Vec { + let mut output = format!("{status_line}\r\n"); + for header in headers { + output.push_str(&header.name); + output.push_str(": "); + output.push_str(&header.value); + output.push_str("\r\n"); + } + match framing { + ResponseFraming::Preserve(BodyLength::ContentLength(length)) + | ResponseFraming::ContentLength(length) => { + write!(&mut output, "Content-Length: {length}\r\n") + .expect("writing to a String cannot fail"); + } + ResponseFraming::Preserve(BodyLength::Chunked) | ResponseFraming::Chunked => { + output.push_str("Transfer-Encoding: chunked\r\n"); + } + ResponseFraming::Preserve(BodyLength::None) => {} + } + if !trailer_names.is_empty() { + output.push_str("Trailer: "); + output.push_str(&trailer_names.join(", ")); + output.push_str("\r\n"); + } + if connection_close { + output.push_str("Connection: close\r\n"); + } + output.push_str("\r\n"); + output.into_bytes() +} + +fn strip_response_integrity_headers(headers: &mut Vec) { + headers.retain(|header| { + !matches!( + header.name.to_ascii_lowercase().as_str(), + "accept-ranges" + | "etag" + | "content-md5" + | "digest" + | "content-digest" + | "repr-digest" + | "signature" + | "signature-input" + ) + }); +} + +struct BufferedResponseReader<'a, R> { + upstream: &'a mut R, + buffered: &'a [u8], + position: usize, +} + +impl<'a, R: AsyncRead + Unpin> BufferedResponseReader<'a, R> { + fn new(upstream: &'a mut R, buffered: &'a [u8]) -> Self { + Self { + upstream, + buffered, + position: 0, + } + } + + 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> { + let mut output = Vec::with_capacity(length); + while output.len() < length { + let remaining = length - output.len(); + let Some(data) = self.read_some(remaining).await? else { + return Err(miette!("HTTP response body ended unexpectedly")); + }; + output.extend_from_slice(&data); + } + Ok(output) + } + + async fn read_line(&mut self) -> Result> { + let mut line = Vec::new(); + loop { + let Some(byte) = self.read_some(1).await? else { + return Err(miette!("HTTP response ended before line terminator")); + }; + line.push(byte[0]); + if line.len() > MAX_HEADER_BYTES { + return Err(miette!("HTTP response line exceeds limit")); + } + if line.ends_with(b"\r\n") { + line.truncate(line.len() - 2); + return Ok(line); + } + } + } +} + +#[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: bool, + unit_limit: usize, + generation_guard: Option<&PolicyGenerationGuard>, +) -> Result> +where + R: AsyncRead + Unpin, + C: AsyncWrite + Unpin, +{ + let mut pending = Vec::with_capacity(unit_limit); + 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 = reader.read_exact_vec(length).await?; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + remaining -= unit.len() as u64; + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + committed, + unit_limit, + ) + .await?; + } + flush_normalized_response_bytes(session, client, pending, committed).await?; + Ok(Vec::new()) + } + BodyLength::Chunked => { + let mut size_line = reader.read_line().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, committed).await?; + return read_response_trailers(reader).await; + } + let mut remaining = chunk_size; + while remaining > 0 { + let length = remaining.min(unit_limit); + let unit = reader.read_exact_vec(length).await?; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + remaining -= unit.len(); + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + committed, + unit_limit, + ) + .await?; + } + if reader.read_exact_vec(2).await?.as_slice() != b"\r\n" { + return Err(miette!("HTTP response chunk is missing its terminator")); + } + size_line = if let Ok(line) = + tokio::time::timeout(RESPONSE_UNIT_COALESCE_TIMEOUT, reader.read_line()).await + { + line? + } else { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + committed, + ) + .await?; + reader.read_line().await? + }; + } + } + BodyLength::None if server_wants_close || event_stream => loop { + let read = reader.read_some(unit_limit); + let next = if pending.is_empty() && event_stream { + read.await? + } else if pending.is_empty() { + match tokio::time::timeout(RELAY_EOF_IDLE_TIMEOUT, read).await { + Ok(result) => result?, + Err(_) => None, + } + } else if let Ok(result) = + tokio::time::timeout(RESPONSE_UNIT_COALESCE_TIMEOUT, read).await + { + result? + } else { + flush_normalized_response_bytes( + session, + client, + std::mem::take(&mut pending), + committed, + ) + .await?; + continue; + }; + let Some(unit) = next else { + flush_normalized_response_bytes(session, client, pending, committed).await?; + return Ok(Vec::new()); + }; + if let Some(guard) = generation_guard { + guard.ensure_current()?; + } + buffer_normalized_response_bytes( + session, + client, + &mut pending, + unit, + committed, + unit_limit, + ) + .await?; + }, + BodyLength::None => { + flush_normalized_response_bytes(session, client, pending, committed).await?; + Ok(Vec::new()) + } + } +} + +async fn buffer_normalized_response_bytes( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + pending: &mut Vec, + data: Vec, + committed: bool, + 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, committed).await?; + } + Ok(()) +} + +async fn flush_normalized_response_bytes( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + pending: Vec, + committed: bool, +) -> Result<()> { + if pending.is_empty() { + return Ok(()); + } + process_response_unit(session, client, pending, committed).await +} + +async fn process_response_unit( + session: &mut openshell_supervisor_middleware::HttpResponseSession, + client: &mut C, + unit: Vec, + committed: bool, +) -> Result<()> { + let output = session + .push_body(unit) + .await + .map_err(|error| miette!("HTTP response middleware failure: {error}"))?; + if committed { + for unit in output { + write_chunk(client, &unit).await?; + } + client.flush().await.into_diagnostic()?; + } else if !output.is_empty() { + return Err(miette!( + "whole-body response middleware released output before finalization" + )); + } + Ok(()) +} + +async fn read_response_trailers( + reader: &mut BufferedResponseReader<'_, R>, +) -> Result> { + let mut trailers = Vec::new(); + loop { + let line = reader.read_line().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"))?; + let name = name.to_ascii_lowercase(); + if matches!( + name.as_str(), + "connection" + | "content-length" + | "keep-alive" + | "proxy-authenticate" + | "proxy-authorization" + | "proxy-connection" + | "te" + | "trailer" + | "transfer-encoding" + | "upgrade" + ) { + continue; + } + 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_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`. @@ -3623,17 +4624,207 @@ 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, HttpResponseBodyMode, HttpResponseBodyResult, + HttpResponseBodyTransform, HttpResponseEvent, HttpResponseEventResult, + HttpResponsePreflightDecision, HttpResponsePreflightInspect, 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_decision, + }; 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, + Stream, + InvalidBodySequence, + InvalidWholeBodySequence, + } + + struct ResponseRelayService { + script: ResponseRelayScript, + } + + #[tonic::async_trait] + impl openshell_supervisor_middleware::InProcessMiddleware for ResponseRelayService { + async fn describe(&self) -> MiddlewareManifest { + MiddlewareManifest { + name: "test/response-relay".into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: SupervisorMiddlewareOperation::HttpResponse as i32, + phase: SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: 4096, + timeout: String::new(), + }], + expected_audience: String::new(), + } + } + + async fn validate_config( + &self, + _middleware_name: &str, + _config: &prost_types::Struct, + ) -> Result<()> { + Ok(()) + } + + async fn evaluate_http_request( + &self, + _request: openshell_supervisor_middleware::HttpRequestView<'_>, + ) -> Result { + Ok(HttpRequestResult { + decision: Decision::Allow as i32, + ..Default::default() + }) + } + + async fn open_http_response_pre_return( + &self, + mut requests: mpsc::Receiver, + ) -> std::result::Result< + openshell_supervisor_middleware::HttpResponseResultStream, + tonic::Status, + > { + let script = self.script; + let (sender, receiver) = mpsc::channel(4); + tokio::spawn(async move { + while let Some(event) = requests.recv().await { + let Some(event) = event.event else { + break; + }; + let result = match event { + http_response_event::Event::Preflight(_) => { + let (body_mode, header_mutations, declared_trailer_names) = match script + { + ResponseRelayScript::HeadersOnly => ( + HttpResponseBodyMode::HeadersOnly, + vec![write_header( + "cache-control", + "private", + ExistingHeaderAction::Overwrite, + )], + Vec::new(), + ), + ResponseRelayScript::WholeBody + | ResponseRelayScript::InvalidWholeBodySequence => { + (HttpResponseBodyMode::WholeBodyBytes, Vec::new(), Vec::new()) + } + ResponseRelayScript::Stream + | ResponseRelayScript::InvalidBodySequence => ( + HttpResponseBodyMode::StreamBytes, + Vec::new(), + vec!["digest".into()], + ), + }; + HttpResponseEventResult { + result: Some( + http_response_event_result::Result::PreflightDecision( + HttpResponsePreflightDecision { + decision: Some( + http_response_preflight_decision::Decision::Inspect( + HttpResponsePreflightInspect { + body_mode: body_mode as i32, + header_mutations, + declared_trailer_names, + ..Default::default() + }, + ), + ), + }, + ), + ), + } + } + http_response_event::Event::Body(body) => { + let Some(http_response_body_unit::Payload::Data(data)) = body.payload + else { + break; + }; + let replacement = match script { + ResponseRelayScript::WholeBody + | ResponseRelayScript::InvalidWholeBodySequence => { + [b"whole:".as_slice(), &data].concat() + } + ResponseRelayScript::Stream + | ResponseRelayScript::InvalidBodySequence => { + data.to_ascii_uppercase() + } + ResponseRelayScript::HeadersOnly => break, + }; + HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: if matches!( + script, + ResponseRelayScript::InvalidBodySequence + | ResponseRelayScript::InvalidWholeBodySequence + ) { + body.sequence + 1 + } else { + body.sequence + }, + decision: Some( + http_response_body_result::Decision::Transform( + HttpResponseBodyTransform { + replacement: Some( + http_response_body_transform::Replacement::Data( + replacement, + ), + ), + }, + ), + ), + ..Default::default() + }, + )), + } + } + http_response_event::Event::Trailers(_) => HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult { + trailer_mutations: if matches!( + script, + ResponseRelayScript::Stream + ) { + vec![write_header( + "digest", + "sha-256=:test:", + ExistingHeaderAction::Overwrite, + )] + } else { + Vec::new() + }, + ..Default::default() + }, + )), + }, + http_response_event::Event::BodyEnd(_) => continue, + 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, @@ -5503,6 +6694,273 @@ mod tests { assert!(!is_bodiless_response("POST", 201)); } + fn response_middleware_fixture( + script: ResponseRelayScript, + ) -> ( + openshell_supervisor_middleware::ChainRunner, + Vec, + ) { + let runner = + openshell_supervisor_middleware::ChainRunner::new(Arc::new(ResponseRelayService { + script, + })); + let chain = vec![openshell_supervisor_middleware::ChainEntry { + name: "response".into(), + implementation: "test/response-relay".into(), + order: 0, + config: prost_types::Struct::default(), + on_error: openshell_supervisor_middleware::OnError::FailClosed, + }]; + (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, + } + } + + async fn run_response_middleware_relay( + response: &'static [u8], + method: &str, + script: ResponseRelayScript, + ) -> (Result, Vec) { + let (runner, chain) = response_middleware_fixture(script); + let (mut upstream_read, mut upstream_write) = tokio::io::duplex(16 * 1024); + let (mut client_read, mut client_write) = tokio::io::duplex(16 * 1024); + tokio::spawn(async move { + upstream_write.write_all(response).await.unwrap(); + upstream_write.shutdown().await.unwrap(); + }); + let outcome = relay_response( + method, + &mut upstream_read, + &mut client_write, + RelayResponseOptions::default(), + Some(response_middleware_context(&runner, &chain, method)), + ) + .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_repairs_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("Transfer-Encoding: chunked\r\n"), + "{delivered}" + ); + assert!( + delivered.ends_with("5\r\nhello\r\n0\r\n\r\n"), + "{delivered}" + ); + } + + #[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_streams_normalized_chunk_payloads_and_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, digest\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: sha-256=:test:\r\n"), + "{delivered}" + ); + assert!(!delivered.contains("ext=yes"), "{delivered}"); + } + + #[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_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_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()), + }], + ); + 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}"); + } + #[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 @@ -5524,6 +6982,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5569,6 +7028,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5619,6 +7079,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5663,6 +7124,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5704,6 +7166,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5742,6 +7205,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5782,6 +7246,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5826,6 +7291,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5869,6 +7335,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5905,6 +7372,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5951,6 +7419,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await @@ -5998,6 +7467,7 @@ mod tests { &mut upstream_read, &mut client_write, RelayResponseOptions::default(), + None, ), ) .await diff --git a/crates/openshell-supervisor-network/src/proxy.rs b/crates/openshell-supervisor-network/src/proxy.rs index 177d640fd8..3d1ccdb2dd 100644 --- a/crates/openshell-supervisor-network/src/proxy.rs +++ b/crates/openshell-supervisor-network/src/proxy.rs @@ -1751,7 +1751,7 @@ async fn handle_tcp_connection( let target = parts.next().unwrap_or(""); if method != "CONNECT" { - return handle_forward_proxy( + return Box::pin(handle_forward_proxy( method, target, &buf[..], @@ -1768,7 +1768,7 @@ async fn handle_tcp_connection( dynamic_credentials, denial_tx.as_ref(), activity_tx.as_ref(), - ) + )) .await; } @@ -6630,7 +6630,7 @@ network_policies: tokio::time::timeout( std::time::Duration::from_secs(30), - handle_forward_proxy( + Box::pin(handle_forward_proxy( "GET", &target, request.as_bytes(), @@ -6647,7 +6647,7 @@ network_policies: None, None, None, - ), + )), ) .await .expect("denied preflight must complete without an upstream response") @@ -6763,7 +6763,7 @@ network_policies: let (mut proxy_connection, _) = proxy_listener.accept().await.unwrap(); let handler = tokio::spawn(async move { - handle_forward_proxy( + Box::pin(handle_forward_proxy( "GET", &target, request.as_bytes(), @@ -6780,7 +6780,7 @@ network_policies: None, None, None, - ) + )) .await }); let scenario = tokio::time::timeout(std::time::Duration::from_secs(60), async { From ce876dc2288fb03831baa6b7d0eb9a9fc880b075 Mon Sep 17 00:00:00 2001 From: Piotr Mlocek Date: Mon, 31 Aug 2026 20:13:30 -0700 Subject: [PATCH 3/6] fix(middleware): harden response delivery semantics Signed-off-by: Piotr Mlocek --- .../src/l7/rest.rs | 222 +++++++++++++++++- 1 file changed, 216 insertions(+), 6 deletions(-) diff --git a/crates/openshell-supervisor-network/src/l7/rest.rs b/crates/openshell-supervisor-network/src/l7/rest.rs index 49c222bdba..0d0470df30 100644 --- a/crates/openshell-supervisor-network/src/l7/rest.rs +++ b/crates/openshell-supervisor-network/src/l7/rest.rs @@ -3531,8 +3531,18 @@ where server_wants_close, &declared_trailers, ); - client.write_all(&head).await.into_diagnostic()?; - client.flush().await.into_diagnostic()?; + if let Err(error) = client.write_all(&head).await { + session + .end(openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect) + .await; + return Err(error).into_diagnostic(); + } + if let Err(error) = client.flush().await { + session + .end(openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect) + .await; + return Err(error).into_diagnostic(); + } true }; @@ -3557,6 +3567,11 @@ where .is_some_and(PolicyGenerationGuard::is_stale) { openshell_core::proto::HttpResponseSessionEndReason::PolicyReload + } else if error + .to_string() + .starts_with("HTTP response client write failed:") + { + openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect } else if error .to_string() .starts_with("HTTP response middleware failure:") @@ -4198,9 +4213,14 @@ async fn process_response_unit( .map_err(|error| miette!("HTTP response middleware failure: {error}"))?; if committed { for unit in output { - write_chunk(client, &unit).await?; + write_chunk(client, &unit) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; } - client.flush().await.into_diagnostic()?; + client + .flush() + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; } else if !output.is_empty() { return Err(miette!( "whole-body response middleware released output before finalization" @@ -6699,6 +6719,19 @@ mod tests { ) -> ( 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 { @@ -6709,7 +6742,7 @@ mod tests { implementation: "test/response-relay".into(), order: 0, config: prost_types::Struct::default(), - on_error: openshell_supervisor_middleware::OnError::FailClosed, + on_error, }]; (runner, chain) } @@ -6745,7 +6778,22 @@ mod tests { method: &str, script: ResponseRelayScript, ) -> (Result, Vec) { - let (runner, chain) = response_middleware_fixture(script); + run_response_middleware_relay_with_error( + response, + method, + script, + openshell_supervisor_middleware::OnError::FailClosed, + ) + .await + } + + async fn run_response_middleware_relay_with_error( + response: &'static [u8], + method: &str, + script: ResponseRelayScript, + on_error: openshell_supervisor_middleware::OnError, + ) -> (Result, Vec) { + let (runner, chain) = response_middleware_fixture_with_error(script, on_error); let (mut upstream_read, mut upstream_write) = tokio::io::duplex(16 * 1024); let (mut client_read, mut client_write) = tokio::io::duplex(16 * 1024); tokio::spawn(async move { @@ -6845,6 +6893,55 @@ mod tests { 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( @@ -6900,6 +6997,88 @@ mod tests { assert!(!delivered.contains("502 Bad Gateway"), "{delivered}"); } + #[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( @@ -9419,4 +9598,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(), + }] + ); + } } From 7040e192c87bb8e2fb31a2d78ce75abe24f444a0 Mon Sep 17 00:00:00 2001 From: Piotr Mlocek Date: Mon, 31 Aug 2026 23:07:44 -0700 Subject: [PATCH 4/6] feat(middleware): extend content guard response example Signed-off-by: Piotr Mlocek --- .../Cargo.lock | 263 ++++++++++++ .../Cargo.toml | 3 + .../README.md | 35 +- .../policy.yaml | 19 + .../src/main.rs | 397 +++++++++++++++++- .../upstream.py | 51 +++ 6 files changed, 754 insertions(+), 14 deletions(-) create mode 100644 examples/supervisor-middleware-content-guard/upstream.py diff --git a/examples/supervisor-middleware-content-guard/Cargo.lock b/examples/supervisor-middleware-content-guard/Cargo.lock index f31d5be9b5..ad79004dbe 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.lock +++ b/examples/supervisor-middleware-content-guard/Cargo.lock @@ -206,6 +206,17 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527" +[[package]] +name = "chacha20" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "65c35e4b699c7e15ccbe7ee35c005e4fc0a278d22238a2857e6ce2dadeda1b06" +dependencies = [ + "cfg-if", + "cpufeatures", + "rand_core", +] + [[package]] name = "clap" version = "4.6.1" @@ -252,6 +263,16 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" +[[package]] +name = "combine" +version = "4.6.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfc320937d09e6de266b31b9afb480f197d7a861be86be7cb2ea7e5d1bfffc5e" +dependencies = [ + "bytes", + "memchr", +] + [[package]] name = "core-foundation" version = "0.10.1" @@ -268,6 +289,27 @@ version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" +[[package]] +name = "cpufeatures" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ca28b0ae3115b884660db4118d803791fd6756b6e88f39c0f3f7859060d7566" +dependencies = [ + "libc", +] + +[[package]] +name = "critical-section" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "790eea4361631c5e7d22598ecd5723ff611904e3344ce8720784c93e3d83d40b" + +[[package]] +name = "data-encoding" +version = "2.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4583a4551df46e2792f82ceeac45e850d2e2d5debba0b91f102385cda5b11f06" + [[package]] name = "displaydoc" version = "0.2.6" @@ -399,6 +441,7 @@ dependencies = [ "cfg-if", "libc", "r-efi", + "rand_core", ] [[package]] @@ -453,6 +496,25 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hickory-proto" +version = "0.26.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bab31817bfb44672a252e97fe81cd0c18d1b2cf892108922f6818820df8c643" +dependencies = [ + "data-encoding", + "idna", + "ipnet", + "jni", + "once_cell", + "rand", + "ring", + "thiserror", + "tinyvec", + "tracing", + "url", +] + [[package]] name = "http" version = "1.4.2" @@ -699,12 +761,71 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn", +] + [[package]] name = "libc" version = "0.2.186" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" +[[package]] +name = "libyml" +version = "0.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3302702afa434ffa30847a83305f0a69d6abd74293b6554c18ec85c7ef30c980" +dependencies = [ + "anyhow", + "version_check", +] + [[package]] name = "linux-raw-sys" version = "0.12.1" @@ -832,6 +953,10 @@ name = "once_cell" version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +dependencies = [ + "critical-section", + "portable-atomic", +] [[package]] name = "once_cell_polyfill" @@ -878,12 +1003,26 @@ dependencies = [ "tower", ] +[[package]] +name = "openshell-policy" +version = "0.0.0" +dependencies = [ + "hickory-proto", + "miette", + "openshell-core", + "prost-types", + "serde", + "serde_json", + "serde_yml", +] + [[package]] name = "openshell-supervisor-middleware-content-guard" version = "0.0.0" dependencies = [ "clap", "openshell-core", + "openshell-policy", "prost-types", "tokio", "tokio-stream", @@ -968,6 +1107,12 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" +[[package]] +name = "portable-atomic" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05c8b63e8d9609db387f0324918f81d68fe27748f084ef092fb35954d0539a85" + [[package]] name = "potential_utf" version = "0.1.5" @@ -1148,6 +1293,23 @@ version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "rand" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +dependencies = [ + "chacha20", + "getrandom 0.4.3", + "rand_core", +] + +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + [[package]] name = "redox_syscall" version = "0.5.18" @@ -1206,6 +1368,15 @@ version = "0.1.27" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b50b8869d9fc858ce7266cce0194bd74df58b9d0e3f6df3a9fc8eb470d95c09d" +[[package]] +name = "rustc_version" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" +dependencies = [ + "semver", +] + [[package]] name = "rustix" version = "1.1.4" @@ -1266,6 +1437,21 @@ dependencies = [ "untrusted", ] +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + [[package]] name = "schannel" version = "0.1.29" @@ -1304,6 +1490,12 @@ dependencies = [ "libc", ] +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + [[package]] name = "serde" version = "1.0.228" @@ -1347,6 +1539,21 @@ dependencies = [ "zmij", ] +[[package]] +name = "serde_yml" +version = "0.0.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59e2dd588bf1597a252c3b920e0143eb99b0f76e4e082f4c92ce34fbc9e71ddd" +dependencies = [ + "indexmap", + "itoa", + "libyml", + "memchr", + "ryu", + "serde", + "version_check", +] + [[package]] name = "shlex" version = "2.0.1" @@ -1363,6 +1570,22 @@ dependencies = [ "libc", ] +[[package]] +name = "simd_cesu8" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11031e251abf8611c80f460e19dbdeb54a66db918e49c65a7065b46ac7aec520" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "slab" version = "0.4.12" @@ -1515,6 +1738,21 @@ dependencies = [ "zerovec", ] +[[package]] +name = "tinyvec" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb4ebadaa0af04fab11ae01eb5f9fdb5f9c5b875506e210e71c07873528baa7f" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + [[package]] name = "tokio" version = "1.52.3" @@ -1775,6 +2013,22 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + [[package]] name = "want" version = "0.3.1" @@ -1790,6 +2044,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "windows-link" version = "0.2.1" diff --git a/examples/supervisor-middleware-content-guard/Cargo.toml b/examples/supervisor-middleware-content-guard/Cargo.toml index ef0d1f47c3..9a9f98381a 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.toml +++ b/examples/supervisor-middleware-content-guard/Cargo.toml @@ -20,6 +20,9 @@ tokio = { version = "1.43", features = ["macros", "rt-multi-thread"] } tokio-stream = "0.1" tonic = { version = "0.14", features = ["transport"] } +[dev-dependencies] +openshell-policy = { path = "../../crates/openshell-policy" } + [[bin]] name = "supervisor-middleware-content-guard" path = "src/main.rs" diff --git a/examples/supervisor-middleware-content-guard/README.md b/examples/supervisor-middleware-content-guard/README.md index 53eda94210..76a69e7ded 100644 --- a/examples/supervisor-middleware-content-guard/README.md +++ b/examples/supervisor-middleware-content-guard/README.md @@ -8,14 +8,14 @@ SPDX-License-Identifier: Apache-2.0 > [!WARNING] > Supervisor middleware is a research preview. Its policy and service contracts may change without compatibility guarantees. Use it only to prototype and evaluate middleware integrations. -This example implements an operator-run supervisor middleware service. It scans UTF-8 HTTP request bodies and complete client-to-upstream WebSocket text messages for configured literal strings, then either replaces every match or denies the request or message. Findings report only aggregate counts and never include configured terms or inspected content. +This example implements one operator-run service with HTTP request, HTTP response, and client WebSocket bindings. Request bodies and WebSocket text messages use the content-guard behavior. The response binding demonstrates headers-only, whole-body, and normalized streaming processing. > [!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. Request and WebSocket handling supports only case-sensitive literal matches in UTF-8 text. The response paths demonstrate framing and body modes with fixed transformations; they do not implement content detection. The example does not address content encodings, application-level normalization, binary WebSocket messages, or adversarial inputs. ## Prerequisites -Install `cargo`, `curl`, `jq`, and `openssl` on the host before running the smoke script. +Install `cargo`, `curl`, `jq`, `openssl`, and Python 3 on the host before running the example. ## Run the smoke example @@ -54,6 +54,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 +85,32 @@ curl -sS https://httpbin.org/anything \ The echoed JSON body contains `[FILTERED]` instead of the configured term. +## HTTP response behavior + +The same service advertises `HTTP_RESPONSE/PRE_RETURN`. Start the included raw HTTP upstream in another terminal: + +```shell +python3 examples/supervisor-middleware-content-guard/upstream.py +``` + +The example policy allows these requests from the sandbox: + +```shell +curl -i http://host.openshell.internal:18081/headers-only +curl -i http://host.openshell.internal:18081/whole-body +curl -i --raw http://host.openshell.internal:18081/stream +``` + +The paths exercise each response mode: + +| Path | Mode | Result | +| --- | --- | --- | +| `/headers-only` | `HEADERS_ONLY` | Adds `x-example-response-mode` and passes the body through. | +| `/whole-body` | `WHOLE_BODY_BYTES` | Prefixes the normalized complete body with `[whole]` plus a space. | +| `/stream` | `STREAM_BYTES` | Uppercases each normalized unit and writes the declared `x-example-body-bytes` trailer. | + +The upstream uses content-length, chunked, and close-delimited responses. OpenShell removes those transport boundaries before invoking middleware and repairs downstream framing after a transformation. Paths outside this table return `SKIP` with `path_not_selected`. + ## 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 +134,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..fb7a6e0d59 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,21 @@ network_policies: path: /anything binaries: - path: /usr/bin/curl + response-framing-demo: + name: Response framing demo + endpoints: + - host: host.openshell.internal + port: 18081 + protocol: rest + rules: + - allow: + method: GET + path: /headers-only + - allow: + method: GET + path: /whole-body + - allow: + method: GET + path: /stream + binaries: + - path: /usr/bin/curl diff --git a/examples/supervisor-middleware-content-guard/src/main.rs b/examples/supervisor-middleware-content-guard/src/main.rs index 8d714264e7..95f40fad17 100644 --- a/examples/supervisor-middleware-content-guard/src/main.rs +++ b/examples/supervisor-middleware-content-guard/src/main.rs @@ -6,20 +6,31 @@ use std::net::SocketAddr; use std::ops::Range; use clap::Parser; -use openshell_core::middleware::WebSocketResponseStream; +use openshell_core::middleware::{HttpResponseResultStream, WebSocketResponseStream}; +use openshell_core::proto::middleware::v1::http_response_pre_return_server::{ + HttpResponsePreReturn, HttpResponsePreReturnServer, +}; use openshell_core::proto::middleware::v1::supervisor_middleware_server::{ SupervisorMiddleware, SupervisorMiddlewareServer, }; use openshell_core::proto::{ - Decision, Finding, HttpRequestEvaluation, HttpRequestResult, MiddlewareBinding, - MiddlewareManifest, SupervisorMiddlewareOperation, SupervisorMiddlewarePhase, - ValidateConfigRequest, ValidateConfigResponse, WebSocketMessage, WebSocketMessageResult, - WebSocketPreflightAction, WebSocketPreflightDecision, WebSocketSessionEvent, - WebSocketSessionEventResult, web_socket_message, web_socket_message_result, - web_socket_session_event, web_socket_session_event_result, + Decision, ExistingHeaderAction, Finding, HeaderMutation, HttpRequestEvaluation, + HttpRequestResult, HttpResponseBodyMode, HttpResponseBodyResult, HttpResponseBodyTransform, + HttpResponseEvent, HttpResponseEventResult, HttpResponsePreflightDecision, + HttpResponsePreflightInspect, HttpResponsePreflightSkip, HttpResponseTrailersResult, + MiddlewareBinding, MiddlewareManifest, SupervisorMiddlewareOperation, + SupervisorMiddlewarePhase, ValidateConfigRequest, ValidateConfigResponse, WebSocketMessage, + WebSocketMessageResult, WebSocketPreflightAction, WebSocketPreflightDecision, + WebSocketSessionEvent, WebSocketSessionEventResult, WriteHeader, header_mutation, + http_response_body_result, http_response_body_transform, http_response_body_unit, + http_response_event, http_response_event_result, http_response_preflight_decision, + web_socket_message, web_socket_message_result, web_socket_session_event, + web_socket_session_event_result, }; use prost_types::Struct; use prost_types::value::Kind; +use tokio::sync::mpsc; +use tokio_stream::wrappers::ReceiverStream; use tokio_stream::{Stream, StreamExt}; use tonic::transport::Server; use tonic::{Request, Response, Status}; @@ -238,6 +249,12 @@ impl SupervisorMiddleware for ContentGuard { max_payload_bytes: MAX_PAYLOAD_BYTES, timeout: String::new(), }, + MiddlewareBinding { + operation: SupervisorMiddlewareOperation::HttpResponse as i32, + phase: SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: MAX_PAYLOAD_BYTES, + timeout: String::new(), + }, ], expected_audience: String::new(), })) @@ -282,6 +299,240 @@ impl SupervisorMiddleware for ContentGuard { } } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ResponseMode { + HeadersOnly, + WholeBody, + Stream, +} + +#[derive(Debug, Default)] +struct ResponseSessionState { + selected: Option, + next_sequence: u64, + input_bytes: u64, +} + +impl ResponseSessionState { + fn preflight( + &mut self, + preflight: openshell_core::proto::HttpResponsePreflight, + ) -> Result { + if self.selected.is_some() { + return Err(Status::failed_precondition("duplicate response preflight")); + } + GuardConfig::parse(preflight.config.as_ref()).map_err(Status::invalid_argument)?; + let path = preflight + .target + .as_ref() + .map(|target| target.path.as_str()) + .unwrap_or_default(); + let Some(selected) = response_mode_for_path(path) else { + return Ok(response_preflight_skip()); + }; + self.selected = Some(selected); + self.next_sequence = 1; + Ok(response_preflight_inspect(selected)) + } + + fn body( + &mut self, + body: openshell_core::proto::HttpResponseBodyUnit, + ) -> Result { + let selected = self + .selected + .ok_or_else(|| Status::failed_precondition("body arrived before preflight"))?; + if selected == ResponseMode::HeadersOnly { + return Err(Status::failed_precondition( + "headers-only sessions do not receive body events", + )); + } + if body.sequence != self.next_sequence { + return Err(Status::invalid_argument(format!( + "expected body sequence {}, received {}", + self.next_sequence, body.sequence + ))); + } + self.next_sequence = self.next_sequence.saturating_add(1); + let Some(http_response_body_unit::Payload::Data(data)) = body.payload else { + return Err(Status::invalid_argument("body data is required")); + }; + self.input_bytes = self.input_bytes.saturating_add(data.len() as u64); + let replacement = match selected { + ResponseMode::WholeBody => [b"[whole] ".as_slice(), &data].concat(), + ResponseMode::Stream => data.to_ascii_uppercase(), + ResponseMode::HeadersOnly => unreachable!(), + }; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::BodyResult( + HttpResponseBodyResult { + sequence: body.sequence, + decision: Some(http_response_body_result::Decision::Transform( + HttpResponseBodyTransform { + replacement: Some(http_response_body_transform::Replacement::Data( + replacement, + )), + }, + )), + ..Default::default() + }, + )), + }) + } + + fn body_end(&self, final_sequence: u64) -> Result<(), Status> { + let expected = self.next_sequence.saturating_sub(1); + if final_sequence != expected { + return Err(Status::invalid_argument(format!( + "expected final body sequence {expected}, received {final_sequence}" + ))); + } + Ok(()) + } + + fn trailers(&self) -> Result { + let selected = self + .selected + .ok_or_else(|| Status::failed_precondition("trailers arrived before preflight"))?; + if selected == ResponseMode::HeadersOnly { + return Err(Status::failed_precondition( + "headers-only sessions do not receive trailers", + )); + } + let trailer_mutations = if selected == ResponseMode::Stream { + vec![write_header( + "x-example-body-bytes", + &self.input_bytes.to_string(), + )] + } else { + Vec::new() + }; + Ok(HttpResponseEventResult { + result: Some(http_response_event_result::Result::TrailersResult( + HttpResponseTrailersResult { + trailer_mutations, + ..Default::default() + }, + )), + }) + } +} + +fn response_mode_for_path(path: &str) -> Option { + match path { + "/headers-only" => Some(ResponseMode::HeadersOnly), + "/whole-body" => Some(ResponseMode::WholeBody), + "/stream" => Some(ResponseMode::Stream), + _ => None, + } +} + +fn response_preflight_skip() -> HttpResponseEventResult { + HttpResponseEventResult { + result: Some(http_response_event_result::Result::PreflightDecision( + HttpResponsePreflightDecision { + decision: Some(http_response_preflight_decision::Decision::Skip( + HttpResponsePreflightSkip { + reason: "path is outside the response example".into(), + reason_code: "path_not_selected".into(), + ..Default::default() + }, + )), + }, + )), + } +} + +fn response_preflight_inspect(selected: ResponseMode) -> HttpResponseEventResult { + let (body_mode, declared_trailer_names) = match selected { + ResponseMode::HeadersOnly => (HttpResponseBodyMode::HeadersOnly, Vec::new()), + ResponseMode::WholeBody => (HttpResponseBodyMode::WholeBodyBytes, Vec::new()), + ResponseMode::Stream => ( + HttpResponseBodyMode::StreamBytes, + vec!["x-example-body-bytes".into()], + ), + }; + let mode_name = match selected { + ResponseMode::HeadersOnly => "headers-only", + ResponseMode::WholeBody => "whole-body", + ResponseMode::Stream => "stream", + }; + HttpResponseEventResult { + result: Some(http_response_event_result::Result::PreflightDecision( + HttpResponsePreflightDecision { + decision: Some(http_response_preflight_decision::Decision::Inspect( + HttpResponsePreflightInspect { + body_mode: body_mode as i32, + header_mutations: vec![write_header("x-example-response-mode", mode_name)], + declared_trailer_names, + ..Default::default() + }, + )), + }, + )), + } +} + +fn write_header(name: &str, value: &str) -> HeaderMutation { + HeaderMutation { + operation: Some(header_mutation::Operation::Write(WriteHeader { + name: name.into(), + value: value.into(), + on_existing: ExistingHeaderAction::Overwrite as i32, + })), + } +} + +#[tonic::async_trait] +impl HttpResponsePreReturn for ContentGuard { + type EvaluateStream = HttpResponseResultStream; + + async fn evaluate( + &self, + request: Request>, + ) -> Result, Status> { + let mut events = request.into_inner(); + let (sender, receiver) = mpsc::channel(4); + tokio::spawn(async move { + let mut state = ResponseSessionState::default(); + while let Some(event) = events.next().await { + let event = match event { + Ok(event) => event, + Err(error) => { + let _ = sender.send(Err(error)).await; + break; + } + }; + let result = match event.event { + Some(http_response_event::Event::Preflight(preflight)) => { + state.preflight(preflight).map(Some) + } + Some(http_response_event::Event::Body(body)) => state.body(body).map(Some), + Some(http_response_event::Event::BodyEnd(body_end)) => { + state.body_end(body_end.final_sequence).map(|()| None) + } + Some(http_response_event::Event::Trailers(_)) => state.trailers().map(Some), + Some(http_response_event::Event::SessionEnd(_)) => break, + None => Err(Status::invalid_argument("response event is required")), + }; + match result { + Ok(Some(result)) => { + if sender.send(Ok(result)).await.is_err() { + break; + } + } + Ok(None) => {} + Err(error) => { + let _ = sender.send(Err(error)).await; + break; + } + } + } + }); + Ok(Response::new(Box::pin(ReceiverStream::new(receiver)))) + } +} + fn validate_phase(phase: i32) -> Result<(), String> { if phase != PHASE as i32 { return Err(format!("unsupported phase '{phase}'")); @@ -470,6 +721,7 @@ async fn main() -> Result<(), Box> { println!("serving {MANIFEST_NAME} on http://{}", cli.bind); Server::builder() .add_service(SupervisorMiddlewareServer::new(ContentGuard)) + .add_service(HttpResponsePreReturnServer::new(ContentGuard)) .serve(cli.bind) .await?; Ok(()) @@ -478,7 +730,10 @@ async fn main() -> Result<(), Box> { #[cfg(test)] mod tests { use super::*; - use openshell_core::proto::{MiddlewareSessionEnd, WebSocketPreflight, WebSocketSessionStart}; + use openshell_core::proto::{ + HttpRequestTarget, HttpResponseBodyUnit, HttpResponsePreflight, MiddlewareSessionEnd, + WebSocketPreflight, WebSocketSessionStart, + }; use prost_types::{ListValue, Value}; use std::collections::BTreeMap; @@ -511,13 +766,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 +783,121 @@ 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 + ); + assert_eq!(manifest.bindings[2].max_payload_bytes, MAX_PAYLOAD_BYTES); + } + + fn response_preflight(path: &str) -> HttpResponsePreflight { + HttpResponsePreflight { + target: Some(HttpRequestTarget { + path: path.into(), + ..Default::default() + }), + config: Some(config("redact", &["prototype-secret"], Some("[FILTERED]"))), + ..Default::default() + } + } + + #[test] + fn paths_select_all_three_response_modes() { + for (path, expected) in [ + ("/headers-only", HttpResponseBodyMode::HeadersOnly), + ("/whole-body", HttpResponseBodyMode::WholeBodyBytes), + ("/stream", HttpResponseBodyMode::StreamBytes), + ] { + let mut state = ResponseSessionState::default(); + let result = state.preflight(response_preflight(path)).unwrap(); + let Some(http_response_event_result::Result::PreflightDecision(decision)) = + result.result + else { + panic!("expected preflight decision"); + }; + let Some(http_response_preflight_decision::Decision::Inspect(inspect)) = + decision.decision + else { + panic!("expected inspect decision"); + }; + assert_eq!(inspect.body_mode, expected as i32); + } + } + + #[test] + fn whole_body_and_stream_transform_differently() { + for (path, expected) in [ + ("/whole-body", b"[whole] hello".as_slice()), + ("/stream", b"HELLO".as_slice()), + ] { + let mut state = ResponseSessionState::default(); + state.preflight(response_preflight(path)).unwrap(); + let result = state + .body(HttpResponseBodyUnit { + sequence: 1, + payload: Some(http_response_body_unit::Payload::Data(b"hello".to_vec())), + }) + .unwrap(); + let Some(http_response_event_result::Result::BodyResult(body)) = result.result else { + panic!("expected body result"); + }; + let Some(http_response_body_result::Decision::Transform(transform)) = body.decision + else { + panic!("expected body transform"); + }; + let Some(http_response_body_transform::Replacement::Data(data)) = transform.replacement + else { + panic!("expected replacement data"); + }; + assert_eq!(data, expected); + } + } + + #[test] + fn stream_declares_and_writes_byte_count_trailer() { + let mut state = ResponseSessionState::default(); + let preflight = state.preflight(response_preflight("/stream")).unwrap(); + let Some(http_response_event_result::Result::PreflightDecision(decision)) = + preflight.result + else { + panic!("expected preflight decision"); + }; + let Some(http_response_preflight_decision::Decision::Inspect(inspect)) = decision.decision + else { + panic!("expected inspect decision"); + }; + assert_eq!(inspect.declared_trailer_names, ["x-example-body-bytes"]); + state + .body(HttpResponseBodyUnit { + sequence: 1, + payload: Some(http_response_body_unit::Payload::Data(b"hello".to_vec())), + }) + .unwrap(); + state.body_end(1).unwrap(); + let trailers = state.trailers().unwrap(); + let Some(http_response_event_result::Result::TrailersResult(trailers)) = trailers.result + else { + panic!("expected trailer result"); + }; + assert_eq!(trailers.trailer_mutations.len(), 1); + } + + #[test] + fn paths_outside_the_response_demo_are_skipped() { + let mut state = ResponseSessionState::default(); + let result = state.preflight(response_preflight("/outside")).unwrap(); + let Some(http_response_event_result::Result::PreflightDecision(decision)) = result.result + else { + panic!("expected preflight decision"); + }; + let Some(http_response_preflight_decision::Decision::Skip(skip)) = decision.decision else { + panic!("expected skip decision"); + }; + assert_eq!(skip.reason_code, "path_not_selected"); } #[tokio::test] @@ -737,4 +1107,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..566d7086f3 --- /dev/null +++ b/examples/supervisor-middleware-content-guard/upstream.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import socketserver + + +class Handler(socketserver.BaseRequestHandler): + def handle(self): + request = b"" + while b"\r\n\r\n" not in request: + block = self.request.recv(4096) + if not block: + return + request += block + path = request.split(b" ", 2)[1] + if path == b"/headers-only": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Content-Length: 12\r\n\r\n" + b"headers-only" + ) + elif path == b"/whole-body": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Transfer-Encoding: chunked\r\n\r\n" + b"6\r\nwhole \r\n4\r\nbody\r\n0\r\n\r\n" + ) + elif path == b"/stream": + response = ( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/plain\r\n" + b"Connection: close\r\n\r\n" + b"stream body" + ) + else: + response = b"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n" + self.request.sendall(response) + + +class DemoServer(socketserver.ThreadingTCPServer): + allow_reuse_address = True + + +with DemoServer(("0.0.0.0", 18081), Handler) as server: + print("response framing demo upstream listening on 0.0.0.0:18081", flush=True) + try: + server.serve_forever() + except KeyboardInterrupt: + pass From c439a3fb786452586a8de4e0cda366366a02162a Mon Sep 17 00:00:00 2001 From: Piotr Mlocek Date: Mon, 31 Aug 2026 23:47:13 -0700 Subject: [PATCH 5/6] fix(middleware): address response relay review feedback Signed-off-by: Piotr Mlocek --- .../src/response.rs | 61 +++- .../src/l7/rest.rs | 325 +++++++++++++++--- 2 files changed, 328 insertions(+), 58 deletions(-) diff --git a/crates/openshell-supervisor-middleware/src/response.rs b/crates/openshell-supervisor-middleware/src/response.rs index a2fe8a7e8e..052516d4c1 100644 --- a/crates/openshell-supervisor-middleware/src/response.rs +++ b/crates/openshell-supervisor-middleware/src/response.rs @@ -72,6 +72,7 @@ pub struct HttpResponseInvocation { pub failed: bool, pub stage_disabled: bool, pub reason_code: Option, + pub failure_category: Option, } pub struct HttpResponsePreflightOutcome { @@ -211,12 +212,18 @@ impl HttpResponseSession { })?; let deadline = Instant::now() + MAX_MIDDLEWARE_CHAIN_TIMEOUT; let output = self.process_units_from(0, vec![data], deadline).await?; - if self.defer_output_until_finish { + if !self.defer_output_until_finish { + return Ok(output); + } + if self.requires_whole_body() { self.deferred_output.extend(output); - Ok(Vec::new()) - } else { - Ok(output) + return Ok(Vec::new()); } + + self.defer_output_until_finish = false; + let mut released = std::mem::take(&mut self.deferred_output); + released.extend(output); + Ok(released) } /// Finalize every body stage, process normalized trailers, and end streams. @@ -595,6 +602,7 @@ impl HttpResponseSession { failed: true, stage_disabled: true, reason_code: None, + failure_category: Some(response_failure_category(reason).into()), }); stage .end(HttpResponseSessionEndReason::MiddlewareFailure) @@ -795,6 +803,7 @@ impl ChainRunner { failed: false, stage_disabled: false, reason_code: (!skip.reason_code.is_empty()).then_some(skip.reason_code), + failure_category: None, }); let mut skipped = HttpResponseStage { entry, @@ -898,6 +907,7 @@ impl ChainRunner { failed: false, stage_disabled: false, reason_code: None, + failure_category: None, }); stages.push(HttpResponseStage { entry, @@ -1338,6 +1348,7 @@ fn body_invocation( failed: false, stage_disabled: false, reason_code: None, + failure_category: None, } } @@ -1412,10 +1423,37 @@ fn collect_preflight_failure( failed: true, stage_disabled: true, reason_code: None, + failure_category: Some(response_failure_category(reason).into()), }); fail_closed.then(|| format!("middleware_failed: {reason}")) } +fn response_failure_category(reason: &str) -> &'static str { + if reason == "middleware_session_capacity_exhausted" { + "session_capacity" + } else if reason.contains("over_capacity") { + "payload_capacity" + } else if reason.contains("timeout") { + "timeout" + } else if reason.contains("stream_closed") + || reason.contains("stream closed") + || reason.contains("transport") + || reason.contains("unavailable") + { + "transport" + } else if matches!( + reason, + "bodyless_response" + | "partial_response" + | "content_coding_not_identity" + | "cache_control_no_transform" + ) { + "response_not_inspectable" + } else { + "invalid_result" + } +} + fn empty_preflight_outcome(headers: Vec) -> HttpResponsePreflightOutcome { HttpResponsePreflightOutcome { allowed: true, @@ -2057,9 +2095,20 @@ mod tests { let pushed = session.push_body(original.clone()).await; assert_eq!(pushed.is_ok(), allowed); if allowed { - assert!(pushed.unwrap().is_empty()); + 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_eq!(finish.body_units, vec![original]); + assert!(finish.body_units.is_empty()); } } } diff --git a/crates/openshell-supervisor-network/src/l7/rest.rs b/crates/openshell-supervisor-network/src/l7/rest.rs index 0d0470df30..d6b1e5f88a 100644 --- a/crates/openshell-supervisor-network/src/l7/rest.rs +++ b/crates/openshell-supervisor-network/src/l7/rest.rs @@ -3471,6 +3471,7 @@ where }; 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 { @@ -3515,23 +3516,33 @@ where let whole_body = session.requires_whole_body(); let unit_limit = session.stream_unit_limit().max(1); - let committed = if whole_body { - false - } else { - let mut declared_trailers = upstream_declared_trailers; - for name in &preflight.declared_trailer_names { - if !declared_trailers.contains(name) { - declared_trailers.push(name.clone()); - } + let chunked_output = supports_chunked_response; + let close_delimited_output = !supports_chunked_response; + let mut declared_trailers = upstream_declared_trailers; + for name in &preflight.declared_trailer_names { + if !declared_trailers.contains(name) { + declared_trailers.push(name.clone()); } - let head = serialize_response_head( - &status_line, - &preflight.headers, - ResponseFraming::Chunked, - server_wants_close, - &declared_trailers, - ); - if let Err(error) = client.write_all(&head).await { + } + let downstream_trailers = if chunked_output { + declared_trailers.as_slice() + } else { + &[] + }; + let streaming_head = serialize_response_head( + &status_line, + &preflight.headers, + if chunked_output { + ResponseFraming::Chunked + } else { + ResponseFraming::Preserve(BodyLength::None) + }, + server_wants_close || close_delimited_output, + downstream_trailers, + ); + let mut committed = !whole_body; + if committed { + if let Err(error) = client.write_all(&streaming_head).await { session .end(openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect) .await; @@ -3543,8 +3554,7 @@ where .await; return Err(error).into_diagnostic(); } - true - }; + } let mut reader = BufferedResponseReader::new(upstream, &buffered[header_end..]); let body_result = relay_normalized_response_body( @@ -3554,7 +3564,9 @@ where body_length, server_wants_close, event_stream, - committed, + &mut committed, + chunked_output, + &streaming_head, unit_limit, middleware.generation_guard, ) @@ -3655,7 +3667,7 @@ where &finish.invocations, ); - if whole_body { + if whole_body && !committed { let mut headers = preflight.headers; if finish.strip_stale_integrity_headers { strip_response_integrity_headers(&mut headers); @@ -3665,16 +3677,20 @@ where .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() { + let framing = if finish.trailers.is_empty() || !supports_chunked_response { ResponseFraming::ContentLength(output_length as u64) } else { ResponseFraming::Chunked }; - let trailer_names: Vec = finish - .trailers - .iter() - .map(|header| header.name.clone()) - .collect(); + 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, @@ -3695,13 +3711,21 @@ where } } else { for unit in &finish.body_units { - write_chunk(client, unit).await?; + if chunked_output { + write_chunk(client, unit).await?; + } else { + client.write_all(unit).await.into_diagnostic()?; + } + } + if chunked_output { + write_response_trailers(client, &finish.trailers).await?; } - write_response_trailers(client, &finish.trailers).await?; } client.flush().await.into_diagnostic()?; Ok(Some( - if server_wants_close || matches!(body_length, BodyLength::None) { + if (committed && close_delimited_output) + || (matches!(body_length, BodyLength::None) && (server_wants_close || event_stream)) + { RelayOutcome::Consumed } else { RelayOutcome::Reusable @@ -3720,6 +3744,13 @@ fn emit_http_response_middleware_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); + } + } } fn http_response_middleware_invocation_events( @@ -3792,6 +3823,51 @@ fn http_response_middleware_invocation_events( .collect() } +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, @@ -4039,7 +4115,9 @@ async fn relay_normalized_response_body( body_length: BodyLength, server_wants_close: bool, event_stream: bool, - committed: bool, + committed: &mut bool, + chunked_output: bool, + commit_head: &[u8], unit_limit: usize, generation_guard: Option<&PolicyGenerationGuard>, ) -> Result> @@ -4048,6 +4126,11 @@ where C: AsyncWrite + Unpin, { let mut pending = Vec::with_capacity(unit_limit); + let mut framing = ResponseOutputState { + committed, + chunked: chunked_output, + commit_head, + }; match body_length { BodyLength::ContentLength(mut remaining) => { while remaining > 0 { @@ -4064,12 +4147,12 @@ where client, &mut pending, unit, - committed, + &mut framing, unit_limit, ) .await?; } - flush_normalized_response_bytes(session, client, pending, committed).await?; + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; Ok(Vec::new()) } BodyLength::Chunked => { @@ -4085,7 +4168,7 @@ where 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, committed).await?; + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; return read_response_trailers(reader).await; } let mut remaining = chunk_size; @@ -4101,7 +4184,7 @@ where client, &mut pending, unit, - committed, + &mut framing, unit_limit, ) .await?; @@ -4118,7 +4201,7 @@ where session, client, std::mem::take(&mut pending), - committed, + &mut framing, ) .await?; reader.read_line().await? @@ -4143,13 +4226,13 @@ where session, client, std::mem::take(&mut pending), - committed, + &mut framing, ) .await?; continue; }; let Some(unit) = next else { - flush_normalized_response_bytes(session, client, pending, committed).await?; + flush_normalized_response_bytes(session, client, pending, &mut framing).await?; return Ok(Vec::new()); }; if let Some(guard) = generation_guard { @@ -4160,31 +4243,37 @@ where client, &mut pending, unit, - committed, + &mut framing, unit_limit, ) .await?; }, BodyLength::None => { - flush_normalized_response_bytes(session, client, pending, committed).await?; + 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], +} + async fn buffer_normalized_response_bytes( session: &mut openshell_supervisor_middleware::HttpResponseSession, client: &mut C, pending: &mut Vec, data: Vec, - committed: bool, + 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, committed).await?; + process_response_unit(session, client, unit, framing).await?; } Ok(()) } @@ -4193,38 +4282,57 @@ async fn flush_normalized_response_bytes( session: &mut openshell_supervisor_middleware::HttpResponseSession, client: &mut C, pending: Vec, - committed: bool, + framing: &mut ResponseOutputState<'_>, ) -> Result<()> { if pending.is_empty() { return Ok(()); } - process_response_unit(session, client, pending, committed).await + process_response_unit(session, client, pending, framing).await } async fn process_response_unit( session: &mut openshell_supervisor_middleware::HttpResponseSession, client: &mut C, unit: Vec, - committed: bool, + framing: &mut ResponseOutputState<'_>, ) -> Result<()> { let output = session .push_body(unit) .await .map_err(|error| miette!("HTTP response middleware failure: {error}"))?; - if committed { + if !*framing.committed && !output.is_empty() { + if session.requires_whole_body() { + return Err(miette!( + "whole-body response middleware released output before finalization" + )); + } + client + .write_all(framing.commit_head) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + client + .flush() + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + *framing.committed = true; + } + if *framing.committed { for unit in output { - write_chunk(client, &unit) - .await - .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + if framing.chunked { + write_chunk(client, &unit) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + } else { + client + .write_all(&unit) + .await + .map_err(|error| miette!("HTTP response client write failed: {error}"))?; + } } client .flush() .await .map_err(|error| miette!("HTTP response client write failed: {error}"))?; - } else if !output.is_empty() { - return Err(miette!( - "whole-body response middleware released output before finalization" - )); } Ok(()) } @@ -4670,6 +4778,7 @@ mod tests { enum ResponseRelayScript { HeadersOnly, WholeBody, + WholeBodyWithTrailer, Stream, InvalidBodySequence, InvalidWholeBodySequence, @@ -4744,6 +4853,11 @@ mod tests { | ResponseRelayScript::InvalidWholeBodySequence => { (HttpResponseBodyMode::WholeBodyBytes, Vec::new(), Vec::new()) } + ResponseRelayScript::WholeBodyWithTrailer => ( + HttpResponseBodyMode::WholeBodyBytes, + Vec::new(), + vec!["digest".into()], + ), ResponseRelayScript::Stream | ResponseRelayScript::InvalidBodySequence => ( HttpResponseBodyMode::StreamBytes, @@ -4777,6 +4891,7 @@ mod tests { }; let replacement = match script { ResponseRelayScript::WholeBody + | ResponseRelayScript::WholeBodyWithTrailer | ResponseRelayScript::InvalidWholeBodySequence => { [b"whole:".as_slice(), &data].concat() } @@ -4820,6 +4935,7 @@ mod tests { trailer_mutations: if matches!( script, ResponseRelayScript::Stream + | ResponseRelayScript::WholeBodyWithTrailer ) { vec![write_header( "digest", @@ -6879,6 +6995,54 @@ mod tests { assert!(!delivered.contains("ext=yes"), "{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("0\r\n\r\n")); + } + #[tokio::test] async fn response_middleware_forwards_interim_head_before_final_preflight() { let (outcome, delivered) = run_response_middleware_relay( @@ -7120,6 +7284,7 @@ mod tests { 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(); @@ -7140,6 +7305,62 @@ mod tests { 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 From cf6faa742ca7a72720cae4eb10a9027dfbc829b0 Mon Sep 17 00:00:00 2001 From: Piotr Mlocek Date: Wed, 2 Sep 2026 16:34:25 -0700 Subject: [PATCH 6/6] fix(middleware): align response runtime with interface Signed-off-by: Piotr Mlocek --- .../src/lib.rs | 16 +- .../src/response.rs | 947 ++++++++---------- .../src/l7/rest.rs | 103 +- .../Cargo.lock | 11 + .../src/main.rs | 123 +-- 5 files changed, 512 insertions(+), 688 deletions(-) diff --git a/crates/openshell-supervisor-middleware/src/lib.rs b/crates/openshell-supervisor-middleware/src/lib.rs index f6a4e720fc..6a3e5ffc9c 100644 --- a/crates/openshell-supervisor-middleware/src/lib.rs +++ b/crates/openshell-supervisor-middleware/src/lib.rs @@ -848,6 +848,7 @@ fn validate_payload_limit(source: &str, binding: &MiddlewareBinding) -> Result Result Err(miette!( - "{source} advertises HTTP_RESPONSE/PRE_RETURN, which is not yet supported" - )), + ) => Ok(SupportedBinding::HttpResponsePreReturn), ( Some(SupervisorMiddlewareOperation::WebsocketMessage), Some(SupervisorMiddlewarePhase::PreCredentials), @@ -3703,7 +3702,7 @@ mod tests { } #[test] - fn manifest_rejects_http_response_pre_return_binding_until_dispatch_is_available() { + fn manifest_accepts_http_response_pre_return_binding_when_dispatch_is_available() { let registration = external_registration(4096); let manifest = MiddlewareManifest { name: "example/response".into(), @@ -3717,13 +3716,8 @@ mod tests { expected_audience: String::new(), }; - let error = validate_external_manifest(®istration, &manifest, 4096, false) - .expect_err("HTTP response pre-return binding must remain unavailable"); - assert!( - error - .to_string() - .contains("HTTP_RESPONSE/PRE_RETURN, which is not yet supported") - ); + validate_external_manifest(®istration, &manifest, 4096, false) + .expect("HTTP response pre-return binding is supported"); } #[test] diff --git a/crates/openshell-supervisor-middleware/src/response.rs b/crates/openshell-supervisor-middleware/src/response.rs index 052516d4c1..5eb1cc411e 100644 --- a/crates/openshell-supervisor-middleware/src/response.rs +++ b/crates/openshell-supervisor-middleware/src/response.rs @@ -3,7 +3,7 @@ //! HTTP response pre-return middleware chain execution. -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::BTreeMap; use std::time::Duration; use futures::StreamExt as _; @@ -12,13 +12,11 @@ use tokio::sync::mpsc; use tokio::time::Instant; use openshell_core::proto::{ - Finding, HeaderMutation, HttpHeader, HttpRequestTarget, HttpResponseBodyEnd, - HttpResponseBodyMode, HttpResponseBodyPassThrough, HttpResponseBodyUnit, HttpResponseEvent, - HttpResponseEventResult, HttpResponsePreflight, HttpResponseSessionEnd, - HttpResponseSessionEndReason, HttpResponseTrailers, RemoveHeader, RequestContext, - header_mutation, http_response_body_result, http_response_body_transform, - http_response_body_unit, http_response_event, http_response_event_result, - http_response_preflight_decision, + Finding, HttpHeader, HttpRequestTarget, HttpResponseBodyMode, HttpResponseBodyPassThrough, + HttpResponseBodyUnit, HttpResponseEvent, HttpResponseEventResult, HttpResponsePreflight, + 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_decision, }; use super::{ @@ -29,7 +27,7 @@ use super::{ MAX_MIDDLEWARE_PREFLIGHT_TIMEOUT, MAX_MIDDLEWARE_REASON_BYTES, MAX_MIDDLEWARE_REASON_CODE_BYTES, MAX_MIDDLEWARE_TARGET_BYTES, MiddlewareDiagnosticPolicy, MiddlewareSessionAdmission, MiddlewareSessionPermit, NamespacedFinding, OnError, headers, - is_stable_reason_code, + is_stable_reason_code, middleware_denial_reason, }; const STREAM_CHANNEL_CAPACITY: usize = 4; @@ -52,11 +50,13 @@ pub struct HttpResponsePreflightInput { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum HttpResponseInvocationOutcome { Skip, + BlockDelivery, HeadersOnly, WholeBody, Stream, PassThrough, Transform, + SkipRemaining, FailOpen, FailClosed, } @@ -130,7 +130,6 @@ struct HttpResponseStage { transport: Option, mode: StageMode, next_sequence: u64, - declared_trailer_names: BTreeSet, whole_body: Vec, } @@ -143,7 +142,7 @@ impl HttpResponseStage { self.is_active() && self.mode != StageMode::HeadersOnly } - async fn end(&mut self, reason: HttpResponseSessionEndReason) { + async fn end(&mut self, reason: MiddlewareSessionEndReason) { if let Some(transport) = self.transport.take() { let _ = tokio::time::timeout( Duration::from_millis(10), @@ -157,7 +156,6 @@ impl HttpResponseStage { pub struct HttpResponseSession { runner: ChainRunner, stages: Vec, - connection_nominated_headers: Vec, findings: Vec, metadata: BTreeMap>, invocations: Vec, @@ -184,7 +182,9 @@ impl HttpResponseSession { stage .entry .max_payload_bytes - .min(MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) + .checked_div(2) + .unwrap_or_default() + .clamp(1, MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) }) .min() .unwrap_or(MAX_HTTP_RESPONSE_STREAM_UNIT_BYTES) @@ -226,7 +226,7 @@ impl HttpResponseSession { Ok(released) } - /// Finalize every body stage, process normalized trailers, and end streams. + /// Finalize every body stage, preserve normalized trailers, and end streams. pub async fn finish( mut self, trailers: Vec, @@ -244,7 +244,7 @@ impl HttpResponseSession { let stage_output = match self.finish_stage(index, deadline).await { Ok(output) => output, Err(failure) => { - self.end_all(HttpResponseSessionEndReason::MiddlewareFailure) + self.end_all(MiddlewareSessionEndReason::MiddlewareFailure) .await; return Err(failure); } @@ -257,15 +257,11 @@ impl HttpResponseSession { } } - let trailers = match self.process_trailers(trailers, deadline).await { - Ok(trailers) => trailers, - Err(failure) => { - self.end_all(HttpResponseSessionEndReason::MiddlewareFailure) - .await; - return Err(failure); - } - }; - self.end_all(HttpResponseSessionEndReason::Normal).await; + let mut trailers = trailers; + if self.body_transformed { + strip_stale_integrity(&mut trailers); + } + self.end_all(MiddlewareSessionEndReason::Normal).await; self.session_admission.take(); Ok(HttpResponseFinish { body_units: released, @@ -277,7 +273,7 @@ impl HttpResponseSession { }) } - pub async fn end(mut self, reason: HttpResponseSessionEndReason) { + pub async fn end(mut self, reason: MiddlewareSessionEndReason) { self.end_all(reason).await; } @@ -345,8 +341,7 @@ impl HttpResponseSession { let sequence = stage.next_sequence; stage.next_sequence += 1; - let input_size = data.len(); - let event = body_event(sequence, data.clone()); + let event = body_event(sequence, data.clone(), false); let result = match exchange(stage, event, deadline).await { Ok(result) => result, Err(reason) => { @@ -355,51 +350,7 @@ impl HttpResponseSession { .await; } }; - match validate_body_result(result, sequence, stage.entry.max_payload_bytes) { - Ok(BodyDecision::PassThrough(findings, metadata)) => { - collect_diagnostics( - stage, - findings, - metadata, - &mut self.findings, - &mut self.metadata, - ); - self.invocations.push(body_invocation( - stage, - HttpResponseInvocationOutcome::PassThrough, - sequence, - input_size, - input_size, - )); - Ok(vec![data]) - } - Ok(BodyDecision::Transform(replacement, findings, metadata)) => { - self.body_transformed = true; - collect_diagnostics( - stage, - findings, - metadata, - &mut self.findings, - &mut self.metadata, - ); - self.invocations.push(body_invocation( - stage, - HttpResponseInvocationOutcome::Transform, - sequence, - input_size, - replacement.len(), - )); - if replacement.is_empty() { - Ok(Vec::new()) - } else { - Ok(vec![replacement]) - } - } - Err(reason) => { - self.handle_stage_failure(index, reason, Some(sequence), data) - .await - } - } + self.apply_body_result(index, result, sequence, data).await } async fn finish_stage( @@ -407,17 +358,22 @@ impl HttpResponseSession { index: usize, deadline: Instant, ) -> Result>, HttpResponseMiddlewareFailure> { - let stage = &mut self.stages[index]; - if !stage.is_body_active() { + if !self.stages[index].is_body_active() { return Ok(Vec::new()); } + let mode = self.stages[index].mode; let mut output = Vec::new(); - if stage.mode == StageMode::WholeBody { - let data = std::mem::take(&mut stage.whole_body); + if mode == StageMode::WholeBody { + let data = std::mem::take(&mut self.stages[index].whole_body); let sequence = 1; - stage.next_sequence = 2; - let input_size = data.len(); - let result = match exchange(stage, body_event(sequence, data.clone()), deadline).await { + 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) => { return self @@ -425,158 +381,131 @@ impl HttpResponseSession { .await; } }; - match validate_body_result(result, sequence, stage.entry.max_payload_bytes) { - Ok(BodyDecision::PassThrough(findings, metadata)) => { - collect_diagnostics( - stage, - findings, - metadata, - &mut self.findings, - &mut self.metadata, - ); - self.invocations.push(body_invocation( - stage, - HttpResponseInvocationOutcome::PassThrough, - sequence, - input_size, - input_size, - )); - output.push(data); - } - Ok(BodyDecision::Transform(replacement, findings, metadata)) => { - self.body_transformed = true; - collect_diagnostics( - stage, - findings, - metadata, - &mut self.findings, - &mut self.metadata, - ); - self.invocations.push(body_invocation( - stage, - HttpResponseInvocationOutcome::Transform, - sequence, - input_size, - replacement.len(), - )); - if !replacement.is_empty() { - output.push(replacement); - } - } - Err(reason) => { - return self - .handle_stage_failure(index, reason, Some(sequence), data) - .await; - } - } + output.extend( + self.apply_body_result(index, result, sequence, data) + .await?, + ); } - let final_sequence = stage.next_sequence.saturating_sub(1); - let body_end_sent = if let Some(transport) = stage.transport.as_ref() { - send_without_result( - &transport.sender, - HttpResponseEvent { - event: Some(http_response_event::Event::BodyEnd(HttpResponseBodyEnd { - final_sequence, - })), - }, + 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, - stage.entry.timeout, ) .await - .is_ok() - } else { - false - }; - if !body_end_sent { - self.handle_stage_failure(index, "middleware_stream_closed", None, Vec::new()) - .await?; + { + Ok(result) => result, + Err(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 process_trailers( + async fn apply_body_result( &mut self, - mut trailers: Vec, - deadline: Instant, - ) -> Result, HttpResponseMiddlewareFailure> { - if self.body_transformed { - strip_stale_integrity(&mut trailers); - } - for index in 0..self.stages.len() { - if !self.stages[index].is_body_active() { - continue; + 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 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) => { - let original = Vec::new(); - self.handle_stage_failure(index, &reason, None, original) - .await?; - continue; - } - }; - let Some(http_response_event_result::Result::TrailersResult(trailer_result)) = - result.result - else { - self.handle_stage_failure(index, "unexpected_response_result", None, Vec::new()) - .await?; - continue; - }; - if let Err(reason) = validate_diagnostics( - &trailer_result.reason, - "", - &trailer_result.findings, - &trailer_result.metadata, - ) { - self.handle_stage_failure(index, reason, None, Vec::new()) - .await?; - continue; + }; + let input_size = original.len(); + let stage = &mut self.stages[index]; + collect_diagnostics( + stage, + decision.findings, + decision.metadata, + &mut self.findings, + &mut self.metadata, + ); + let reason_code = (!decision.reason_code.is_empty()).then_some(decision.reason_code); + match decision.action { + BodyAction::PassThrough => { + let output_size = original.len(); + self.invocations.push(body_invocation_with_reason( + stage, + HttpResponseInvocationOutcome::PassThrough, + sequence, + input_size, + output_size, + reason_code, + )); + Ok((!original.is_empty()) + .then_some(original) + .into_iter() + .collect()) } - if let Err(reason) = - validate_trailer_mutations(&self.stages[index], &trailer_result.trailer_mutations) - { - self.handle_stage_failure(index, reason, None, Vec::new()) - .await?; - continue; + 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()) } - match headers::apply( - headers::HeaderAuthority::ResponseTrailers, - &trailers, - &self.connection_nominated_headers, - &trailer_result.trailer_mutations, - ) { - Ok(updated) => trailers = updated, - Err(error) => { - let reason = self.stages[index].entry.service.as_ref().map_or_else( - || error.to_string(), - |service| { - service - .diagnostic_policy - .header_mutation_error_reason(&error) - }, - ); - self.handle_stage_failure(index, &reason, None, Vec::new()) - .await?; - continue; - } + 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, + )); + Ok((!output.is_empty()).then_some(output).into_iter().collect()) + } + BodyAction::BlockDelivery => { + let denial_reason = + middleware_denial_reason(&stage.entry.entry.name, reason_code.as_deref()); + self.invocations.push(body_invocation_with_reason( + stage, + HttpResponseInvocationOutcome::BlockDelivery, + sequence, + input_size, + 0, + reason_code, + )); + self.end_all(MiddlewareSessionEndReason::MiddlewareDenial) + .await; + Err(HttpResponseMiddlewareFailure { + reason: denial_reason, + }) } - let findings = trailer_result.findings; - let metadata = trailer_result.metadata; - collect_diagnostics( - &self.stages[index], - findings, - metadata, - &mut self.findings, - &mut self.metadata, - ); } - Ok(trailers) } async fn handle_stage_failure( @@ -605,7 +534,7 @@ impl HttpResponseSession { failure_category: Some(response_failure_category(reason).into()), }); stage - .end(HttpResponseSessionEndReason::MiddlewareFailure) + .end(MiddlewareSessionEndReason::MiddlewareFailure) .await; if stage.entry.on_error() == OnError::FailOpen { if original.is_empty() { @@ -620,7 +549,7 @@ impl HttpResponseSession { } } - async fn end_all(&mut self, reason: HttpResponseSessionEndReason) { + async fn end_all(&mut self, reason: MiddlewareSessionEndReason) { for stage in &mut self.stages { stage.end(reason).await; } @@ -654,14 +583,13 @@ impl ChainRunner { let mut findings = Vec::new(); let mut metadata = BTreeMap::new(); let mut invocations = Vec::new(); - let mut declared_trailer_names = BTreeSet::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, HttpResponseSessionEndReason::MiddlewareFailure).await; + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; return Ok(failed_preflight_outcome( headers, reason, @@ -681,6 +609,12 @@ impl ChainRunner { 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(), + ), + deferral_permitted: entry.on_error() == OnError::FailClosed, }; let timeout = entry.timeout.min(MAX_MIDDLEWARE_PREFLIGHT_TIMEOUT); let opened = tokio::time::timeout(timeout, async { @@ -711,7 +645,7 @@ impl ChainRunner { if let Some(reason) = collect_preflight_failure(&entry, &reason, &mut invocations) { - end_stages(&mut stages, HttpResponseSessionEndReason::MiddlewareFailure) + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure) .await; return Ok(failed_preflight_outcome( headers, @@ -727,7 +661,7 @@ impl ChainRunner { if let Some(reason) = collect_preflight_failure(&entry, "middleware_timeout", &mut invocations) { - end_stages(&mut stages, HttpResponseSessionEndReason::MiddlewareFailure) + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure) .await; return Ok(failed_preflight_outcome( headers, @@ -748,7 +682,7 @@ impl ChainRunner { "unexpected_response_result", &mut invocations, ) { - end_stages(&mut stages, HttpResponseSessionEndReason::MiddlewareFailure).await; + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; return Ok(failed_preflight_outcome( headers, reason, @@ -759,37 +693,34 @@ impl ChainRunner { } continue; }; - match decision.decision { - Some(http_response_preflight_decision::Decision::Skip(skip)) => { - let invalid = validate_diagnostics( - &skip.reason, - &skip.reason_code, - &skip.findings, - &skip.metadata, - ); - if let Err(reason) = invalid { - if let Some(reason) = - collect_preflight_failure(&entry, reason, &mut invocations) - { - end_stages( - &mut stages, - HttpResponseSessionEndReason::MiddlewareFailure, - ) - .await; - return Ok(failed_preflight_outcome( - headers, - reason, - findings, - metadata, - invocations, - )); - } - continue; - } + if let Err(reason) = validate_diagnostics( + &decision.reason, + &decision.reason_code, + &decision.findings, + &decision.metadata, + ) { + if let Some(reason) = collect_preflight_failure(&entry, reason, &mut invocations) { + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; + return Ok(failed_preflight_outcome( + headers, + reason, + findings, + metadata, + invocations, + )); + } + continue; + } + let reason_code = + (!decision.reason_code.is_empty()).then(|| decision.reason_code.clone()); + let decision_findings = decision.findings; + let decision_metadata = decision.metadata; + match decision.action { + Some(http_response_preflight_decision::Action::Skip(_)) => { collect_preflight_diagnostics( &entry, - skip.findings, - skip.metadata, + decision_findings, + decision_metadata, &mut findings, &mut metadata, ); @@ -802,7 +733,7 @@ impl ChainRunner { output_size: None, failed: false, stage_disabled: false, - reason_code: (!skip.reason_code.is_empty()).then_some(skip.reason_code), + reason_code, failure_category: None, }); let mut skipped = HttpResponseStage { @@ -810,21 +741,14 @@ impl ChainRunner { transport: Some(HttpResponseStageTransport { sender, responses }), mode: StageMode::HeadersOnly, next_sequence: 1, - declared_trailer_names: BTreeSet::new(), whole_body: Vec::new(), }; - skipped - .end(HttpResponseSessionEndReason::StageSkipped) - .await; + skipped.end(MiddlewareSessionEndReason::StageSkipped).await; } - Some(http_response_preflight_decision::Decision::Inspect(inspect)) => { - let mode = match validate_inspect( - &entry, - &inspect, - original_restriction.as_deref(), - input.declared_body_length, - &input.connection_nominated_headers, - ) { + Some(http_response_preflight_decision::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) = @@ -832,7 +756,7 @@ impl ChainRunner { { end_stages( &mut stages, - HttpResponseSessionEndReason::MiddlewareFailure, + MiddlewareSessionEndReason::MiddlewareFailure, ) .await; return Ok(failed_preflight_outcome( @@ -862,7 +786,7 @@ impl ChainRunner { { end_stages( &mut stages, - HttpResponseSessionEndReason::MiddlewareFailure, + MiddlewareSessionEndReason::MiddlewareFailure, ) .await; return Ok(failed_preflight_outcome( @@ -880,16 +804,10 @@ impl ChainRunner { if mode == StageMode::Stream { strip_stale_integrity(&mut headers); } - let declared = normalize_declared_trailers( - &inspect.declared_trailer_names, - &input.connection_nominated_headers, - ) - .expect("inspect validation normalized trailer declarations"); - declared_trailer_names.extend(declared.iter().cloned()); collect_preflight_diagnostics( &entry, - inspect.findings, - inspect.metadata, + decision_findings, + decision_metadata, &mut findings, &mut metadata, ); @@ -906,7 +824,7 @@ impl ChainRunner { output_size: None, failed: false, stage_disabled: false, - reason_code: None, + reason_code, failure_category: None, }); stages.push(HttpResponseStage { @@ -914,17 +832,52 @@ impl ChainRunner { transport: Some(HttpResponseStageTransport { sender, responses }), mode, next_sequence: 1, - declared_trailer_names: declared, whole_body: Vec::new(), }); } + Some(http_response_preflight_decision::Action::BlockDelivery(_)) => { + collect_preflight_diagnostics( + &entry, + decision_findings, + decision_metadata, + &mut findings, + &mut metadata, + ); + invocations.push(HttpResponseInvocation { + config_name: entry.entry.name.clone(), + implementation: entry.entry.implementation.clone(), + outcome: HttpResponseInvocationOutcome::BlockDelivery, + sequence: None, + input_size: 0, + output_size: None, + failed: false, + stage_disabled: false, + reason_code: reason_code.clone(), + failure_category: None, + }); + stages.push(HttpResponseStage { + entry: entry.clone(), + transport: Some(HttpResponseStageTransport { sender, responses }), + mode: StageMode::HeadersOnly, + next_sequence: 1, + whole_body: Vec::new(), + }); + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareDenial).await; + return Ok(failed_preflight_outcome( + headers, + middleware_denial_reason(&entry.entry.name, reason_code.as_deref()), + findings, + metadata, + invocations, + )); + } None => { if let Some(reason) = collect_preflight_failure( &entry, "invalid_preflight_decision", &mut invocations, ) { - end_stages(&mut stages, HttpResponseSessionEndReason::MiddlewareFailure) + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure) .await; return Ok(failed_preflight_outcome( headers, @@ -944,7 +897,7 @@ impl ChainRunner { allowed: true, reason: String::new(), headers, - declared_trailer_names: declared_trailer_names.into_iter().collect(), + declared_trailer_names: Vec::new(), session: None, findings, metadata, @@ -959,11 +912,10 @@ impl ChainRunner { allowed: true, reason: String::new(), headers, - declared_trailer_names: declared_trailer_names.into_iter().collect(), + declared_trailer_names: Vec::new(), session: Some(HttpResponseSession { runner: self.clone(), stages, - connection_nominated_headers: input.connection_nominated_headers, findings: Vec::new(), metadata: BTreeMap::new(), invocations: Vec::new(), @@ -980,13 +932,23 @@ impl ChainRunner { } } -enum BodyDecision { - PassThrough(Vec, std::collections::HashMap), - Transform( - Vec, - Vec, - std::collections::HashMap, - ), +enum BodyAction { + PassThrough, + Transform(Vec), + BlockDelivery, + SkipRemaining(CurrentBodyAction), +} + +enum CurrentBodyAction { + PassThrough, + Transform(Vec), +} + +struct BodyDecision { + action: BodyAction, + reason_code: String, + findings: Vec, + metadata: std::collections::HashMap, } fn validate_body_result( @@ -1000,39 +962,63 @@ fn validate_body_result( if body.sequence != sequence { return Err("response_body_sequence_mismatch"); } - validate_diagnostics(&body.reason, "", &body.findings, &body.metadata)?; - match body.decision { - Some(http_response_body_result::Decision::PassThrough(HttpResponseBodyPassThrough {})) => { - Ok(BodyDecision::PassThrough(body.findings, body.metadata)) + 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::Decision::Transform(transform)) => { - let Some(http_response_body_transform::Replacement::Data(replacement)) = - transform.replacement - else { - return Err("response_body_replacement_missing"); + 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"), }; - if replacement.len() > max_payload_bytes { - return Err("response_body_replacement_over_capacity"); - } - Ok(BodyDecision::Transform( - replacement, - body.findings, - body.metadata, - )) + BodyAction::SkipRemaining(current) } - None => Err("invalid_response_body_decision"), + None => return Err("invalid_response_body_decision"), + }; + Ok(BodyDecision { + action, + reason_code: body.reason_code, + findings: body.findings, + metadata: body.metadata, + }) +} + +fn validate_replacement( + replacement: Option, + max_payload_bytes: usize, +) -> Result, &'static str> { + let Some(http_response_body_transform::Replacement::Data(replacement)) = replacement else { + return Err("response_body_replacement_missing"); + }; + if replacement.len() > max_payload_bytes { + return Err("response_body_replacement_over_capacity"); } + Ok(replacement) } fn validate_inspect( entry: &DescribedChainEntry, inspect: &openshell_core::proto::HttpResponsePreflightInspect, - body_restriction: Option<&str>, - declared_body_length: Option, - connection_nominated_headers: &[String], + permitted_modes: &[i32], ) -> Result { - validate_diagnostics(&inspect.reason, "", &inspect.findings, &inspect.metadata) - .map_err(str::to_string)?; let mode = match HttpResponseBodyMode::try_from(inspect.body_mode) { Ok(HttpResponseBodyMode::HeadersOnly) => StageMode::HeadersOnly, Ok(HttpResponseBodyMode::WholeBodyBytes) => StageMode::WholeBody, @@ -1041,25 +1027,10 @@ fn validate_inspect( return Err("invalid_response_body_mode".into()); } }; - if mode != StageMode::HeadersOnly - && let Some(restriction) = body_restriction - { - return Err(restriction.to_string()); + if !permitted_modes.contains(&inspect.body_mode) { + return Err("response_body_mode_not_permitted".into()); } - if mode == StageMode::WholeBody - && declared_body_length.is_some_and(|length| length > entry.max_payload_bytes as u64) - { - return Err("whole_body_over_capacity".into()); - } - if mode == StageMode::HeadersOnly && !inspect.declared_trailer_names.is_empty() { - return Err("response_trailer_declaration_without_body".into()); - } - if inspect - .header_mutations - .len() - .saturating_add(inspect.declared_trailer_names.len()) - > headers::MAX_HEADER_MUTATIONS - { + if inspect.header_mutations.len() > headers::MAX_HEADER_MUTATIONS { return Err("header_mutation_count_over_capacity".into()); } let encoded_mutations = inspect @@ -1068,75 +1039,15 @@ fn validate_inspect( .fold(0usize, |total, mutation| { total.saturating_add(mutation.encoded_len()) }); - let declared_bytes = inspect - .declared_trailer_names - .iter() - .fold(0usize, |total, name| total.saturating_add(name.len())); - if encoded_mutations.saturating_add(declared_bytes) > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES - { + if encoded_mutations > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES { return Err("header_mutation_bytes_over_capacity".into()); } - normalize_declared_trailers( - &inspect.declared_trailer_names, - connection_nominated_headers, - )?; if entry.max_payload_bytes == 0 && mode != StageMode::HeadersOnly { return Err("response_payload_limit_invalid".into()); } Ok(mode) } -fn normalize_declared_trailers( - names: &[String], - connection_nominated_headers: &[String], -) -> Result, String> { - let mut normalized = BTreeSet::new(); - for name in names { - let name = headers::normalize_name(name).map_err(|error| error.to_string())?; - let validation = HeaderMutation { - operation: Some(header_mutation::Operation::Remove(RemoveHeader { - name: name.clone(), - })), - }; - headers::apply( - headers::HeaderAuthority::ResponseTrailers, - &[], - connection_nominated_headers, - &[validation], - ) - .map_err(|error| error.to_string())?; - if !normalized.insert(name) { - return Err("response_trailer_declaration_duplicate".into()); - } - } - Ok(normalized) -} - -fn validate_trailer_mutations( - stage: &HttpResponseStage, - mutations: &[HeaderMutation], -) -> Result<(), &'static str> { - if mutations.len() > headers::MAX_HEADER_MUTATIONS { - return Err("header_mutation_count_over_capacity"); - } - if mutations.iter().fold(0usize, |total, mutation| { - total.saturating_add(mutation.encoded_len()) - }) > MAX_MIDDLEWARE_HEADER_MUTATION_WIRE_BYTES - { - return Err("header_mutation_bytes_over_capacity"); - } - for mutation in mutations { - if let Some(header_mutation::Operation::Write(write)) = mutation.operation.as_ref() { - let name = - headers::normalize_name(&write.name).map_err(|_| "header_mutation_invalid_name")?; - if !stage.declared_trailer_names.contains(&name) { - return Err("response_trailer_name_not_declared"); - } - } - } - Ok(()) -} - 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")); @@ -1240,6 +1151,40 @@ fn body_restriction(input: &HttpResponsePreflightInput) -> Option { None } +fn permitted_body_modes( + input: &HttpResponsePreflightInput, + entry: &DescribedChainEntry, + body_restriction: Option<&str>, +) -> Vec { + let mut modes = vec![HttpResponseBodyMode::HeadersOnly as i32]; + if body_restriction.is_some() { + return modes; + } + if input + .declared_body_length + .is_none_or(|length| length <= entry.max_payload_bytes as u64) + && !is_open_ended_response(input) + { + modes.push(HttpResponseBodyMode::WholeBodyBytes as i32); + } + if entry.max_payload_bytes >= 2 { + modes.push(HttpResponseBodyMode::StreamBytes as i32); + } + modes +} + +fn is_open_ended_response(input: &HttpResponsePreflightInput) -> bool { + input.headers.iter().any(|header| { + header.name.eq_ignore_ascii_case("content-type") + && matches!( + header.value.split(';').next().map(str::trim), + Some(value) + if value.eq_ignore_ascii_case("text/event-stream") + || value.eq_ignore_ascii_case("multipart/x-mixed-replace") + ) + }) +} + fn strip_stale_integrity(headers: &mut Vec) { headers.retain(|header| { !matches!( @@ -1298,45 +1243,34 @@ async fn exchange( } } -async fn send_without_result( - sender: &mpsc::Sender, - event: HttpResponseEvent, - chain_deadline: Instant, - stage_timeout: Duration, -) -> Result<(), ()> { - let remaining = chain_deadline.saturating_duration_since(Instant::now()); - let timeout = stage_timeout.min(remaining); - tokio::time::timeout(timeout, sender.send(event)) - .await - .map_err(|_| ())? - .map_err(|_| ()) -} - -fn body_event(sequence: u64, data: Vec) -> HttpResponseEvent { +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: HttpResponseSessionEndReason) -> HttpResponseEvent { +fn session_end_event(reason: MiddlewareSessionEndReason) -> HttpResponseEvent { HttpResponseEvent { event: Some(http_response_event::Event::SessionEnd( - HttpResponseSessionEnd { + MiddlewareSessionEnd { reason: reason as i32, + protocol_error: None, }, )), } } -fn body_invocation( +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(), @@ -1347,7 +1281,7 @@ fn body_invocation( output_size: Some(output_size), failed: false, stage_disabled: false, - reason_code: None, + reason_code, failure_category: None, } } @@ -1397,7 +1331,6 @@ fn collect_preflight_diagnostics( transport: None, mode: StageMode::HeadersOnly, next_sequence: 1, - declared_trailer_names: BTreeSet::new(), whole_body: Vec::new(), }; collect_diagnostics(&stage, findings, metadata, all_findings, all_metadata); @@ -1518,7 +1451,7 @@ fn response_session_capacity_exhausted( } } -async fn end_stages(stages: &mut [HttpResponseStage], reason: HttpResponseSessionEndReason) { +async fn end_stages(stages: &mut [HttpResponseStage], reason: MiddlewareSessionEndReason) { for stage in stages { stage.end(reason).await; } @@ -1530,10 +1463,10 @@ mod tests { use openshell_core::middleware::{HttpRequestView, InProcessMiddleware}; use openshell_core::proto::{ - Decision, ExistingHeaderAction, HttpRequestResult, HttpResponseBodyResult, + Decision, ExistingHeaderAction, HeaderMutation, HttpRequestResult, HttpResponseBodyResult, HttpResponseBodyTransform, HttpResponsePreflightDecision, HttpResponsePreflightInspect, - HttpResponsePreflightSkip, HttpResponseTrailersResult, MiddlewareBinding, - MiddlewareManifest, WriteHeader, http_response_preflight_decision, + HttpResponsePreflightSkip, MiddlewareBinding, MiddlewareManifest, WriteHeader, + header_mutation, http_response_preflight_decision, }; use tokio_stream::wrappers::ReceiverStream; use tokio_stream::wrappers::TcpListenerStream; @@ -1629,8 +1562,8 @@ mod tests { result: Some( http_response_event_result::Result::PreflightDecision( HttpResponsePreflightDecision { - decision: Some( - http_response_preflight_decision::Decision::Inspect( + action: Some( + http_response_preflight_decision::Action::Inspect( HttpResponsePreflightInspect { body_mode: HttpResponseBodyMode::HeadersOnly as i32, @@ -1638,10 +1571,10 @@ mod tests { "cache-control", "remote", )], - ..Default::default() }, ), ), + ..Default::default() }, ), ), @@ -1741,71 +1674,58 @@ mod tests { result: Some( http_response_event_result::Result::PreflightDecision( HttpResponsePreflightDecision { - decision: Some( - http_response_preflight_decision::Decision::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() - }, + action: Some( + http_response_preflight_decision::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, declared_trailer_names) = - match selected_script { - Script::HeadersOnly => ( - HttpResponseBodyMode::HeadersOnly, - vec![write_header("cache-control", "private")], - Vec::new(), - ), - Script::Stream | Script::InvalidSequence => ( - HttpResponseBodyMode::StreamBytes, - Vec::new(), - vec!["digest".into()], - ), - Script::HangBody - | Script::LargeStream - | Script::UndeclaredTrailer => ( - HttpResponseBodyMode::StreamBytes, - Vec::new(), - Vec::new(), - ), - Script::WholeBody => ( - HttpResponseBodyMode::WholeBodyBytes, - Vec::new(), - Vec::new(), - ), - Script::Configured - | Script::Skip - | Script::InvalidSkipReason => unreachable!(), - }; + 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::UndeclaredTrailer => { + (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::PreflightDecision( HttpResponsePreflightDecision { - decision: Some( - http_response_preflight_decision::Decision::Inspect( + action: Some( + http_response_preflight_decision::Action::Inspect( HttpResponsePreflightInspect { body_mode: body_mode as i32, header_mutations, - declared_trailer_names, - ..Default::default() }, ), ), + ..Default::default() }, ), ), @@ -1843,38 +1763,20 @@ mod tests { } else { body.sequence }, - decision: Some( - http_response_body_result::Decision::Transform( - HttpResponseBodyTransform { - replacement: Some( - http_response_body_transform::Replacement::Data( - replacement, - ), + action: Some(http_response_body_result::Action::Transform( + HttpResponseBodyTransform { + replacement: Some( + http_response_body_transform::Replacement::Data( + replacement, ), - }, - ), - ), + ), + }, + )), ..Default::default() }, )), } } - http_response_event::Event::Trailers(_) => HttpResponseEventResult { - result: Some(http_response_event_result::Result::TrailersResult( - HttpResponseTrailersResult { - trailer_mutations: if matches!( - selected_script, - Script::Stream | Script::UndeclaredTrailer - ) { - vec![write_header("digest", "sha-256=:test:")] - } else { - Vec::new() - }, - ..Default::default() - }, - )), - }, - http_response_event::Event::BodyEnd(_) => continue, http_response_event::Event::SessionEnd(_) => break, }; if sender.send(Ok(result)).await.is_err() { @@ -1992,7 +1894,7 @@ mod tests { } #[tokio::test] - async fn stream_mode_transforms_lockstep_units_and_declared_trailer() { + async fn stream_mode_transforms_lockstep_units_and_preserves_trailers() { let runner = ChainRunner::new(Arc::new(ResponseService { script: Script::Stream, })); @@ -2000,7 +1902,7 @@ mod tests { .preflight_http_response(&[entry(OnError::FailClosed)], input(200)) .await .expect("response preflight"); - assert_eq!(outcome.declared_trailer_names, vec!["digest"]); + assert!(outcome.declared_trailer_names.is_empty()); let mut session = outcome.session.take().expect("streaming session"); assert_eq!( @@ -2010,15 +1912,16 @@ mod tests { .expect("transform stream unit"), vec![b"HELLO".to_vec()] ); - let finish = session.finish(Vec::new()).await.expect("finish stream"); + 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, - vec![HttpHeader { - name: "digest".into(), - value: "sha-256=:test:".into(), - }] - ); + assert_eq!(finish.trailers, original_trailers); } #[tokio::test] @@ -2182,26 +2085,28 @@ mod tests { } #[tokio::test] - async fn undeclared_response_trailer_obeys_on_error() { - for (on_error, allowed) in [(OnError::FailOpen, true), (OnError::FailClosed, false)] { - let runner = ChainRunner::new(Arc::new(ResponseService { - script: Script::UndeclaredTrailer, - })); - let mut outcome = runner - .preflight_http_response(&[entry(on_error)], 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 finish = session.finish(Vec::new()).await; - assert_eq!(finish.is_ok(), allowed); - if let Ok(finish) = finish { - assert!(finish.trailers.is_empty()); - } - } + async fn response_trailers_bypass_middleware_in_v1() { + let runner = ChainRunner::new(Arc::new(ResponseService { + script: Script::UndeclaredTrailer, + })); + 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, trailers); } #[tokio::test] diff --git a/crates/openshell-supervisor-network/src/l7/rest.rs b/crates/openshell-supervisor-network/src/l7/rest.rs index d6b1e5f88a..95c9efa093 100644 --- a/crates/openshell-supervisor-network/src/l7/rest.rs +++ b/crates/openshell-supervisor-network/src/l7/rest.rs @@ -3544,13 +3544,13 @@ where if committed { if let Err(error) = client.write_all(&streaming_head).await { session - .end(openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect) + .end(openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect) .await; return Err(error).into_diagnostic(); } if let Err(error) = client.flush().await { session - .end(openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect) + .end(openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect) .await; return Err(error).into_diagnostic(); } @@ -3578,19 +3578,19 @@ where .generation_guard .is_some_and(PolicyGenerationGuard::is_stale) { - openshell_core::proto::HttpResponseSessionEndReason::PolicyReload + openshell_core::proto::MiddlewareSessionEndReason::PolicyReload } else if error .to_string() .starts_with("HTTP response client write failed:") { - openshell_core::proto::HttpResponseSessionEndReason::ClientDisconnect + openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect } else if error .to_string() .starts_with("HTTP response middleware failure:") { - openshell_core::proto::HttpResponseSessionEndReason::MiddlewareFailure + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareFailure } else { - openshell_core::proto::HttpResponseSessionEndReason::UpstreamError + openshell_core::proto::MiddlewareSessionEndReason::UpstreamDisconnect }; session.end(end_reason).await; if committed { @@ -3624,7 +3624,7 @@ where && let Err(error) = guard.ensure_current() { session - .end(openshell_core::proto::HttpResponseSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; return Err(error); } @@ -4755,11 +4755,10 @@ mod tests { use openshell_core::proto::{ Decision, HttpRequestResult, HttpResponseBodyMode, HttpResponseBodyResult, HttpResponseBodyTransform, HttpResponseEvent, HttpResponseEventResult, - HttpResponsePreflightDecision, HttpResponsePreflightInspect, 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_decision, + HttpResponsePreflightDecision, HttpResponsePreflightInspect, 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_decision, }; use openshell_core::secrets::SecretResolver; use std::pin::Pin; @@ -4838,8 +4837,7 @@ mod tests { }; let result = match event { http_response_event::Event::Preflight(_) => { - let (body_mode, header_mutations, declared_trailer_names) = match script - { + let (body_mode, header_mutations) = match script { ResponseRelayScript::HeadersOnly => ( HttpResponseBodyMode::HeadersOnly, vec![write_header( @@ -4847,38 +4845,30 @@ mod tests { "private", ExistingHeaderAction::Overwrite, )], - Vec::new(), ), ResponseRelayScript::WholeBody - | ResponseRelayScript::InvalidWholeBodySequence => { - (HttpResponseBodyMode::WholeBodyBytes, Vec::new(), Vec::new()) + | ResponseRelayScript::InvalidWholeBodySequence + | ResponseRelayScript::WholeBodyWithTrailer => { + (HttpResponseBodyMode::WholeBodyBytes, Vec::new()) } - ResponseRelayScript::WholeBodyWithTrailer => ( - HttpResponseBodyMode::WholeBodyBytes, - Vec::new(), - vec!["digest".into()], - ), ResponseRelayScript::Stream - | ResponseRelayScript::InvalidBodySequence => ( - HttpResponseBodyMode::StreamBytes, - Vec::new(), - vec!["digest".into()], - ), + | ResponseRelayScript::InvalidBodySequence => { + (HttpResponseBodyMode::StreamBytes, Vec::new()) + } }; HttpResponseEventResult { result: Some( http_response_event_result::Result::PreflightDecision( HttpResponsePreflightDecision { - decision: Some( - http_response_preflight_decision::Decision::Inspect( + action: Some( + http_response_preflight_decision::Action::Inspect( HttpResponsePreflightInspect { body_mode: body_mode as i32, header_mutations, - declared_trailer_names, - ..Default::default() }, ), ), + ..Default::default() }, ), ), @@ -4913,43 +4903,20 @@ mod tests { } else { body.sequence }, - decision: Some( - http_response_body_result::Decision::Transform( - HttpResponseBodyTransform { - replacement: Some( - http_response_body_transform::Replacement::Data( - replacement, - ), + action: Some(http_response_body_result::Action::Transform( + HttpResponseBodyTransform { + replacement: Some( + http_response_body_transform::Replacement::Data( + replacement, ), - }, - ), - ), + ), + }, + )), ..Default::default() }, )), } } - http_response_event::Event::Trailers(_) => HttpResponseEventResult { - result: Some(http_response_event_result::Result::TrailersResult( - HttpResponseTrailersResult { - trailer_mutations: if matches!( - script, - ResponseRelayScript::Stream - | ResponseRelayScript::WholeBodyWithTrailer - ) { - vec![write_header( - "digest", - "sha-256=:test:", - ExistingHeaderAction::Overwrite, - )] - } else { - Vec::new() - }, - ..Default::default() - }, - )), - }, - http_response_event::Event::BodyEnd(_) => continue, http_response_event::Event::SessionEnd(_) => break, }; if sender.send(Ok(result)).await.is_err() { @@ -6973,7 +6940,7 @@ mod tests { } #[tokio::test] - async fn response_middleware_streams_normalized_chunk_payloads_and_trailers() { + 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", @@ -6982,16 +6949,10 @@ mod tests { .await; assert!(matches!(outcome.unwrap(), RelayOutcome::Reusable)); let delivered = String::from_utf8(delivered).unwrap(); - assert!( - delivered.contains("Trailer: x-upstream, digest\r\n"), - "{delivered}" - ); + 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: sha-256=:test:\r\n"), - "{delivered}" - ); + assert!(!delivered.contains("digest:"), "{delivered}"); assert!(!delivered.contains("ext=yes"), "{delivered}"); } diff --git a/examples/supervisor-middleware-content-guard/Cargo.lock b/examples/supervisor-middleware-content-guard/Cargo.lock index ad79004dbe..7c80c308c0 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.lock +++ b/examples/supervisor-middleware-content-guard/Cargo.lock @@ -979,6 +979,8 @@ dependencies = [ "prost-types", "protoc-bin-vendored", "rustix", + "rustls", + "rustls-pemfile", "serde", "serde_json", "thiserror", @@ -1417,6 +1419,15 @@ dependencies = [ "security-framework", ] +[[package]] +name = "rustls-pemfile" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "rustls-pki-types" version = "1.14.1" diff --git a/examples/supervisor-middleware-content-guard/src/main.rs b/examples/supervisor-middleware-content-guard/src/main.rs index 95f40fad17..2d766ce541 100644 --- a/examples/supervisor-middleware-content-guard/src/main.rs +++ b/examples/supervisor-middleware-content-guard/src/main.rs @@ -17,15 +17,14 @@ use openshell_core::proto::{ Decision, ExistingHeaderAction, Finding, HeaderMutation, HttpRequestEvaluation, HttpRequestResult, HttpResponseBodyMode, HttpResponseBodyResult, HttpResponseBodyTransform, HttpResponseEvent, HttpResponseEventResult, HttpResponsePreflightDecision, - HttpResponsePreflightInspect, HttpResponsePreflightSkip, HttpResponseTrailersResult, - MiddlewareBinding, MiddlewareManifest, SupervisorMiddlewareOperation, - SupervisorMiddlewarePhase, ValidateConfigRequest, ValidateConfigResponse, WebSocketMessage, - WebSocketMessageResult, WebSocketPreflightAction, WebSocketPreflightDecision, - WebSocketSessionEvent, WebSocketSessionEventResult, WriteHeader, header_mutation, - http_response_body_result, http_response_body_transform, http_response_body_unit, - http_response_event, http_response_event_result, http_response_preflight_decision, - web_socket_message, web_socket_message_result, web_socket_session_event, - web_socket_session_event_result, + HttpResponsePreflightInspect, HttpResponsePreflightSkip, MiddlewareBinding, MiddlewareManifest, + SupervisorMiddlewareOperation, SupervisorMiddlewarePhase, ValidateConfigRequest, + ValidateConfigResponse, WebSocketMessage, WebSocketMessageResult, WebSocketPreflightAction, + WebSocketPreflightDecision, WebSocketSessionEvent, WebSocketSessionEventResult, WriteHeader, + header_mutation, http_response_body_result, http_response_body_transform, + http_response_body_unit, http_response_event, http_response_event_result, + http_response_preflight_decision, 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; @@ -310,7 +309,7 @@ enum ResponseMode { struct ResponseSessionState { selected: Option, next_sequence: u64, - input_bytes: u64, + body_ended: bool, } impl ResponseSessionState { @@ -332,6 +331,7 @@ impl ResponseSessionState { }; self.selected = Some(selected); self.next_sequence = 1; + self.body_ended = false; Ok(response_preflight_inspect(selected)) } @@ -347,6 +347,11 @@ impl ResponseSessionState { "headers-only sessions do not receive body events", )); } + if self.body_ended { + return Err(Status::failed_precondition( + "body arrived after end_of_stream", + )); + } if body.sequence != self.next_sequence { return Err(Status::invalid_argument(format!( "expected body sequence {}, received {}", @@ -357,7 +362,7 @@ impl ResponseSessionState { let Some(http_response_body_unit::Payload::Data(data)) = body.payload else { return Err(Status::invalid_argument("body data is required")); }; - self.input_bytes = self.input_bytes.saturating_add(data.len() as u64); + self.body_ended = body.end_of_stream; let replacement = match selected { ResponseMode::WholeBody => [b"[whole] ".as_slice(), &data].concat(), ResponseMode::Stream => data.to_ascii_uppercase(), @@ -367,7 +372,7 @@ impl ResponseSessionState { result: Some(http_response_event_result::Result::BodyResult( HttpResponseBodyResult { sequence: body.sequence, - decision: Some(http_response_body_result::Decision::Transform( + action: Some(http_response_body_result::Action::Transform( HttpResponseBodyTransform { replacement: Some(http_response_body_transform::Replacement::Data( replacement, @@ -379,43 +384,6 @@ impl ResponseSessionState { )), }) } - - fn body_end(&self, final_sequence: u64) -> Result<(), Status> { - let expected = self.next_sequence.saturating_sub(1); - if final_sequence != expected { - return Err(Status::invalid_argument(format!( - "expected final body sequence {expected}, received {final_sequence}" - ))); - } - Ok(()) - } - - fn trailers(&self) -> Result { - let selected = self - .selected - .ok_or_else(|| Status::failed_precondition("trailers arrived before preflight"))?; - if selected == ResponseMode::HeadersOnly { - return Err(Status::failed_precondition( - "headers-only sessions do not receive trailers", - )); - } - let trailer_mutations = if selected == ResponseMode::Stream { - vec![write_header( - "x-example-body-bytes", - &self.input_bytes.to_string(), - )] - } else { - Vec::new() - }; - Ok(HttpResponseEventResult { - result: Some(http_response_event_result::Result::TrailersResult( - HttpResponseTrailersResult { - trailer_mutations, - ..Default::default() - }, - )), - }) - } } fn response_mode_for_path(path: &str) -> Option { @@ -431,26 +399,22 @@ fn response_preflight_skip() -> HttpResponseEventResult { HttpResponseEventResult { result: Some(http_response_event_result::Result::PreflightDecision( HttpResponsePreflightDecision { - decision: Some(http_response_preflight_decision::Decision::Skip( - HttpResponsePreflightSkip { - reason: "path is outside the response example".into(), - reason_code: "path_not_selected".into(), - ..Default::default() - }, + action: Some(http_response_preflight_decision::Action::Skip( + HttpResponsePreflightSkip {}, )), + reason: "path is outside the response example".into(), + reason_code: "path_not_selected".into(), + ..Default::default() }, )), } } fn response_preflight_inspect(selected: ResponseMode) -> HttpResponseEventResult { - let (body_mode, declared_trailer_names) = match selected { - ResponseMode::HeadersOnly => (HttpResponseBodyMode::HeadersOnly, Vec::new()), - ResponseMode::WholeBody => (HttpResponseBodyMode::WholeBodyBytes, Vec::new()), - ResponseMode::Stream => ( - HttpResponseBodyMode::StreamBytes, - vec!["x-example-body-bytes".into()], - ), + let body_mode = match selected { + ResponseMode::HeadersOnly => HttpResponseBodyMode::HeadersOnly, + ResponseMode::WholeBody => HttpResponseBodyMode::WholeBodyBytes, + ResponseMode::Stream => HttpResponseBodyMode::StreamBytes, }; let mode_name = match selected { ResponseMode::HeadersOnly => "headers-only", @@ -460,14 +424,13 @@ fn response_preflight_inspect(selected: ResponseMode) -> HttpResponseEventResult HttpResponseEventResult { result: Some(http_response_event_result::Result::PreflightDecision( HttpResponsePreflightDecision { - decision: Some(http_response_preflight_decision::Decision::Inspect( + action: Some(http_response_preflight_decision::Action::Inspect( HttpResponsePreflightInspect { body_mode: body_mode as i32, header_mutations: vec![write_header("x-example-response-mode", mode_name)], - declared_trailer_names, - ..Default::default() }, )), + ..Default::default() }, )), } @@ -508,10 +471,6 @@ impl HttpResponsePreReturn for ContentGuard { state.preflight(preflight).map(Some) } Some(http_response_event::Event::Body(body)) => state.body(body).map(Some), - Some(http_response_event::Event::BodyEnd(body_end)) => { - state.body_end(body_end.final_sequence).map(|()| None) - } - Some(http_response_event::Event::Trailers(_)) => state.trailers().map(Some), Some(http_response_event::Event::SessionEnd(_)) => break, None => Err(Status::invalid_argument("response event is required")), }; @@ -819,8 +778,7 @@ mod tests { else { panic!("expected preflight decision"); }; - let Some(http_response_preflight_decision::Decision::Inspect(inspect)) = - decision.decision + let Some(http_response_preflight_decision::Action::Inspect(inspect)) = decision.action else { panic!("expected inspect decision"); }; @@ -840,13 +798,13 @@ mod tests { .body(HttpResponseBodyUnit { sequence: 1, payload: Some(http_response_body_unit::Payload::Data(b"hello".to_vec())), + end_of_stream: true, }) .unwrap(); let Some(http_response_event_result::Result::BodyResult(body)) = result.result else { panic!("expected body result"); }; - let Some(http_response_body_result::Decision::Transform(transform)) = body.decision - else { + let Some(http_response_body_result::Action::Transform(transform)) = body.action else { panic!("expected body transform"); }; let Some(http_response_body_transform::Replacement::Data(data)) = transform.replacement @@ -858,7 +816,7 @@ mod tests { } #[test] - fn stream_declares_and_writes_byte_count_trailer() { + fn stream_tracks_end_of_stream_on_the_final_unit() { let mut state = ResponseSessionState::default(); let preflight = state.preflight(response_preflight("/stream")).unwrap(); let Some(http_response_event_result::Result::PreflightDecision(decision)) = @@ -866,24 +824,19 @@ mod tests { else { panic!("expected preflight decision"); }; - let Some(http_response_preflight_decision::Decision::Inspect(inspect)) = decision.decision + let Some(http_response_preflight_decision::Action::Inspect(inspect)) = decision.action else { panic!("expected inspect decision"); }; - assert_eq!(inspect.declared_trailer_names, ["x-example-body-bytes"]); + assert_eq!(inspect.body_mode, HttpResponseBodyMode::StreamBytes as i32); state .body(HttpResponseBodyUnit { sequence: 1, payload: Some(http_response_body_unit::Payload::Data(b"hello".to_vec())), + end_of_stream: true, }) .unwrap(); - state.body_end(1).unwrap(); - let trailers = state.trailers().unwrap(); - let Some(http_response_event_result::Result::TrailersResult(trailers)) = trailers.result - else { - panic!("expected trailer result"); - }; - assert_eq!(trailers.trailer_mutations.len(), 1); + assert!(state.body_ended); } #[test] @@ -894,10 +847,10 @@ mod tests { else { panic!("expected preflight decision"); }; - let Some(http_response_preflight_decision::Decision::Skip(skip)) = decision.decision else { + let Some(http_response_preflight_decision::Action::Skip(_)) = decision.action else { panic!("expected skip decision"); }; - assert_eq!(skip.reason_code, "path_not_selected"); + assert_eq!(decision.reason_code, "path_not_selected"); } #[tokio::test]