diff --git a/.github/workflows/branch-checks.yml b/.github/workflows/branch-checks.yml index 2d9dc9914f..e927c877b7 100644 --- a/.github/workflows/branch-checks.yml +++ b/.github/workflows/branch-checks.yml @@ -154,12 +154,14 @@ jobs: cargo fmt --all -- --check cargo fmt --manifest-path e2e/rust/Cargo.toml --all -- --check cargo fmt --manifest-path examples/governance-interceptor/Cargo.toml --all -- --check + cargo fmt --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml --all -- --check - name: Lint run: | cargo clippy --workspace --all-targets -- -D warnings cargo clippy --manifest-path e2e/rust/Cargo.toml --all-targets -- -D warnings cargo check --manifest-path examples/governance-interceptor/Cargo.toml --all-targets + cargo check --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml --all-targets - name: Test env: diff --git a/crates/openshell-core/src/middleware.rs b/crates/openshell-core/src/middleware.rs index 2b3fb18982..d1d59f2a8c 100644 --- a/crates/openshell-core/src/middleware.rs +++ b/crates/openshell-core/src/middleware.rs @@ -11,11 +11,17 @@ use tokio::sync::mpsc; use tonic::{Request, Response, Status}; use crate::proto::{ - HttpHeader, HttpRequestEvaluation, HttpRequestResult, HttpRequestTarget, MiddlewareManifest, - RequestContext, SupervisorMiddlewarePhase, ValidateConfigRequest, ValidateConfigResponse, - WebSocketSessionEvent, WebSocketSessionEventResult, + HttpHeader, HttpRequestEvaluation, HttpRequestResult, HttpRequestTarget, HttpResponseEvent, + HttpResponseEventResult, MiddlewareManifest, RequestContext, SupervisorMiddlewarePhase, + ValidateConfigRequest, ValidateConfigResponse, WebSocketSessionEvent, + WebSocketSessionEventResult, }; +/// Transport-neutral result stream for one HTTP response middleware stage. +pub type HttpResponseResultStream = Pin< + Box> + Send + 'static>, +>; + /// Transport-neutral response stream for one WebSocket middleware stage. pub type WebSocketResponseStream = Pin< Box< @@ -47,6 +53,15 @@ pub trait SupervisorMiddlewareEndpoint: Send + Sync { &self, requests: mpsc::Receiver, ) -> Result; + + async fn open_http_response_pre_return( + &self, + _requests: mpsc::Receiver, + ) -> Result { + Err(Status::unimplemented( + "middleware does not implement HTTP response pre-return evaluation", + )) + } } /// Borrowed request state exposed to one in-process middleware invocation. @@ -242,6 +257,18 @@ pub trait InProcessMiddleware: Send + Sync { "middleware does not implement WebSocket sessions", )) } + + /// Open one HTTP response pre-return stream. + /// + /// Request-only implementations may keep the default unsupported response. + async fn open_http_response_pre_return( + &self, + _requests: mpsc::Receiver, + ) -> std::result::Result { + Err(Status::unimplemented( + "middleware does not implement HTTP response pre-return evaluation", + )) + } } /// Default timeout for one supervisor middleware RPC. diff --git a/crates/openshell-supervisor-middleware/src/lib.rs b/crates/openshell-supervisor-middleware/src/lib.rs index bb27d26270..39aca27845 100644 --- a/crates/openshell-supervisor-middleware/src/lib.rs +++ b/crates/openshell-supervisor-middleware/src/lib.rs @@ -33,7 +33,8 @@ use tokio::sync::{OnceCell, OwnedSemaphorePermit, Semaphore}; use tonic::{Request, Response as TonicResponse, Status as TonicStatus}; pub use openshell_core::middleware::{ - HttpRequestView, InProcessMiddleware, SupervisorMiddlewareEndpoint, WebSocketResponseStream, + HttpRequestView, HttpResponseResultStream, InProcessMiddleware, SupervisorMiddlewareEndpoint, + WebSocketResponseStream, }; pub type MiddlewareService = dyn SupervisorMiddleware; @@ -180,6 +181,13 @@ impl InProcessMiddleware for EndpointInProcessAdapter { ) -> std::result::Result { self.endpoint.open_websocket_session(requests).await } + + async fn open_http_response_pre_return( + &self, + requests: tokio::sync::mpsc::Receiver, + ) -> std::result::Result { + self.endpoint.open_http_response_pre_return(requests).await + } } /// Adapt a transport-neutral endpoint to the in-process registry contract. @@ -835,6 +843,12 @@ fn supported_binding(source: &str, binding: &MiddlewareBinding) -> Result Ok(SupportedBinding::HttpPreCredentials), + ( + Some(SupervisorMiddlewareOperation::HttpResponse), + Some(SupervisorMiddlewarePhase::PreReturn), + ) => Err(miette!( + "{source} advertises HTTP_RESPONSE/PRE_RETURN, which is not yet supported" + )), ( Some(SupervisorMiddlewareOperation::WebsocketMessage), Some(SupervisorMiddlewarePhase::PreCredentials), @@ -843,7 +857,7 @@ fn supported_binding(source: &str, binding: &MiddlewareBinding) -> Result Err(miette!( - "{source} advertises WEBSOCKET_MESSAGE/PRE_RETURN, which is reserved for PR 2" + "{source} advertises WEBSOCKET_MESSAGE/PRE_RETURN, which is not yet supported" )), _ => Err(miette!( "{source} advertises an unsupported middleware operation/phase pair" @@ -1435,6 +1449,20 @@ impl ChainRunner { .entries) } + pub async fn describe_http_response_chain( + &self, + entries: &[ChainEntry], + ) -> Result> { + Ok(self + .describe_chain_for( + entries, + SupervisorMiddlewareOperation::HttpResponse, + SupervisorMiddlewarePhase::PreReturn, + ) + .await? + .entries) + } + async fn describe_chain_for( &self, entries: &[ChainEntry], @@ -3657,6 +3685,30 @@ mod tests { ); } + #[test] + fn manifest_rejects_http_response_pre_return_binding_until_dispatch_is_available() { + let registration = external_registration(4096); + let manifest = MiddlewareManifest { + name: "example/response".into(), + service_version: "test".into(), + bindings: vec![MiddlewareBinding { + operation: SupervisorMiddlewareOperation::HttpResponse as i32, + phase: SupervisorMiddlewarePhase::PreReturn as i32, + max_payload_bytes: 4096, + timeout: "500ms".into(), + }], + 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") + ); + } + #[test] fn manifest_accepts_forward_websocket_binding_and_reserves_return_phase() { let binding = |phase| MiddlewareBinding { @@ -3676,8 +3728,8 @@ mod tests { manifest.bindings = vec![binding(SupervisorMiddlewarePhase::PreReturn)]; let error = validate_manifest_bindings("test WebSocket service", &manifest, None) - .expect_err("return-path binding stays reserved for PR 2"); - assert!(error.to_string().contains("reserved for PR 2")); + .expect_err("return-path WebSocket binding is not yet supported"); + assert!(error.to_string().contains("not yet supported")); } #[test] @@ -4666,7 +4718,7 @@ mod tests { close_on_first_message: bool, messages: Arc, session_ends: Option< - tokio::sync::mpsc::UnboundedSender, + tokio::sync::mpsc::UnboundedSender, >, } @@ -4757,7 +4809,7 @@ mod tests { Some(web_socket_session_event::Event::SessionEnd(end)) => { if let Some(session_ends) = &session_ends && let Ok(reason) = - openshell_core::proto::WebSocketSessionEndReason::try_from( + openshell_core::proto::MiddlewareSessionEndReason::try_from( end.reason, ) { @@ -5042,14 +5094,14 @@ mod tests { assert!(!text.invocations[0].failed); session - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; } } #[tokio::test] async fn explicit_websocket_preflight_denial_is_authoritative_for_both_error_modes() { - use openshell_core::proto::WebSocketSessionEndReason; + use openshell_core::proto::MiddlewareSessionEndReason; for on_error in [OnError::FailOpen, OnError::FailClosed] { let (session_ends_tx, mut session_ends_rx) = tokio::sync::mpsc::unbounded_channel(); @@ -5085,7 +5137,7 @@ mod tests { assert!(!outcome.allowed); assert_eq!( outcome.terminal_reason, - Some(WebSocketSessionEndReason::MiddlewareDenial) + Some(MiddlewareSessionEndReason::MiddlewareDenial) ); assert_eq!( outcome.reason, @@ -5123,7 +5175,7 @@ mod tests { assert!(!outcome.invocations[0].failed); assert_eq!( session_ends_rx.recv().await, - Some(WebSocketSessionEndReason::MiddlewareDenial) + Some(MiddlewareSessionEndReason::MiddlewareDenial) ); assert!( session_ends_rx.try_recv().is_err(), @@ -5134,7 +5186,7 @@ mod tests { #[tokio::test] async fn mixed_websocket_preflight_denial_ends_every_opened_stage() { - use openshell_core::proto::WebSocketSessionEndReason; + use openshell_core::proto::MiddlewareSessionEndReason; let (first_end_tx, mut first_end_rx) = tokio::sync::mpsc::unbounded_channel(); let (denier_end_tx, mut denier_end_rx) = tokio::sync::mpsc::unbounded_channel(); @@ -5176,7 +5228,7 @@ mod tests { assert!(!outcome.allowed); assert_eq!( outcome.terminal_reason, - Some(WebSocketSessionEndReason::MiddlewareDenial) + Some(MiddlewareSessionEndReason::MiddlewareDenial) ); assert_eq!( outcome @@ -5193,7 +5245,7 @@ mod tests { for receiver in [&mut first_end_rx, &mut denier_end_rx, &mut last_end_rx] { assert_eq!( receiver.recv().await, - Some(WebSocketSessionEndReason::MiddlewareDenial) + Some(MiddlewareSessionEndReason::MiddlewareDenial) ); assert!( receiver.try_recv().is_err(), @@ -5270,7 +5322,7 @@ mod tests { "middleware_failed: request_message_over_capacity" ); session - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; } @@ -5327,7 +5379,7 @@ mod tests { assert!(!redacted.invocations[0].stage_disabled); session - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; } @@ -5373,7 +5425,7 @@ mod tests { ); assert!(outcome.invocations[0].transformed); session - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; } @@ -5459,7 +5511,7 @@ mod tests { assert!(target.query.is_empty()); assert_eq!(observed.requested_subprotocols, ["realtime"]); session - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; let _ = shutdown_tx.send(()); server_task @@ -5570,7 +5622,7 @@ mod tests { drop(work); session - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; let _ = shutdown_tx.send(()); server_task @@ -5830,7 +5882,7 @@ mod tests { ); assert_eq!( session_ends_rx.recv().await, - Some(openshell_core::proto::WebSocketSessionEndReason::StageSkipped) + Some(openshell_core::proto::MiddlewareSessionEndReason::StageSkipped) ); assert!( session_ends_rx.try_recv().is_err(), @@ -5876,7 +5928,7 @@ mod tests { sessions .pop() .expect("retained session") - .end(openshell_core::proto::WebSocketSessionEndReason::NormalClose) + .end(openshell_core::proto::MiddlewareSessionEndReason::Normal) .await; assert_eq!(runner.registry.session_admission.available_permits(), 1); @@ -5918,7 +5970,7 @@ mod tests { sessions .pop() .expect("retained old-generation session") - .end(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; let admitted = replacement .preflight_websocket(&chain, websocket_preflight_input("new-generation-admitted")) diff --git a/crates/openshell-supervisor-middleware/src/remote.rs b/crates/openshell-supervisor-middleware/src/remote.rs index edc1e8066c..9443038100 100644 --- a/crates/openshell-supervisor-middleware/src/remote.rs +++ b/crates/openshell-supervisor-middleware/src/remote.rs @@ -3,12 +3,14 @@ use miette::{IntoDiagnostic, Result, WrapErr}; use openshell_core::middleware::{ - HttpRequestView, SupervisorMiddlewareEndpoint, WebSocketResponseStream, + HttpRequestView, HttpResponseResultStream, SupervisorMiddlewareEndpoint, + WebSocketResponseStream, }; +use openshell_core::proto::middleware::v1::http_response_pre_return_client::HttpResponsePreReturnClient; use openshell_core::proto::middleware::v1::supervisor_middleware_client::SupervisorMiddlewareClient; use openshell_core::proto::{ - HttpRequestEvaluation, HttpRequestResult, MiddlewareManifest, ValidateConfigRequest, - ValidateConfigResponse, WebSocketSessionEvent, + HttpRequestEvaluation, HttpRequestResult, HttpResponseEvent, MiddlewareManifest, + ValidateConfigRequest, ValidateConfigResponse, WebSocketSessionEvent, }; use openshell_extension_core::{ BearerTokenInterceptor, BearerTokenSlot, ExtensionChannelConfig, ExtensionServerTrust, @@ -106,6 +108,7 @@ impl GrpcMiddlewareService { #[derive(Clone)] pub struct RemoteMiddlewareService { client: SupervisorMiddlewareClient, + response_client: HttpResponsePreReturnClient, } impl RemoteMiddlewareService { @@ -133,7 +136,10 @@ impl RemoteMiddlewareService { let channel = InterceptedService::new(channel, interceptor); Ok(Self { - client: SupervisorMiddlewareClient::new(channel) + client: SupervisorMiddlewareClient::new(channel.clone()) + .max_decoding_message_size(MIDDLEWARE_GRPC_MESSAGE_BYTES) + .max_encoding_message_size(MIDDLEWARE_GRPC_MESSAGE_BYTES), + response_client: HttpResponsePreReturnClient::new(channel) .max_decoding_message_size(MIDDLEWARE_GRPC_MESSAGE_BYTES) .max_encoding_message_size(MIDDLEWARE_GRPC_MESSAGE_BYTES), }) @@ -179,4 +185,18 @@ impl SupervisorMiddlewareEndpoint for RemoteMiddlewareService { .into_inner(); Ok(Box::pin(responses)) } + + async fn open_http_response_pre_return( + &self, + receiver: tokio::sync::mpsc::Receiver, + ) -> std::result::Result { + let mut client = self.response_client.clone(); + let responses = client + .evaluate(Request::new(tokio_stream::wrappers::ReceiverStream::new( + receiver, + ))) + .await? + .into_inner(); + Ok(Box::pin(responses)) + } } diff --git a/crates/openshell-supervisor-middleware/src/websocket.rs b/crates/openshell-supervisor-middleware/src/websocket.rs index 1fd95021a8..e7c9956c8d 100644 --- a/crates/openshell-supervisor-middleware/src/websocket.rs +++ b/crates/openshell-supervisor-middleware/src/websocket.rs @@ -12,11 +12,12 @@ use tokio::sync::mpsc; use tokio::time::Instant; use openshell_core::proto::{ - Decision, HttpRequestTarget, RequestContext, SupervisorMiddlewarePhase, WebSocketMessage, + Decision, HttpRequestTarget, MiddlewareSessionEnd, MiddlewareSessionEndReason, + MiddlewareSessionProtocolError, RequestContext, SupervisorMiddlewarePhase, WebSocketMessage, WebSocketMessageResult, WebSocketPreflight, WebSocketPreflightAction, - WebSocketPreflightDecision, WebSocketSessionEnd, WebSocketSessionEndReason, - WebSocketSessionEvent, WebSocketSessionStart, web_socket_message, web_socket_message_result, - web_socket_session_event, web_socket_session_event_result, + WebSocketPreflightDecision, WebSocketProtocolError, WebSocketSessionEvent, + WebSocketSessionStart, middleware_session_protocol_error, web_socket_message, + web_socket_message_result, web_socket_session_event, web_socket_session_event_result, }; use super::{ @@ -106,7 +107,7 @@ pub struct WebSocketPreflightResult { pub allowed: bool, /// Typed terminal reason when preflight denied the upgrade. `None` means /// the request may continue, including voluntary skip and fail-open. - pub terminal_reason: Option, + pub terminal_reason: Option, pub reason: String, pub denial: Option, pub session: Option, @@ -122,7 +123,7 @@ pub struct WebSocketPreflightResult { pub struct WebSocketSessionStartOutcome { pub allowed: bool, /// Typed terminal reason when session start cannot continue. - pub terminal_reason: Option, + pub terminal_reason: Option, pub reason: String, pub invocations: Vec, } @@ -156,10 +157,11 @@ impl WebSocketStage { } async fn disable(&mut self) { - self.end(WebSocketSessionEndReason::MiddlewareFailure).await; + self.end(MiddlewareSessionEndReason::MiddlewareFailure) + .await; } - async fn end(&mut self, reason: WebSocketSessionEndReason) { + async fn end(&mut self, reason: MiddlewareSessionEndReason) { if let Some(transport) = self.transport.take() { let _ = tokio::time::timeout( Duration::from_millis(10), @@ -295,10 +297,10 @@ impl ChainRunner { } if let Some(denial) = denial { - end_stages(&mut stages, WebSocketSessionEndReason::MiddlewareDenial).await; + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareDenial).await; return Ok(WebSocketPreflightResult { allowed: false, - terminal_reason: Some(WebSocketSessionEndReason::MiddlewareDenial), + terminal_reason: Some(MiddlewareSessionEndReason::MiddlewareDenial), reason: middleware_denial_reason( &denial.config_name, denial.reason_code.as_deref(), @@ -315,10 +317,10 @@ impl ChainRunner { } if let Some(reason) = fail_closed_reason { - end_stages(&mut stages, WebSocketSessionEndReason::MiddlewareFailure).await; + end_stages(&mut stages, MiddlewareSessionEndReason::MiddlewareFailure).await; return Ok(WebSocketPreflightResult { allowed: false, - terminal_reason: Some(WebSocketSessionEndReason::MiddlewareFailure), + terminal_reason: Some(MiddlewareSessionEndReason::MiddlewareFailure), reason, denial: None, session: None, @@ -401,7 +403,7 @@ fn session_capacity_exhausted( .collect(); WebSocketPreflightResult { allowed: !fail_closed, - terminal_reason: fail_closed.then_some(WebSocketSessionEndReason::MiddlewareFailure), + terminal_reason: fail_closed.then_some(MiddlewareSessionEndReason::MiddlewareFailure), reason: if fail_closed { format!("middleware_failed: {reason}") } else { @@ -444,7 +446,7 @@ impl WebSocketSession { if selected_subprotocol.len() > MAX_SELECTED_SUBPROTOCOL_BYTES { return WebSocketSessionStartOutcome { allowed: false, - terminal_reason: Some(WebSocketSessionEndReason::MiddlewareFailure), + terminal_reason: Some(MiddlewareSessionEndReason::MiddlewareFailure), reason: "middleware_failed: selected_subprotocol_over_capacity".to_string(), invocations: Vec::new(), }; @@ -480,7 +482,7 @@ impl WebSocketSession { allowed: fail_closed.is_none(), terminal_reason: fail_closed .as_ref() - .map(|_| WebSocketSessionEndReason::MiddlewareFailure), + .map(|_| MiddlewareSessionEndReason::MiddlewareFailure), reason: fail_closed.unwrap_or_default(), invocations, } @@ -804,7 +806,7 @@ impl WebSocketSession { } } - pub async fn end(mut self, reason: WebSocketSessionEndReason) { + pub async fn end(mut self, reason: MiddlewareSessionEndReason) { end_stages(&mut self.stages, reason).await; self.reconcile_lifecycle(); } @@ -812,7 +814,7 @@ impl WebSocketSession { impl Drop for WebSocketSession { fn drop(&mut self) { - end_stages_now(&mut self.stages, WebSocketSessionEndReason::Cancellation); + end_stages_now(&mut self.stages, MiddlewareSessionEndReason::Cancellation); self.reconcile_lifecycle(); } } @@ -966,7 +968,7 @@ async fn open_stage(entry: DescribedChainEntry, input: WebSocketPreflightInput) }; let Some(response) = response else { let _ = sender.try_send(session_end_request( - WebSocketSessionEndReason::MiddlewareFailure, + MiddlewareSessionEndReason::MiddlewareFailure, )); return OpenStage::Failed(entry, "missing_preflight_decision".into()); }; @@ -974,7 +976,7 @@ async fn open_stage(entry: DescribedChainEntry, input: WebSocketPreflightInput) response.result else { let _ = sender.try_send(session_end_request( - WebSocketSessionEndReason::MiddlewareFailure, + MiddlewareSessionEndReason::MiddlewareFailure, )); return OpenStage::Failed(entry, "invalid_preflight_decision".into()); }; @@ -982,7 +984,7 @@ async fn open_stage(entry: DescribedChainEntry, input: WebSocketPreflightInput) Ok(decision) => decision, Err(reason) => { let _ = sender.try_send(session_end_request( - WebSocketSessionEndReason::MiddlewareFailure, + MiddlewareSessionEndReason::MiddlewareFailure, )); return OpenStage::Failed(entry, reason.into()); } @@ -1015,7 +1017,9 @@ async fn open_stage(entry: DescribedChainEntry, input: WebSocketPreflightInput) WebSocketPreflightAction::Skip => { let outcome = preflight_stage_outcome(&entry, WebSocketInvocationOutcome::Skip, decision); - let _ = sender.try_send(session_end_request(WebSocketSessionEndReason::StageSkipped)); + let _ = sender.try_send(session_end_request( + MiddlewareSessionEndReason::StageSkipped, + )); OpenStage::Skip(outcome) } WebSocketPreflightAction::Unspecified => { @@ -1299,13 +1303,13 @@ fn failure_invocation( } } -async fn end_stages(stages: &mut [WebSocketStage], reason: WebSocketSessionEndReason) { +async fn end_stages(stages: &mut [WebSocketStage], reason: MiddlewareSessionEndReason) { for stage in stages { stage.end(reason).await; } } -fn end_stages_now(stages: &mut [WebSocketStage], reason: WebSocketSessionEndReason) { +fn end_stages_now(stages: &mut [WebSocketStage], reason: MiddlewareSessionEndReason) { for stage in stages { if let Some(transport) = stage.transport.take() { let _ = transport.sender.try_send(session_end_request(reason)); @@ -1313,11 +1317,19 @@ fn end_stages_now(stages: &mut [WebSocketStage], reason: WebSocketSessionEndReas } } -fn session_end_request(reason: WebSocketSessionEndReason) -> WebSocketSessionEvent { +fn session_end_request(reason: MiddlewareSessionEndReason) -> WebSocketSessionEvent { + let protocol_error = (reason == MiddlewareSessionEndReason::ProtocolError).then_some({ + MiddlewareSessionProtocolError { + domain: Some(middleware_session_protocol_error::Domain::WebSocket( + WebSocketProtocolError {}, + )), + } + }); WebSocketSessionEvent { event: Some(web_socket_session_event::Event::SessionEnd( - WebSocketSessionEnd { + MiddlewareSessionEnd { reason: reason as i32, + protocol_error, }, )), } @@ -1327,6 +1339,30 @@ fn session_end_request(reason: WebSocketSessionEndReason) -> WebSocketSessionEve mod tests { use super::*; + #[test] + fn session_end_refines_only_protocol_errors() { + let protocol_event = session_end_request(MiddlewareSessionEndReason::ProtocolError); + let Some(web_socket_session_event::Event::SessionEnd(protocol_end)) = protocol_event.event + else { + panic!("expected session end"); + }; + assert_eq!( + MiddlewareSessionEndReason::try_from(protocol_end.reason), + Ok(MiddlewareSessionEndReason::ProtocolError) + ); + assert!(matches!( + protocol_end.protocol_error.and_then(|detail| detail.domain), + Some(middleware_session_protocol_error::Domain::WebSocket(_)) + )); + + let normal_event = session_end_request(MiddlewareSessionEndReason::Normal); + let Some(web_socket_session_event::Event::SessionEnd(normal_end)) = normal_event.event + else { + panic!("expected session end"); + }; + assert!(normal_end.protocol_error.is_none()); + } + #[test] fn protobuf_rejects_invalid_utf8_text_payload() { let encoded_text_with_invalid_utf8 = [0x12, 0x01, 0xff]; @@ -1401,8 +1437,8 @@ mod tests { panic!("disabled stage must receive session end"); }; assert_eq!( - WebSocketSessionEndReason::try_from(end.reason), - Ok(WebSocketSessionEndReason::MiddlewareFailure) + MiddlewareSessionEndReason::try_from(end.reason), + Ok(MiddlewareSessionEndReason::MiddlewareFailure) ); assert!( requests.try_recv().is_err(), diff --git a/crates/openshell-supervisor-network/src/l7/relay.rs b/crates/openshell-supervisor-network/src/l7/relay.rs index 2697fedb3c..9e4949b191 100644 --- a/crates/openshell-supervisor-network/src/l7/relay.rs +++ b/crates/openshell-supervisor-network/src/l7/relay.rs @@ -855,7 +855,7 @@ where Ok(None) => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::Cancellation) + .end(openshell_core::proto::MiddlewareSessionEndReason::Cancellation) .await; } return Ok(()); @@ -875,14 +875,14 @@ where RelayOutcome::Reusable => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected) + .end(openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure) .await; } } RelayOutcome::Consumed => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected) + .end(openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure) .await; } return Ok(()); @@ -1217,7 +1217,7 @@ where emit_policy_reload(guard, host, port, &options.policy_name); if let Some(session) = options.middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; } send_websocket_close(client, upstream, 1012).await; @@ -1594,7 +1594,7 @@ where Ok(None) => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::Cancellation) + .end(openshell_core::proto::MiddlewareSessionEndReason::Cancellation) .await; } return Ok(()); @@ -1614,14 +1614,14 @@ where RelayOutcome::Reusable => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected) + .end(openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure) .await; } } RelayOutcome::Consumed => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected) + .end(openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure) .await; } debug!( @@ -1720,7 +1720,7 @@ pub(crate) async fn finalize_websocket_pre_upgrade( emit_policy_reload(guard, host, port, policy_name); if let Some(session) = session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; } Err(error) @@ -1731,9 +1731,9 @@ pub(crate) async fn finalize_websocket_pre_upgrade( Err(error) => { let reason = if guard.is_stale() { emit_policy_reload(guard, host, port, policy_name); - openshell_core::proto::WebSocketSessionEndReason::PolicyReload + openshell_core::proto::MiddlewareSessionEndReason::PolicyReload } else { - openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected + openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure }; if let Some(session) = session.take() { session.end(reason).await; diff --git a/crates/openshell-supervisor-network/src/l7/websocket.rs b/crates/openshell-supervisor-network/src/l7/websocket.rs index 6cd8aa4818..b7d5c67dbd 100644 --- a/crates/openshell-supervisor-network/src/l7/websocket.rs +++ b/crates/openshell-supervisor-network/src/l7/websocket.rs @@ -97,7 +97,8 @@ pub enum WebSocketAssemblyAdmissionOutcome { #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum WebSocketTerminationCause { - PeerDisconnect, + DownstreamDisconnect, + UpstreamDisconnect, PolicyReload, CapacityExhausted, MiddlewareDenial, @@ -111,7 +112,7 @@ enum WebSocketTerminationCause { impl WebSocketTerminationCause { fn close_code(self) -> Option { match self { - Self::PeerDisconnect => None, + Self::DownstreamDisconnect | Self::UpstreamDisconnect => None, Self::PolicyReload => Some(1012), Self::CapacityExhausted => Some(1013), Self::MiddlewareDenial | Self::MiddlewareFailure | Self::PolicyDenial => Some(1008), @@ -121,21 +122,24 @@ impl WebSocketTerminationCause { } } - fn session_end_reason(self) -> openshell_core::proto::WebSocketSessionEndReason { + fn session_end_reason(self) -> openshell_core::proto::MiddlewareSessionEndReason { match self { - Self::PeerDisconnect => { - openshell_core::proto::WebSocketSessionEndReason::PeerDisconnect + Self::DownstreamDisconnect => { + openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect } - Self::PolicyReload => openshell_core::proto::WebSocketSessionEndReason::PolicyReload, + Self::UpstreamDisconnect => { + openshell_core::proto::MiddlewareSessionEndReason::UpstreamDisconnect + } + Self::PolicyReload => openshell_core::proto::MiddlewareSessionEndReason::PolicyReload, Self::MiddlewareDenial => { - openshell_core::proto::WebSocketSessionEndReason::MiddlewareDenial + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareDenial } - Self::PolicyDenial => openshell_core::proto::WebSocketSessionEndReason::PolicyDenial, + Self::PolicyDenial => openshell_core::proto::MiddlewareSessionEndReason::PolicyDenial, Self::CapacityExhausted | Self::MiddlewareFailure => { - openshell_core::proto::WebSocketSessionEndReason::MiddlewareFailure + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareFailure } Self::InvalidUtf8 | Self::ProtocolError | Self::MessageTooBig => { - openshell_core::proto::WebSocketSessionEndReason::ProtocolError + openshell_core::proto::MiddlewareSessionEndReason::ProtocolError } } } @@ -189,7 +193,8 @@ enum AssemblyTimeoutKind { #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum FrameErrorKind { - PeerDisconnect, + DownstreamDisconnect, + UpstreamDisconnect, Protocol(FrameFailureClass), InvalidUtf8, MessageTooBig, @@ -209,20 +214,27 @@ impl std::fmt::Display for FrameError { } impl FrameError { - fn peer_io(context: &str, error: std::io::Error) -> Self { + fn client_io(error: std::io::Error) -> Self { Self { - kind: FrameErrorKind::PeerDisconnect, - error: miette!("{context}: {error}"), + kind: FrameErrorKind::DownstreamDisconnect, + error: miette!("websocket client read failed: {error}"), } } - fn peer_disconnect(error: miette::Report) -> Self { + fn client_disconnect(error: miette::Report) -> Self { Self { - kind: FrameErrorKind::PeerDisconnect, + kind: FrameErrorKind::DownstreamDisconnect, error, } } + fn upstream_io(context: &str, error: std::io::Error) -> Self { + Self { + kind: FrameErrorKind::UpstreamDisconnect, + error: miette!("{context}: {error}"), + } + } + fn protocol(failure_class: FrameFailureClass, error: miette::Report) -> Self { Self { kind: FrameErrorKind::Protocol(failure_class), @@ -263,7 +275,12 @@ impl FrameError { impl From for WebSocketTermination { fn from(frame_error: FrameError) -> Self { let (cause, failure_class) = match frame_error.kind { - FrameErrorKind::PeerDisconnect => (WebSocketTerminationCause::PeerDisconnect, None), + FrameErrorKind::DownstreamDisconnect => { + (WebSocketTerminationCause::DownstreamDisconnect, None) + } + FrameErrorKind::UpstreamDisconnect => { + (WebSocketTerminationCause::UpstreamDisconnect, None) + } FrameErrorKind::Protocol(failure_class) => ( WebSocketTerminationCause::ProtocolError, Some(failure_class), @@ -381,7 +398,7 @@ impl TextMessageAssembly { Err(FrameError::assembly_timeout(AssemblyTimeoutKind::Idle)) } result = reader.read(buffer) => { - result.map_err(|error| FrameError::peer_io("websocket client read failed", error)) + result.map_err(FrameError::client_io) }, } } @@ -395,7 +412,7 @@ impl TextMessageAssembly { while filled < buffer.len() { let read = self.read_some(reader, &mut buffer[filled..]).await?; if read == 0 { - return Err(FrameError::peer_disconnect(miette!( + return Err(FrameError::client_disconnect(miette!( "websocket payload ended before declared length" ))); } @@ -524,22 +541,35 @@ where &mut options, ); let server_to_client = async { - tokio::io::copy(&mut upstream_read, &mut client_write) - .await - .map_err(|error| { + let mut buf = vec![0u8; COPY_BUF_SIZE]; + loop { + let read = upstream_read.read(&mut buf).await.map_err(|error| { terminate( - WebSocketTerminationCause::PeerDisconnect, - miette!("websocket upstream relay ended: {error}"), + WebSocketTerminationCause::UpstreamDisconnect, + miette!("websocket upstream read failed: {error}"), ) })?; - client_write.flush().await.map_err(|error| { - terminate( - WebSocketTerminationCause::PeerDisconnect, - miette!("websocket client relay ended: {error}"), - ) - })?; + if read == 0 { + break; + } + client_write + .write_all(&buf[..read]) + .await + .map_err(|error| { + terminate( + WebSocketTerminationCause::DownstreamDisconnect, + miette!("websocket client write failed: {error}"), + ) + })?; + client_write.flush().await.map_err(|error| { + terminate( + WebSocketTerminationCause::DownstreamDisconnect, + miette!("websocket client flush failed: {error}"), + ) + })?; + } Ok::<_, WebSocketTermination>( - openshell_core::proto::WebSocketSessionEndReason::PeerDisconnect, + openshell_core::proto::MiddlewareSessionEndReason::UpstreamDisconnect, ) }; @@ -619,7 +649,7 @@ async fn relay_client_to_server( host: &str, port: u16, options: &mut RelayOptions<'_>, -) -> WebSocketRelayResult +) -> WebSocketRelayResult where R: AsyncRead + Unpin, W: AsyncWrite + Unpin, @@ -637,9 +667,9 @@ where else { let _ = writer.shutdown().await; return Ok(if close_seen { - openshell_core::proto::WebSocketSessionEndReason::NormalClose + openshell_core::proto::MiddlewareSessionEndReason::Normal } else { - openshell_core::proto::WebSocketSessionEndReason::PeerDisconnect + openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect }); }; @@ -886,7 +916,7 @@ async fn read_exact_for_assembly( .read_exact(buffer) .await .map(|_| ()) - .map_err(|error| FrameError::peer_io("websocket client read failed", error)), + .map_err(FrameError::client_io), } } @@ -897,10 +927,7 @@ async fn read_frame_header( let mut first = [0u8; 1]; let first_read = match assembly { Some(assembly) => assembly.read_some(reader, &mut first).await, - None => reader - .read(&mut first) - .await - .map_err(|error| FrameError::peer_io("websocket client read failed", error)), + None => reader.read(&mut first).await.map_err(FrameError::client_io), }; let first = match first_read { Ok(0) => return Ok(None), @@ -1600,15 +1627,15 @@ where writer .write_all(&frame.raw_header) .await - .map_err(|error| FrameError::peer_io("websocket upstream write failed", error))?; + .map_err(|error| FrameError::upstream_io("websocket upstream write failed", error))?; writer .write_all(&raw_payload) .await - .map_err(|error| FrameError::peer_io("websocket upstream write failed", error))?; + .map_err(|error| FrameError::upstream_io("websocket upstream write failed", error))?; writer .flush() .await - .map_err(|error| FrameError::peer_io("websocket upstream flush failed", error))?; + .map_err(|error| FrameError::upstream_io("websocket upstream flush failed", error))?; Ok(()) } @@ -1654,7 +1681,7 @@ where writer .write_all(&frame.raw_header) .await - .map_err(|error| FrameError::peer_io("websocket upstream write failed", error))?; + .map_err(|error| FrameError::upstream_io("websocket upstream write failed", error))?; let mut remaining = frame.payload_len; let mut buf = [0u8; COPY_BUF_SIZE]; while remaining > 0 { @@ -1664,22 +1691,22 @@ where let n = reader .read(&mut buf[..to_read]) .await - .map_err(|error| FrameError::peer_io("websocket client read failed", error))?; + .map_err(FrameError::client_io)?; if n == 0 { - return Err(FrameError::peer_disconnect(miette!( + return Err(FrameError::client_disconnect(miette!( "websocket payload ended before declared length" ))); } writer .write_all(&buf[..n]) .await - .map_err(|error| FrameError::peer_io("websocket upstream write failed", error))?; + .map_err(|error| FrameError::upstream_io("websocket upstream write failed", error))?; remaining -= n as u64; } writer .flush() .await - .map_err(|error| FrameError::peer_io("websocket upstream flush failed", error))?; + .map_err(|error| FrameError::upstream_io("websocket upstream flush failed", error))?; Ok(()) } @@ -1762,19 +1789,19 @@ async fn write_text_frame_guarded( tokio::time::timeout(TEXT_MESSAGE_FORWARD_TOTAL_TIMEOUT, async { writer.write_all(header).await.map_err(|error| { terminate( - WebSocketTerminationCause::PeerDisconnect, + WebSocketTerminationCause::UpstreamDisconnect, miette!("websocket upstream write failed: {error}"), ) })?; writer.write_all(payload).await.map_err(|error| { terminate( - WebSocketTerminationCause::PeerDisconnect, + WebSocketTerminationCause::UpstreamDisconnect, miette!("websocket upstream write failed: {error}"), ) })?; writer.flush().await.map_err(|error| { terminate( - WebSocketTerminationCause::PeerDisconnect, + WebSocketTerminationCause::UpstreamDisconnect, miette!("websocket upstream flush failed: {error}"), ) }) @@ -1782,7 +1809,7 @@ async fn write_text_frame_guarded( .await .map_err(|_| { terminate( - WebSocketTerminationCause::PeerDisconnect, + WebSocketTerminationCause::UpstreamDisconnect, miette!("websocket upstream forwarding total timeout"), ) })??; @@ -2216,7 +2243,7 @@ network_policies: #[test] fn termination_causes_map_to_protocol_close_codes_and_session_reasons() { - use openshell_core::proto::WebSocketSessionEndReason as EndReason; + use openshell_core::proto::MiddlewareSessionEndReason as EndReason; assert_eq!( WebSocketTerminationCause::InvalidUtf8.close_code(), @@ -2238,7 +2265,14 @@ network_policies: WebSocketTerminationCause::PolicyReload.close_code(), Some(1012) ); - assert_eq!(WebSocketTerminationCause::PeerDisconnect.close_code(), None); + assert_eq!( + WebSocketTerminationCause::DownstreamDisconnect.close_code(), + None + ); + assert_eq!( + WebSocketTerminationCause::UpstreamDisconnect.close_code(), + None + ); assert_eq!( WebSocketTerminationCause::InvalidUtf8.session_end_reason(), @@ -2260,6 +2294,14 @@ network_policies: WebSocketTerminationCause::PolicyReload.session_end_reason(), EndReason::PolicyReload ); + assert_eq!( + WebSocketTerminationCause::DownstreamDisconnect.session_end_reason(), + EndReason::DownstreamDisconnect + ); + assert_eq!( + WebSocketTerminationCause::UpstreamDisconnect.session_end_reason(), + EndReason::UpstreamDisconnect + ); } fn resolver() -> (HashMap, SecretResolver) { @@ -2603,7 +2645,7 @@ network_policies: .await .expect("join forwarding") .expect_err("non-reading upstream must time out"); - assert_eq!(error.cause, WebSocketTerminationCause::PeerDisconnect); + assert_eq!(error.cause, WebSocketTerminationCause::UpstreamDisconnect); assert!( error .error @@ -3486,7 +3528,7 @@ network_policies: enum ObservedWebSocketRequest { SessionStart, Message { sequence: u64, payload: String }, - SessionEnd(openshell_core::proto::WebSocketSessionEndReason), + SessionEnd(openshell_core::proto::MiddlewareSessionEndReason), } #[derive(Clone, Default)] @@ -3648,7 +3690,7 @@ network_policies: Some(web_socket_session_event::Event::SessionEnd(end)) => { if let Some(observed) = &observed && let Ok(reason) = - openshell_core::proto::WebSocketSessionEndReason::try_from( + openshell_core::proto::MiddlewareSessionEndReason::try_from( end.reason, ) { @@ -3898,7 +3940,7 @@ network_policies: assert!(matches!( observed.recv().await, Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::ProtocolError, + openshell_core::proto::MiddlewareSessionEndReason::ProtocolError, )) )); assert!( @@ -4018,7 +4060,7 @@ network_policies: assert!(matches!( observed.recv().await, Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::PeerDisconnect, + openshell_core::proto::MiddlewareSessionEndReason::DownstreamDisconnect, )) )); @@ -4056,7 +4098,7 @@ network_policies: assert!(matches!( observed.recv().await, Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::MiddlewareFailure, + openshell_core::proto::MiddlewareSessionEndReason::MiddlewareFailure, )) )); assert!( @@ -4258,7 +4300,7 @@ network_policies: assert!(error.to_string().contains("policy generation is stale")); match observed.recv().await { Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::PolicyReload, + openshell_core::proto::MiddlewareSessionEndReason::PolicyReload, )) => {} Some(ObservedWebSocketRequest::Message { payload, .. }) => { panic!("stale message leaked {} bytes to middleware", payload.len()); @@ -4383,7 +4425,7 @@ network_policies: } assert_eq!( end_reason, - Some(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + Some(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) ); let _ = shutdown_tx.send(()); @@ -4423,7 +4465,7 @@ network_policies: assert!(matches!( observed.recv().await, Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::PolicyReload + openshell_core::proto::MiddlewareSessionEndReason::PolicyReload )) )); @@ -4527,7 +4569,7 @@ network_policies: assert!(matches!( observed.recv().await, Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::PolicyReload + openshell_core::proto::MiddlewareSessionEndReason::PolicyReload )) )); @@ -4637,7 +4679,7 @@ network_policies: ); match observed.recv().await { Some(ObservedWebSocketRequest::SessionEnd( - openshell_core::proto::WebSocketSessionEndReason::PolicyDenial, + openshell_core::proto::MiddlewareSessionEndReason::PolicyDenial, )) => {} Some(ObservedWebSocketRequest::Message { payload, .. }) => { panic!( diff --git a/crates/openshell-supervisor-network/src/opa.rs b/crates/openshell-supervisor-network/src/opa.rs index 63aa2c2c70..60fb96cb0a 100644 --- a/crates/openshell-supervisor-network/src/opa.rs +++ b/crates/openshell-supervisor-network/src/opa.rs @@ -7999,7 +7999,7 @@ network_policies: old_sessions .pop() .expect("old-generation session") - .end(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; let admitted = current_runner .preflight_websocket( diff --git a/crates/openshell-supervisor-network/src/proxy.rs b/crates/openshell-supervisor-network/src/proxy.rs index dc2736a4ea..177d640fd8 100644 --- a/crates/openshell-supervisor-network/src/proxy.rs +++ b/crates/openshell-supervisor-network/src/proxy.rs @@ -5809,7 +5809,7 @@ async fn handle_forward_proxy( ); if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::Cancellation) + .end(openshell_core::proto::MiddlewareSessionEndReason::Cancellation) .await; } respond( @@ -5869,7 +5869,7 @@ async fn handle_forward_proxy( ); if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::Cancellation) + .end(openshell_core::proto::MiddlewareSessionEndReason::Cancellation) .await; } if e.is_endpoint_mismatch() { @@ -5912,7 +5912,7 @@ async fn handle_forward_proxy( emit_l7_tunnel_close_after_policy_change(&host_lc, port, e); if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; } respond( @@ -5957,7 +5957,7 @@ async fn handle_forward_proxy( ocsf_emit!(event); if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected) + .end(openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure) .await; } respond( @@ -5986,7 +5986,7 @@ async fn handle_forward_proxy( emit_l7_tunnel_close_after_policy_change(&host_lc, port, e); if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::PolicyReload) + .end(openshell_core::proto::MiddlewareSessionEndReason::PolicyReload) .await; } respond( @@ -6039,7 +6039,7 @@ async fn handle_forward_proxy( if let Some(error) = report.downcast_ref::() { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::Cancellation) + .end(openshell_core::proto::MiddlewareSessionEndReason::Cancellation) .await; } crate::l7::relay::reject_credential_resolution(client, &l7_ctx, error).await?; @@ -6081,7 +6081,7 @@ async fn handle_forward_proxy( | crate::l7::provider::RelayOutcome::Consumed => { if let Some(session) = middleware_session.take() { session - .end(openshell_core::proto::WebSocketSessionEndReason::UpstreamRejected) + .end(openshell_core::proto::MiddlewareSessionEndReason::UpstreamFailure) .await; } } diff --git a/examples/supervisor-middleware-content-guard/Cargo.lock b/examples/supervisor-middleware-content-guard/Cargo.lock index 9bb8d1feee..f31d5be9b5 100644 --- a/examples/supervisor-middleware-content-guard/Cargo.lock +++ b/examples/supervisor-middleware-content-guard/Cargo.lock @@ -853,6 +853,7 @@ dependencies = [ "prost", "prost-types", "protoc-bin-vendored", + "rustix", "serde", "serde_json", "thiserror", diff --git a/examples/supervisor-middleware-content-guard/src/main.rs b/examples/supervisor-middleware-content-guard/src/main.rs index c527537e74..8d714264e7 100644 --- a/examples/supervisor-middleware-content-guard/src/main.rs +++ b/examples/supervisor-middleware-content-guard/src/main.rs @@ -478,7 +478,7 @@ async fn main() -> Result<(), Box> { #[cfg(test)] mod tests { use super::*; - use openshell_core::proto::{WebSocketPreflight, WebSocketSessionEnd, WebSocketSessionStart}; + use openshell_core::proto::{MiddlewareSessionEnd, WebSocketPreflight, WebSocketSessionStart}; use prost_types::{ListValue, Value}; use std::collections::BTreeMap; @@ -550,7 +550,7 @@ mod tests { )), })), event(web_socket_session_event::Event::SessionEnd( - WebSocketSessionEnd::default(), + MiddlewareSessionEnd::default(), )), ]); let mut results = ContentGuard::websocket_stream(events); diff --git a/proto/supervisor_middleware.proto b/proto/supervisor_middleware.proto index bcd9c8ddb3..f1389acaf1 100644 --- a/proto/supervisor_middleware.proto +++ b/proto/supervisor_middleware.proto @@ -8,9 +8,9 @@ package openshell.middleware.v1; import "google/protobuf/empty.proto"; import "google/protobuf/struct.proto"; -// SupervisorMiddleware lets an operator-run service inspect and transform -// sandbox HTTP requests and client WebSocket text messages before OpenShell -// injects credentials. +// SupervisorMiddleware discovers and configures one operator-run middleware. +// It evaluates HTTP requests and WebSocket messages before credentials. +// Phase-specific services share the same registration. service SupervisorMiddleware { // Describe returns the service manifest and declared bindings. rpc Describe(google.protobuf.Empty) returns (MiddlewareManifest); @@ -33,9 +33,16 @@ service SupervisorMiddleware { returns (stream WebSocketSessionEventResult); } -// MiddlewareManifest describes one middleware service and the bindings it -// exposes. The service is the operator-run gRPC server implementing -// SupervisorMiddleware. +// HttpResponsePreReturn evaluates one response for one middleware stage before +// OpenShell returns it to the sandbox. +service HttpResponsePreReturn { + // Evaluate starts with preflight, followed by selected body units. It may end + // with one best-effort session_end. + rpc Evaluate(stream HttpResponseEvent) + returns (stream HttpResponseEventResult); +} + +// MiddlewareManifest describes one middleware service and its bindings. message MiddlewareManifest { // Human-readable middleware service name used only for diagnostics. This is // not required to match an operator-owned registration name. @@ -57,13 +64,11 @@ message MiddlewareManifest { message MiddlewareBinding { // Supported operation. SupervisorMiddlewareOperation operation = 1; - // Supported evaluation phase. PR 1 supports PRE_CREDENTIALS. PRE_RETURN is - // reserved for the return-path follow-up and is rejected by current - // manifest validation. + // Supported phase. Current manifest validation accepts PRE_CREDENTIALS only; + // PRE_RETURN remains unavailable until the matching relay paths dispatch it. SupervisorMiddlewarePhase phase = 2; - // Maximum logical payload or replacement this binding can process. For - // HTTP_REQUEST this is the request body; for WEBSOCKET_MESSAGE this is one - // complete message. Required for every payload-bearing operation. + // Maximum request body, WebSocket message, or response body unit/replacement. + // Required for payload-bearing operations. uint64 max_payload_bytes = 3; // Optional binding-specific RPC timeout. Empty uses the operator-configured // service timeout, or the 500ms platform default when that is also omitted. @@ -113,7 +118,7 @@ message HttpRequestEvaluation { string middleware_name = 7; } -// HttpHeader is one request header line. +// HttpHeader is one HTTP header line. message HttpHeader { // Lowercased header name. string name = 1; @@ -121,11 +126,257 @@ message HttpHeader { string value = 2; } +// One ordered response event. A stream starts with preflight, may continue with +// body units, and may end with one best-effort session_end. +message HttpResponseEvent { + oneof event { + // Initial response head and request context. + HttpResponsePreflight preflight = 1; + // Next normalized body unit. + HttpResponseBodyUnit body = 2; + // Optional terminal notification. + MiddlewareSessionEnd session_end = 3; + } +} + +// Each preflight and body event requires one ordered result. session_end has no +// result. +message HttpResponseEventResult { + oneof result { + // Result for preflight. + HttpResponsePreflightDecision preflight_decision = 1; + // Result for the next body unit. + HttpResponseBodyResult body_result = 2; + } +} + +// HttpResponsePreflight exposes the current final response head to one stage. +message HttpResponsePreflight { + // Request identity. request_id links request and response evaluations. + // Limited to 4 KiB encoded. + RequestContext context = 1; + // Admitted request target with a redacted query. Limited to 32 KiB encoded. + HttpRequestTarget target = 2; + // Final non-informational upstream status. Upgrades are not evaluated. + uint32 status_code = 3; + // Response headers after prior stages, in wire order. Repeated names remain + // separate. Credential, routing, and hop-by-hop headers are omitted. + // Content-Length, Content-Encoding, and Content-Range retain their read-only + // upstream values. OpenShell may recompute or remove Content-Length later. + // Limited to 128 lines and 64 KiB encoded. + repeated HttpHeader headers = 4; + // Built-in middleware name or operator-owned registration name. + string middleware_name = 5; + // Validated service configuration. Limited to 64 KiB encoded. + google.protobuf.Struct config = 6; + // Effective minimum of platform, registration, and binding limits. Applies to + // whole-body input/replacement and each stream replacement. Stream inputs use + // at most half this limit. + uint64 max_payload_bytes = 7; + // Modes computed once from the original response head. HEADERS_ONLY is always + // present and is the only mode for bodyless, partial, encoded, or no-transform + // responses. Oversized or open-ended responses omit WHOLE_BODY_BYTES. + // Selecting an unlisted mode fails according to on_error. + repeated HttpResponseBodyMode permitted_body_modes = 8; + // Allows STREAM_BYTES replacements to defer bytes across units. Set only for + // fail-closed stages; fail-open stages cannot defer. + bool deferral_permitted = 9; +} + +// Selects skip, inspect, or block. Diagnostic fields apply to every action. +message HttpResponsePreflightDecision { + oneof action { + // Deliver unchanged without invoking on_error. + HttpResponsePreflightSkip skip = 1; + // Inspect with the selected body mode and mutations. + HttpResponsePreflightInspect inspect = 2; + // Prevent delivery to the sandbox. + HttpResponseBlockDelivery block_delivery = 7; + } + // Service diagnostic, never sent to the sandbox or security logs. Maximum + // 4 KiB. + string reason = 3; + // Optional audit code using HttpRequestResult.reason_code format. Returned to + // the sandbox only for block_delivery. + string reason_code = 4; + // Up to 32 audit-safe findings, each limited to 4 KiB encoded. + repeated Finding findings = 5; + // Non-secret diagnostic metadata, limited to 64 entries and 32 KiB. + map metadata = 6; +} + +// Ends this stage successfully without body inspection. +message HttpResponsePreflightSkip {} + +// Blocks delivery as a successful decision regardless of on_error. Preflight +// and WHOLE_BODY_BYTES blocks replace the uncommitted response with a platform +// error. STREAM_BYTES blocks abort delivery after any sent prefix. The upstream +// request is not undone. A block wins over other actions and failures, stops +// later evaluation, and ends opened stages with MIDDLEWARE_DENIAL. +message HttpResponseBlockDelivery {} + +// Selects body inspection and response-header mutations. +message HttpResponsePreflightInspect { + // Required mode from permitted_body_modes. Invalid values fail according to + // on_error. + HttpResponseBodyMode body_mode = 1; + // Ordered mutations applied atomically before the next stage. Only visible + // end-to-end headers may change. Routing, credential, framing, coding, range, + // and hop-by-hop headers are protected; integrity headers may only be removed. + // Limited to 64 operations, 32 KiB of name/value data, and 64 KiB encoded. + repeated HeaderMutation header_mutations = 2; +} + +// Controls which response-body units a stage receives. +enum HttpResponseBodyMode { + // Invalid value handled according to on_error. + HTTP_RESPONSE_BODY_MODE_UNSPECIFIED = 0; + // Inspect only the response head. + HTTP_RESPONSE_BODY_MODE_HEADERS_ONLY = 1; + // Buffer the normalized body as one final unit before committing the head. + // Input and replacement must fit max_payload_bytes. Capacity and deadline + // failures use whole_body_over_capacity and whole_body_accumulation_timeout. + HTTP_RESPONSE_BODY_MODE_WHOLE_BODY_BYTES = 2; + // Receive normalized units ending with end_of_stream. Each replacement must + // fit max_payload_bytes. The full body may exceed it and has no accumulation + // deadline. + HTTP_RESPONSE_BODY_MODE_STREAM_BYTES = 3; +} + +// One normalized body unit. Boundaries have no transport or application +// meaning. +message HttpResponseBodyUnit { + // Contiguous and stage-local, starting at 1. + uint64 sequence = 1; + oneof payload { + // Bytes without transfer framing. WHOLE_BODY_BYTES represents an empty body + // with present empty data. STREAM_BYTES units are at most the smaller of + // 64 KiB and half max_payload_bytes, and may be shorter to preserve flushing. + bytes data = 2; + } + // Marks the final unit. A completed inspection receives exactly one, including + // an empty sequence-1 unit for an empty body. OpenShell does not read ahead, + // so it may send an empty final unit after the last data unit. The final unit + // requires a result containing any deferred bytes. Interrupted streams may + // end with session_end instead. + bool end_of_stream = 3; +} + +// Result for one body unit. Units are processed in lockstep; V1 does not +// support ownership transfer. +message HttpResponseBodyResult { + // Must match the next unit. Zero, gaps, duplicates, and regressions fail. + uint64 sequence = 1; + // Exactly one explicit action is required. + oneof action { + // Forward the input unit unchanged. + HttpResponseBodyPassThrough pass_through = 2; + // Replace the complete input unit. + HttpResponseBodyTransform transform = 3; + // Stop delivery. See HttpResponseBlockDelivery. + HttpResponseBlockDelivery block_delivery = 8; + // Finalize this unit and stop inspecting. + HttpResponseBodySkipRemaining skip_remaining = 9; + } + // Service diagnostic, never sent to the sandbox or security logs. Maximum + // 4 KiB. + string reason = 4; + // Optional audit code using preflight reason_code format. Never sent to the + // sandbox. + string reason_code = 5; + // Up to 32 audit-safe findings, each limited to 4 KiB encoded. + repeated Finding findings = 6; + // Non-secret diagnostic metadata, limited to 64 entries and 32 KiB. + map metadata = 7; +} + +// Preserves the input unit. +message HttpResponseBodyPassThrough {} + +// Finalizes this unit and ends the stage. Later units bypass it but continue +// through other stages. This stage receives no end_of_stream. A deferred tail +// must be in transform. For WHOLE_BODY_BYTES, this equals its nested action. +message HttpResponseBodySkipRemaining { + // Exactly one action for the current unit. + oneof current { + // Forward the current unit unchanged. + HttpResponseBodyPassThrough pass_through = 1; + // Replace the current unit and any deferred bytes. + HttpResponseBodyTransform transform = 2; + } +} + +// Replaces the complete input unit. +message HttpResponseBodyTransform { + // Required replacement, limited to max_payload_bytes. Present empty data + // deletes the input unit. When deferral_permitted, up to half the limit may be + // held for a later replacement; otherwise this must account for all input. + oneof replacement { + // Normalized replacement bytes. + bytes data = 1; + } +} + +// Stable reason OpenShell ended a middleware stage stream. +enum MiddlewareSessionEndReason { + // Invalid reason. + MIDDLEWARE_SESSION_END_REASON_UNSPECIFIED = 0; + // Evaluation completed. + MIDDLEWARE_SESSION_END_REASON_NORMAL = 1; + // The sandbox peer disconnected. + MIDDLEWARE_SESSION_END_REASON_DOWNSTREAM_DISCONNECT = 2; + // A policy reload replaced the active middleware chain. + MIDDLEWARE_SESSION_END_REASON_POLICY_RELOAD = 3; + // A stage denied the operation or blocked the response. + MIDDLEWARE_SESSION_END_REASON_MIDDLEWARE_DENIAL = 4; + // A selected stage failed. + MIDDLEWARE_SESSION_END_REASON_MIDDLEWARE_FAILURE = 5; + // A proxied or middleware protocol was violated. + MIDDLEWARE_SESSION_END_REASON_PROTOCOL_ERROR = 6; + // Evaluation was canceled for another reason. + MIDDLEWARE_SESSION_END_REASON_CANCELLATION = 7; + // Upstream rejected or failed before a valid response or upgrade. + MIDDLEWARE_SESSION_END_REASON_UPSTREAM_FAILURE = 8; + // Network policy denied the operation. + MIDDLEWARE_SESSION_END_REASON_POLICY_DENIAL = 9; + // The stage successfully declined inspection during preflight. + MIDDLEWARE_SESSION_END_REASON_STAGE_SKIPPED = 10; + // Upstream disconnected after a valid response or upgrade. + MIDDLEWARE_SESSION_END_REASON_UPSTREAM_DISCONNECT = 11; +} + +// Best-effort terminal notification. A stage receives at most one and sends no +// result. +message MiddlewareSessionEnd { + // Terminal reason. Producers never send UNSPECIFIED. + MiddlewareSessionEndReason reason = 1; + // Set only for PROTOCOL_ERROR. Missing or unknown details mean a generic + // protocol error. + MiddlewareSessionProtocolError protocol_error = 2; +} + +// Details for a protocol-error session end. +message MiddlewareSessionProtocolError { + oneof domain { + // WebSocket protocol violation. + WebSocketProtocolError web_socket = 1; + // Middleware event/result protocol violation. + MiddlewareExchangeProtocolError middleware_exchange = 2; + } +} + +// WebSocket protocol error details, reserved for future categories. +message WebSocketProtocolError {} + +// Middleware exchange error details, reserved for future categories. +message MiddlewareExchangeProtocolError {} + // Supervisor operation selected for middleware evaluation. enum SupervisorMiddlewareOperation { SUPERVISOR_MIDDLEWARE_OPERATION_UNSPECIFIED = 0; SUPERVISOR_MIDDLEWARE_OPERATION_HTTP_REQUEST = 1; SUPERVISOR_MIDDLEWARE_OPERATION_WEBSOCKET_MESSAGE = 2; + SUPERVISOR_MIDDLEWARE_OPERATION_HTTP_RESPONSE = 3; } // Ordered phase within a supervisor operation. @@ -135,24 +386,6 @@ enum SupervisorMiddlewarePhase { SUPERVISOR_MIDDLEWARE_PHASE_PRE_RETURN = 2; } -// Why OpenShell is ending a middleware stream. -enum WebSocketSessionEndReason { - WEB_SOCKET_SESSION_END_REASON_UNSPECIFIED = 0; - WEB_SOCKET_SESSION_END_REASON_NORMAL_CLOSE = 1; - WEB_SOCKET_SESSION_END_REASON_PEER_DISCONNECT = 2; - WEB_SOCKET_SESSION_END_REASON_POLICY_RELOAD = 3; - WEB_SOCKET_SESSION_END_REASON_MIDDLEWARE_DENIAL = 4; - WEB_SOCKET_SESSION_END_REASON_MIDDLEWARE_FAILURE = 5; - WEB_SOCKET_SESSION_END_REASON_PROTOCOL_ERROR = 6; - WEB_SOCKET_SESSION_END_REASON_CANCELLATION = 7; - WEB_SOCKET_SESSION_END_REASON_UPSTREAM_REJECTED = 8; - WEB_SOCKET_SESSION_END_REASON_POLICY_DENIAL = 9; - // The middleware stage voluntarily declined inspection during preflight. - // This is a successful stage-local outcome, not a cancellation or denial of - // the WebSocket upgrade. - WEB_SOCKET_SESSION_END_REASON_STAGE_SKIPPED = 10; -} - // WebSocketSessionEvent is one ordered event in a stage-local stream. // Message sequence numbers identify logical messages session-wide. A stage // receives a strictly increasing subset of those numbers; gaps are valid when @@ -162,7 +395,7 @@ message WebSocketSessionEvent { WebSocketPreflight preflight = 1; WebSocketSessionStart session_start = 2; WebSocketMessage message = 3; - WebSocketSessionEnd session_end = 4; + MiddlewareSessionEnd session_end = 4; } } @@ -203,12 +436,6 @@ message WebSocketMessage { } } -// WebSocketSessionEnd is OpenShell's best-effort terminal notification for one -// opened stage stream. A stage receives at most one such notification. -message WebSocketSessionEnd { - WebSocketSessionEndReason reason = 1; -} - // WebSocketPreflightAction is the service's one-time scoping decision. enum WebSocketPreflightAction { // Invalid response value handled according to the policy failure mode. diff --git a/tasks/rust.toml b/tasks/rust.toml index 854c2ac939..5c95660f72 100644 --- a/tasks/rust.toml +++ b/tasks/rust.toml @@ -15,6 +15,7 @@ run = [ "cargo clippy --workspace --all-targets -- -D warnings", "cargo clippy --manifest-path e2e/rust/Cargo.toml --all-targets -- -D warnings", "cargo check --manifest-path examples/governance-interceptor/Cargo.toml --all-targets", + "cargo check --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml --all-targets", ] run_windows = "powershell -NoProfile -ExecutionPolicy Bypass -File tasks/scripts/windows-msvc.ps1 lint native" hide = true @@ -25,6 +26,7 @@ run = [ "cargo fmt --all", "cargo fmt --manifest-path e2e/rust/Cargo.toml --all", "cargo fmt --manifest-path examples/governance-interceptor/Cargo.toml --all", + "cargo fmt --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml --all", ] hide = true @@ -34,6 +36,7 @@ run = [ "cargo fmt --all -- --check", "cargo fmt --manifest-path e2e/rust/Cargo.toml --all -- --check", "cargo fmt --manifest-path examples/governance-interceptor/Cargo.toml --all -- --check", + "cargo fmt --manifest-path examples/supervisor-middleware-content-guard/Cargo.toml --all -- --check", ] hide = true