Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .github/workflows/branch-checks.yml
Original file line number Diff line number Diff line change
Expand Up @@ -158,12 +158,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:
Expand Down
33 changes: 30 additions & 3 deletions crates/openshell-core/src/middleware.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<dyn tokio_stream::Stream<Item = Result<HttpResponseEventResult, Status>> + Send + 'static>,
>;

/// Transport-neutral response stream for one WebSocket middleware stage.
pub type WebSocketResponseStream = Pin<
Box<
Expand Down Expand Up @@ -47,6 +53,15 @@ pub trait SupervisorMiddlewareEndpoint: Send + Sync {
&self,
requests: mpsc::Receiver<WebSocketSessionEvent>,
) -> Result<WebSocketResponseStream, Status>;

async fn open_http_response_pre_return(
&self,
_requests: mpsc::Receiver<HttpResponseEvent>,
) -> Result<HttpResponseResultStream, Status> {
Err(Status::unimplemented(
"middleware does not implement HTTP response pre-return evaluation",
))
}
}

/// Borrowed request state exposed to one in-process middleware invocation.
Expand Down Expand Up @@ -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<HttpResponseEvent>,
) -> std::result::Result<HttpResponseResultStream, Status> {
Err(Status::unimplemented(
"middleware does not implement HTTP response pre-return evaluation",
))
}
}

/// Default timeout for one supervisor middleware RPC.
Expand Down
94 changes: 73 additions & 21 deletions crates/openshell-supervisor-middleware/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<EvaluateWebSocketSessionStream = WebSocketResponseStream>;
Expand Down Expand Up @@ -180,6 +181,13 @@ impl InProcessMiddleware for EndpointInProcessAdapter {
) -> std::result::Result<WebSocketResponseStream, tonic::Status> {
self.endpoint.open_websocket_session(requests).await
}

async fn open_http_response_pre_return(
&self,
requests: tokio::sync::mpsc::Receiver<openshell_core::proto::HttpResponseEvent>,
) -> std::result::Result<HttpResponseResultStream, tonic::Status> {
self.endpoint.open_http_response_pre_return(requests).await
}
}

/// Adapt a transport-neutral endpoint to the in-process registry contract.
Expand Down Expand Up @@ -835,6 +843,12 @@ fn supported_binding(source: &str, binding: &MiddlewareBinding) -> Result<Suppor
Some(SupervisorMiddlewareOperation::HttpRequest),
Some(SupervisorMiddlewarePhase::PreCredentials),
) => Ok(SupportedBinding::HttpPreCredentials),
(
Some(SupervisorMiddlewareOperation::HttpResponse),
Comment thread
pimlock marked this conversation as resolved.
Some(SupervisorMiddlewarePhase::PreReturn),
) => Err(miette!(
"{source} advertises HTTP_RESPONSE/PRE_RETURN, which is not yet supported"
)),
(
Some(SupervisorMiddlewareOperation::WebsocketMessage),
Some(SupervisorMiddlewarePhase::PreCredentials),
Expand All @@ -843,7 +857,7 @@ fn supported_binding(source: &str, binding: &MiddlewareBinding) -> Result<Suppor
Some(SupervisorMiddlewareOperation::WebsocketMessage),
Some(SupervisorMiddlewarePhase::PreReturn),
) => 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"
Expand Down Expand Up @@ -1435,6 +1449,20 @@ impl ChainRunner {
.entries)
}

pub async fn describe_http_response_chain(
&self,
entries: &[ChainEntry],
) -> Result<Vec<DescribedChainEntry>> {
Ok(self
.describe_chain_for(
entries,
SupervisorMiddlewareOperation::HttpResponse,
SupervisorMiddlewarePhase::PreReturn,
)
.await?
.entries)
}

async fn describe_chain_for(
&self,
entries: &[ChainEntry],
Expand Down Expand Up @@ -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(&registration, &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 {
Expand All @@ -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]
Expand Down Expand Up @@ -4713,7 +4765,7 @@ mod tests {
close_on_first_message: bool,
messages: Arc<std::sync::atomic::AtomicUsize>,
session_ends: Option<
tokio::sync::mpsc::UnboundedSender<openshell_core::proto::WebSocketSessionEndReason>,
tokio::sync::mpsc::UnboundedSender<openshell_core::proto::MiddlewareSessionEndReason>,
>,
}

Expand Down Expand Up @@ -4804,7 +4856,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,
)
{
Expand Down Expand Up @@ -5089,14 +5141,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();
Expand Down Expand Up @@ -5132,7 +5184,7 @@ mod tests {
assert!(!outcome.allowed);
assert_eq!(
outcome.terminal_reason,
Some(WebSocketSessionEndReason::MiddlewareDenial)
Some(MiddlewareSessionEndReason::MiddlewareDenial)
);
assert_eq!(
outcome.reason,
Expand Down Expand Up @@ -5170,7 +5222,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(),
Expand All @@ -5181,7 +5233,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();
Expand Down Expand Up @@ -5223,7 +5275,7 @@ mod tests {
assert!(!outcome.allowed);
assert_eq!(
outcome.terminal_reason,
Some(WebSocketSessionEndReason::MiddlewareDenial)
Some(MiddlewareSessionEndReason::MiddlewareDenial)
);
assert_eq!(
outcome
Expand All @@ -5240,7 +5292,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(),
Expand Down Expand Up @@ -5317,7 +5369,7 @@ mod tests {
"middleware_failed: request_message_over_capacity"
);
session
.end(openshell_core::proto::WebSocketSessionEndReason::NormalClose)
.end(openshell_core::proto::MiddlewareSessionEndReason::Normal)
.await;
}

Expand Down Expand Up @@ -5374,7 +5426,7 @@ mod tests {
assert!(!redacted.invocations[0].stage_disabled);

session
.end(openshell_core::proto::WebSocketSessionEndReason::NormalClose)
.end(openshell_core::proto::MiddlewareSessionEndReason::Normal)
.await;
}

Expand Down Expand Up @@ -5420,7 +5472,7 @@ mod tests {
);
assert!(outcome.invocations[0].transformed);
session
.end(openshell_core::proto::WebSocketSessionEndReason::NormalClose)
.end(openshell_core::proto::MiddlewareSessionEndReason::Normal)
.await;
}

Expand Down Expand Up @@ -5506,7 +5558,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
Expand Down Expand Up @@ -5617,7 +5669,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
Expand Down Expand Up @@ -5879,7 +5931,7 @@ mod tests {
tokio::time::timeout(Duration::from_secs(1), session_ends_rx.recv())
.await
.expect("skipped stage must receive session_end"),
Some(openshell_core::proto::WebSocketSessionEndReason::StageSkipped)
Some(openshell_core::proto::MiddlewareSessionEndReason::StageSkipped)
);
assert!(
session_ends_rx.try_recv().is_err(),
Expand Down Expand Up @@ -5925,7 +5977,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);

Expand Down Expand Up @@ -5967,7 +6019,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"))
Expand Down
28 changes: 24 additions & 4 deletions crates/openshell-supervisor-middleware/src/remote.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -106,6 +108,7 @@ impl GrpcMiddlewareService {
#[derive(Clone)]
pub struct RemoteMiddlewareService {
client: SupervisorMiddlewareClient<ExtensionChannel>,
response_client: HttpResponsePreReturnClient<ExtensionChannel>,
}

impl RemoteMiddlewareService {
Expand Down Expand Up @@ -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),
})
Expand Down Expand Up @@ -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<HttpResponseEvent>,
) -> std::result::Result<HttpResponseResultStream, Status> {
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))
}
}
Loading
Loading