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: 1 addition & 1 deletion crates/openshell-supervisor-network/src/l7/middleware.rs
Original file line number Diff line number Diff line change
Expand Up @@ -521,7 +521,7 @@ pub async fn apply_middleware_chain_with_request_id<C: AsyncRead + AsyncWrite +
req,
client,
ctx,
"https",
ctx.request_scheme(),
chain,
runner,
generation_guard,
Expand Down
126 changes: 123 additions & 3 deletions crates/openshell-supervisor-network/src/l7/relay.rs
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,23 @@ pub struct L7EvalContext {
pub(crate) endpoint_observation_tx: Option<EndpointObservationSender>,
}

impl L7EvalContext {
/// Whether the inspected transport is plaintext rather than TLS-terminated.
fn is_plaintext(&self) -> bool {
self.request_default_port == Some(80)
}

/// HTTP scheme of the inspected transport.
pub(crate) fn request_scheme(&self) -> &'static str {
if self.is_plaintext() { "http" } else { "https" }
}

/// WebSocket scheme of the inspected transport.
pub(crate) fn websocket_scheme(&self) -> &'static str {
if self.is_plaintext() { "ws" } else { "wss" }
}
}

fn request_default_port(ctx: &L7EvalContext) -> Option<u16> {
ctx.request_default_port
}
Expand Down Expand Up @@ -707,7 +724,7 @@ where
Some(http_response_middleware_relay(
&request,
ctx,
"https",
ctx.request_scheme(),
request_id,
response_chain,
engine.middleware_runner(),
Expand Down Expand Up @@ -1445,7 +1462,7 @@ where
chain,
engine.middleware_runner(),
ctx,
"wss",
ctx.websocket_scheme(),
)
.await;
let preflight = match preflight {
Expand Down Expand Up @@ -2187,7 +2204,7 @@ where
chain,
engine.middleware_runner(),
ctx,
"wss",
ctx.websocket_scheme(),
)
.await;
let preflight = match preflight {
Expand Down Expand Up @@ -7308,6 +7325,109 @@ network_policies:
}
}

#[test]
fn inspected_schemes_follow_transport() {
for (port, http, ws) in [
(Some(80), "http", "ws"),
(Some(443), "https", "wss"),
(None, "https", "wss"),
] {
let ctx = L7EvalContext {
request_default_port: port,
..Default::default()
};
assert_eq!(ctx.request_scheme(), http);
assert_eq!(ctx.websocket_scheme(), ws);
}
}

/// Records the scheme each request reaches middleware with.
struct SchemeRecorder(Arc<std::sync::Mutex<Vec<String>>>);

#[tonic::async_trait]
impl openshell_core::middleware::InProcessMiddleware for SchemeRecorder {
async fn describe(&self) -> openshell_core::proto::MiddlewareManifest {
LimitService {
name: "test/scheme",
max_body_bytes: 1024,
replacement: None,
}
.describe()
.await
}

async fn validate_config(
&self,
_middleware_name: &str,
_config: &prost_types::Struct,
) -> Result<()> {
Ok(())
}

async fn evaluate_http_request(
&self,
request: openshell_core::middleware::HttpRequestView<'_>,
) -> Result<openshell_core::proto::HttpRequestResult> {
self.0.lock().unwrap().push(request.target().scheme.clone());
Ok(openshell_core::proto::HttpRequestResult {
decision: openshell_core::proto::Decision::Allow as i32,
..Default::default()
})
}
}

#[tokio::test]
async fn middleware_request_scheme_matches_inspected_transport() {
use openshell_supervisor_middleware::{ChainEntry, ChainRunner, OnError};

let (_config, tunnel_engine, ctx) =
middleware_relay_context("openshell/regex", "fail_closed");
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let registry = openshell_supervisor_middleware::MiddlewareRegistry::connect_services(
vec![Arc::new(SchemeRecorder(seen.clone()))],
Vec::new(),
)
.await
.expect("connect named middleware services");
let runner = ChainRunner::from_registry(registry);
let chain = vec![ChainEntry {
name: "scheme".into(),
implementation: "test/scheme".into(),
order: 0,
config: prost_types::Struct::default(),
on_error: OnError::FailClosed,
}];

for (port, scheme) in [(Some(80), "http"), (Some(443), "https")] {
let ctx = L7EvalContext {
request_default_port: port,
..ctx.clone()
};
let req = crate::l7::provider::L7Request {
action: "GET".into(),
target: "/v1/messages".into(),
query_params: std::collections::HashMap::new(),
raw_header: b"GET /v1/messages HTTP/1.1\r\nHost: api.example.test\r\n\r\n".to_vec(),
body_length: crate::l7::provider::BodyLength::None,
};
let (_app, mut client) = tokio::io::duplex(8192);
let result = apply_middleware_chain_with_request_id(
req,
&mut client,
&ctx,
chain.clone(),
&runner,
tunnel_engine.generation_guard(),
openshell_supervisor_middleware::TransformedBodyPolicy::NotPolicyRelevant,
"test-request-id",
)
.await
.expect("apply middleware chain");
assert!(matches!(result, MiddlewareApplyResult::Allowed(_)));
assert_eq!(seen.lock().unwrap().pop().as_deref(), Some(scheme));
}
}

#[tokio::test]
async fn all_unresolved_fail_open_forwards_body_unbuffered() {
// A chain whose only entry is an unregistered binding has no resolvable
Expand Down
Loading