Skip to content

Commit 3a8fb17

Browse files
committed
feat(transport): add :* port wildcard to origins
1 parent 4a1043b commit 3a8fb17

2 files changed

Lines changed: 178 additions & 19 deletions

File tree

‎crates/rmcp/src/transport/streamable_http_server/tower.rs‎

Lines changed: 142 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -113,16 +113,21 @@ pub struct StreamableHttpServerConfig {
113113
/// Defaults to an empty list, which disables Origin validation for backward
114114
/// compatibility. A non-empty list enables validation. Requests carrying
115115
/// an `Origin` header must match per RFC 6454 `(scheme, host, port)`;
116-
/// missing-`Origin` requests still pass. An entry that omits the port
117-
/// permits any port; an entry with an explicit port matches only that
118-
/// port, resolving an origin's omitted port to the scheme default
119-
/// (443 for `https`, 80 for `http`). Entries must include a scheme;
116+
/// missing-`Origin` requests still pass. Entries must include a scheme;
120117
/// `"null"` matches the browser's `Origin: null`.
121118
///
119+
/// Ports:
120+
/// - `:*` matches any port.
121+
/// - An explicit port matches only that port. An `Origin` without a port
122+
/// uses the scheme default (443 for `https`, 80 for `http`).
123+
/// - An entry without a port currently matches any port. This is
124+
/// deprecated: a future release will match only the scheme default
125+
/// port. Use `:*` or an explicit port instead.
126+
///
122127
/// Call [`StreamableHttpServerConfig::enforce_origin_validation`] to enable
123128
/// validation with an empty list, rejecting every present Origin value.
124129
/// examples:
125-
/// allowed_origins = ["https://app.example.com", "http://localhost:8080"]
130+
/// allowed_origins = ["https://app.example.com:443", "http://localhost:*"]
126131
pub allowed_origins: Vec<String>,
127132
validate_empty_origin_allowlist: bool,
128133
/// Optional external session store for cross-instance recovery.
@@ -865,14 +870,70 @@ fn default_port(scheme: &str) -> Option<u16> {
865870
}
866871
}
867872

873+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
874+
enum AllowedPort {
875+
/// `:*`
876+
Any,
877+
/// No port in the entry. Matches any port until the deprecation ends.
878+
Unspecified,
879+
Exact(u16),
880+
}
881+
882+
impl AllowedPort {
883+
fn matches(self, scheme: &str, port: Option<u16>) -> bool {
884+
match self {
885+
AllowedPort::Any | AllowedPort::Unspecified => true,
886+
// RFC 6454 §6.2 omits the default port when serializing an origin.
887+
AllowedPort::Exact(allowed) => port.or_else(|| default_port(scheme)) == Some(allowed),
888+
}
889+
}
890+
}
891+
892+
#[derive(Debug, Clone, PartialEq, Eq)]
893+
enum AllowedOrigin {
894+
Null,
895+
Tuple {
896+
scheme: String,
897+
host: String,
898+
port: AllowedPort,
899+
},
900+
}
901+
902+
fn parse_allowed_origin(value: &str) -> Option<AllowedOrigin> {
903+
let value = value.trim();
904+
if let Some(base) = value.strip_suffix(":*") {
905+
let NormalizedOrigin::Tuple {
906+
scheme,
907+
host,
908+
port: None,
909+
} = parse_origin_value(base)?
910+
else {
911+
return None;
912+
};
913+
return Some(AllowedOrigin::Tuple {
914+
scheme,
915+
host,
916+
port: AllowedPort::Any,
917+
});
918+
}
919+
Some(match parse_origin_value(value)? {
920+
NormalizedOrigin::Null => AllowedOrigin::Null,
921+
NormalizedOrigin::Tuple { scheme, host, port } => AllowedOrigin::Tuple {
922+
scheme,
923+
host,
924+
port: port.map_or(AllowedPort::Unspecified, AllowedPort::Exact),
925+
},
926+
})
927+
}
928+
868929
fn origin_is_allowed(origin: &NormalizedOrigin, allowed_origins: &[String]) -> bool {
869930
allowed_origins
870931
.iter()
871-
.filter_map(|raw| parse_origin_value(raw))
932+
.filter_map(|raw| parse_allowed_origin(raw))
872933
.any(|allowed| match (&allowed, origin) {
873-
(NormalizedOrigin::Null, NormalizedOrigin::Null) => true,
934+
(AllowedOrigin::Null, NormalizedOrigin::Null) => true,
874935
(
875-
NormalizedOrigin::Tuple {
936+
AllowedOrigin::Tuple {
876937
scheme: a_scheme,
877938
host: a_host,
878939
port: a_port,
@@ -882,21 +943,82 @@ fn origin_is_allowed(origin: &NormalizedOrigin, allowed_origins: &[String]) -> b
882943
host: o_host,
883944
port: o_port,
884945
},
885-
) => {
886-
a_scheme == o_scheme
887-
&& a_host == o_host
888-
&& match a_port {
889-
// An omitted configured port permits any port.
890-
None => true,
891-
// RFC 6454 §6.2 omits the default port when serializing an
892-
// origin, so an absent incoming port means the scheme default.
893-
Some(a_port) => o_port.or_else(|| default_port(o_scheme)) == Some(*a_port),
894-
}
895-
}
946+
) => a_scheme == o_scheme && a_host == o_host && a_port.matches(o_scheme, *o_port),
896947
_ => false,
897948
})
898949
}
899950

951+
fn warn_on_portless_allowed_origins(allowed_origins: &[String]) {
952+
for raw in allowed_origins {
953+
if let Some(AllowedOrigin::Tuple {
954+
port: AllowedPort::Unspecified,
955+
..
956+
}) = parse_allowed_origin(raw)
957+
{
958+
tracing::warn!(
959+
allowed_origin = raw.as_str(),
960+
"allowed origin without a port matches any port; a future release will match \
961+
only the scheme default port. Use `:*` or an explicit port instead",
962+
);
963+
}
964+
}
965+
}
966+
967+
#[cfg(test)]
968+
mod parse_allowed_origin_tests {
969+
use super::*;
970+
971+
fn tuple(scheme: &str, host: &str, port: AllowedPort) -> Option<AllowedOrigin> {
972+
Some(AllowedOrigin::Tuple {
973+
scheme: scheme.to_string(),
974+
host: host.to_string(),
975+
port,
976+
})
977+
}
978+
979+
#[test]
980+
fn wildcard_port_parses_as_any() {
981+
assert_eq!(
982+
parse_allowed_origin("https://client.example:*"),
983+
tuple("https", "client.example", AllowedPort::Any)
984+
);
985+
}
986+
987+
#[test]
988+
fn wildcard_port_on_ipv6_host_parses_as_any() {
989+
assert_eq!(
990+
parse_allowed_origin("http://[::1]:*"),
991+
tuple("http", "::1", AllowedPort::Any)
992+
);
993+
}
994+
995+
#[test]
996+
fn portless_entry_parses_as_unspecified() {
997+
assert_eq!(
998+
parse_allowed_origin("https://client.example"),
999+
tuple("https", "client.example", AllowedPort::Unspecified)
1000+
);
1001+
}
1002+
1003+
#[test]
1004+
fn explicit_port_parses_as_exact() {
1005+
assert_eq!(
1006+
parse_allowed_origin("https://client.example:8443"),
1007+
tuple("https", "client.example", AllowedPort::Exact(8443))
1008+
);
1009+
}
1010+
1011+
#[test]
1012+
fn wildcard_after_explicit_port_is_rejected() {
1013+
assert_eq!(parse_allowed_origin("https://client.example:443:*"), None);
1014+
}
1015+
1016+
#[test]
1017+
fn wildcard_on_null_is_rejected() {
1018+
assert_eq!(parse_allowed_origin("null:*"), None);
1019+
}
1020+
}
1021+
9001022
fn bad_request_response(message: &str) -> BoxResponse {
9011023
let body = Full::from(message.to_string()).boxed();
9021024

@@ -1158,6 +1280,7 @@ where
11581280
session_manager: Arc<M>,
11591281
config: StreamableHttpServerConfig,
11601282
) -> Self {
1283+
warn_on_portless_allowed_origins(&config.allowed_origins);
11611284
let pending_restores = config
11621285
.session_store
11631286
.is_some()

‎crates/rmcp/tests/test_custom_headers.rs‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1365,4 +1365,40 @@ mod origin_validation {
13651365
.await;
13661366
assert_eq!(response.status(), http::StatusCode::OK);
13671367
}
1368+
1369+
#[tokio::test]
1370+
async fn wildcard_port_allows_non_default_origin_port() {
1371+
let service = service_with_allowed_origins(&["https://client.example:*"]);
1372+
let response = service
1373+
.handle(init_request(Some("https://client.example:8443")))
1374+
.await;
1375+
assert_eq!(response.status(), http::StatusCode::OK);
1376+
}
1377+
1378+
#[tokio::test]
1379+
async fn wildcard_port_allows_port_less_origin() {
1380+
let service = service_with_allowed_origins(&["https://client.example:*"]);
1381+
let response = service
1382+
.handle(init_request(Some("https://client.example")))
1383+
.await;
1384+
assert_eq!(response.status(), http::StatusCode::OK);
1385+
}
1386+
1387+
#[tokio::test]
1388+
async fn wildcard_port_forbids_other_host() {
1389+
let service = service_with_allowed_origins(&["https://client.example:*"]);
1390+
let response = service
1391+
.handle(init_request(Some("https://attacker.example:8443")))
1392+
.await;
1393+
assert_eq!(response.status(), http::StatusCode::FORBIDDEN);
1394+
}
1395+
1396+
#[tokio::test]
1397+
async fn wildcard_port_forbids_scheme_mismatch() {
1398+
let service = service_with_allowed_origins(&["https://client.example:*"]);
1399+
let response = service
1400+
.handle(init_request(Some("http://client.example:8443")))
1401+
.await;
1402+
assert_eq!(response.status(), http::StatusCode::FORBIDDEN);
1403+
}
13681404
}

0 commit comments

Comments
 (0)