use thiserror::Error;
use url::Url;
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct CanonicalResourceUrl {
url: Url,
canonical: String,
}
#[non_exhaustive]
#[derive(Debug, Error, PartialEq, Eq)]
pub enum CanonicalUrlError {
#[error("url is malformed: {0}")]
Malformed(#[from] url::ParseError),
#[error("url must use https scheme")]
InsecureScheme,
#[error("url must not contain a fragment")]
Fragment,
#[error("url must not contain a query string")]
QueryString,
#[error("url must not have a trailing slash (use bare-root form)")]
TrailingSlash,
}
impl CanonicalResourceUrl {
#[must_use = "parsing returns a canonicalised URL; callers must persist or compare it"]
pub fn parse(s: &str) -> Result<Self, CanonicalUrlError> {
let url = Url::parse(s)?;
if url.scheme() != "https" {
return Err(CanonicalUrlError::InsecureScheme);
}
if url.fragment().is_some() {
return Err(CanonicalUrlError::Fragment);
}
if url.query().is_some() {
return Err(CanonicalUrlError::QueryString);
}
let path = url.path();
if path.len() > 1 && path.ends_with('/') {
return Err(CanonicalUrlError::TrailingSlash);
}
let raw = url.as_str();
let canonical = if path == "/" && raw.ends_with('/') {
raw.trim_end_matches('/').to_owned()
} else {
raw.to_owned()
};
Ok(Self { url, canonical })
}
#[must_use]
pub fn as_str(&self) -> &str {
&self.canonical
}
#[must_use]
pub fn into_url(self) -> Url {
self.url
}
}
pub const MAX_ACCEPTED_AUDIENCE_HOSTS: usize = 5;
#[non_exhaustive]
#[derive(Clone, Debug)]
pub struct CanonicalUrlConfig {
issuer: CanonicalResourceUrl,
primary_resource: CanonicalResourceUrl,
accepted_resources: Vec<CanonicalResourceUrl>,
}
#[non_exhaustive]
#[derive(Debug, Error)]
pub enum CanonicalUrlConfigError {
#[error("oauth.canonical_host is required when oauth.mcp_enabled is true")]
Missing,
#[error("oauth.accepted_audience_hosts exceeds cap of {MAX_ACCEPTED_AUDIENCE_HOSTS}")]
TooManyAliases,
#[error("canonical host invalid: {0}")]
InvalidHost(#[from] CanonicalUrlError),
}
impl CanonicalUrlConfig {
pub fn new(
canonical_host: String,
accepted_aliases: Vec<String>,
) -> Result<Self, CanonicalUrlConfigError> {
if canonical_host.is_empty() {
return Err(CanonicalUrlConfigError::Missing);
}
if accepted_aliases.len() > MAX_ACCEPTED_AUDIENCE_HOSTS {
return Err(CanonicalUrlConfigError::TooManyAliases);
}
let issuer = CanonicalResourceUrl::parse(&format!("https://{canonical_host}"))?;
let primary_resource =
CanonicalResourceUrl::parse(&format!("https://{canonical_host}/mcp"))?;
let mut accepted_resources = vec![primary_resource.clone()];
for alias in accepted_aliases {
let r = CanonicalResourceUrl::parse(&format!("https://{alias}/mcp"))?;
if accepted_resources.iter().any(|p| p == &r) {
continue;
}
accepted_resources.push(r);
}
Ok(Self {
issuer,
primary_resource,
accepted_resources,
})
}
#[must_use]
pub fn issuer(&self) -> &CanonicalResourceUrl {
&self.issuer
}
#[must_use]
pub fn primary_resource(&self) -> &CanonicalResourceUrl {
&self.primary_resource
}
#[must_use]
pub fn accepts_audience(&self, aud: &str) -> bool {
self.accepted_resources.iter().any(|r| r.as_str() == aud)
}
#[must_use]
pub fn accepted_resources_strings(&self) -> Vec<String> {
self.accepted_resources
.iter()
.map(|r| r.as_str().to_owned())
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_simple_host() {
let u = CanonicalResourceUrl::parse("https://controller.example.com/mcp").unwrap();
assert_eq!(u.as_str(), "https://controller.example.com/mcp");
}
#[test]
fn rejects_fragment() {
let err = CanonicalResourceUrl::parse("https://controller.example.com/mcp#x").unwrap_err();
assert!(matches!(err, CanonicalUrlError::Fragment));
}
#[test]
fn rejects_query() {
let err =
CanonicalResourceUrl::parse("https://controller.example.com/mcp?x=1").unwrap_err();
assert!(matches!(err, CanonicalUrlError::QueryString));
}
#[test]
fn rejects_http_scheme() {
let err = CanonicalResourceUrl::parse("http://controller.example.com/mcp").unwrap_err();
assert!(matches!(err, CanonicalUrlError::InsecureScheme));
}
#[test]
fn rejects_trailing_slash() {
let err = CanonicalResourceUrl::parse("https://controller.example.com/mcp/").unwrap_err();
assert!(matches!(err, CanonicalUrlError::TrailingSlash));
}
#[test]
fn lowercases_host() {
let u = CanonicalResourceUrl::parse("https://Controller.Example.Com/mcp").unwrap();
assert_eq!(u.as_str(), "https://controller.example.com/mcp");
}
#[test]
fn config_derives_issuer_and_resource_from_host() {
let cfg = CanonicalUrlConfig::new("controller.example.com".into(), vec![]).unwrap();
assert_eq!(cfg.issuer().as_str(), "https://controller.example.com");
assert_eq!(
cfg.primary_resource().as_str(),
"https://controller.example.com/mcp"
);
assert!(cfg.accepts_audience("https://controller.example.com/mcp"));
}
#[test]
fn config_rejects_too_many_aliases() {
let aliases: Vec<String> = (0..=MAX_ACCEPTED_AUDIENCE_HOSTS)
.map(|i| format!("alias{i}.example.com"))
.collect();
let err = CanonicalUrlConfig::new("controller.example.com".into(), aliases).unwrap_err();
assert!(matches!(err, CanonicalUrlConfigError::TooManyAliases));
}
#[test]
fn config_accepts_alias_audience() {
let cfg = CanonicalUrlConfig::new(
"controller.example.com".into(),
vec!["legacy.example.com".into()],
)
.unwrap();
assert!(cfg.accepts_audience("https://controller.example.com/mcp"));
assert!(cfg.accepts_audience("https://legacy.example.com/mcp"));
assert!(!cfg.accepts_audience("https://intruder.example.com/mcp"));
}
#[test]
fn config_missing_host_errors() {
let err = CanonicalUrlConfig::new(String::new(), vec![]).unwrap_err();
assert!(matches!(err, CanonicalUrlConfigError::Missing));
}
#[test]
fn accepted_resources_strings_includes_primary_and_aliases() {
let cfg = CanonicalUrlConfig::new(
"controller.example.com".into(),
vec!["legacy.example.com".into()],
)
.unwrap();
let strings = cfg.accepted_resources_strings();
assert!(strings.contains(&"https://controller.example.com/mcp".to_owned()));
assert!(strings.contains(&"https://legacy.example.com/mcp".to_owned()));
assert_eq!(strings.len(), 2);
}
#[test]
fn accepted_resources_strings_no_aliases() {
let cfg = CanonicalUrlConfig::new("controller.example.com".into(), vec![]).unwrap();
let strings = cfg.accepted_resources_strings();
assert_eq!(
strings,
vec!["https://controller.example.com/mcp".to_owned()]
);
}
}