diff --git a/src/mcp.rs b/src/mcp.rs index 3a9525e..008e7e4 100644 --- a/src/mcp.rs +++ b/src/mcp.rs @@ -1,28 +1,81 @@ use std::{ + fmt::Display, + str::FromStr, sync::Arc, time::{SystemTime, UNIX_EPOCH}, }; use anyhow::{Context, Error, anyhow}; use async_trait::async_trait; -use axum::{Router, middleware::from_fn_with_state, response::Redirect}; -use ic_bn_lib::tasks::{Run, TaskManager}; +use axum::{ + Router, + extract::{Request, State}, + middleware::{Next, from_fn_with_state}, + response::{IntoResponse, Redirect, Response}, +}; +use candid::Principal; +use ic_bn_lib::{ + http::extract_authority, + tasks::{Run, TaskManager}, +}; use imcp2::{ Agent, IiInstance, McpConfig, McpServer, SharedClients, auth_callbacks_router, metrics::{Metrics, write_request_metrics}, }; +use itertools::Itertools; use prometheus::Registry; -use strum::{Display, EnumString}; use tokio_util::sync::CancellationToken; +use tower::ServiceExt; +use url::Url; -#[derive(EnumString, Clone, Copy, Display)] -#[strum(serialize_all = "snake_case")] +use crate::cli::McpCli; + +#[derive(Clone)] pub enum IiType { Prod, Beta, + Custom(IiInstance), } -use crate::cli::McpCli; +impl Display for IiType { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Prod => write!(f, "prod"), + Self::Beta => write!(f, "beta"), + Self::Custom(instance) => write!( + f, + "{}:{}:{}", + instance.name, instance.ii_canister, instance.ii_url + ), + } + } +} + +impl FromStr for IiType { + type Err = Error; + + fn from_str(s: &str) -> Result { + Ok(match s { + "prod" => Self::Prod, + "beta" => Self::Beta, + _ => { + let (canister_id, url) = s.splitn(2, ':').collect_tuple().ok_or_else(|| { + anyhow!("invalid custom II instance format, expected canister_id:url") + })?; + + let canister_id = + Principal::from_str(canister_id).context("invalid canister id")?; + let url = Url::parse(url).context("invalid URL")?; + + Self::Custom(IiInstance { + name: "custom", + ii_canister: canister_id, + ii_url: url.to_string(), + }) + } + }) + } +} struct McpWrapper(McpServer); @@ -36,18 +89,38 @@ impl Run for McpWrapper { } } +pub struct McpState { + hostname: String, + router: Router, +} + +pub async fn middleware( + State(state): State>, + request: Request, + next: Next, +) -> Response { + // If the request is for the MCP hostname, route it to the MCP router directly + if let Some(authority) = extract_authority(&request) + && authority.eq_ignore_ascii_case(&state.hostname) + { + return state.router.clone().oneshot(request).await.into_response(); + } + + next.run(request).await.into_response() +} + /// Inject MCP routes into Router pub fn setup_mcp( cli: &McpCli, agent: Agent, registry: &Registry, tasks: &mut TaskManager, -) -> Result { +) -> Result { let ii_instance = match cli.mcp_ii_instance.as_ref().unwrap() { - IiType::Beta => IiInstance::beta(), - IiType::Prod => IiInstance::prod(), - } - .map_err(Error::msg)?; + IiType::Beta => IiInstance::beta().map_err(Error::msg)?, + IiType::Prod => IiInstance::prod().map_err(Error::msg)?, + IiType::Custom(instance) => instance.clone(), + }; let state_dir = cli .mcp_state_dir @@ -84,6 +157,14 @@ pub fn setup_mcp( .context("unable to create MCP metrics")?; let mcp_redirect_url = cli.mcp_root_redirect.to_string(); + let hostname = cli + .mcp_public_url + .as_ref() + .unwrap() + .host_str() + .unwrap() + .to_string(); + let router = Router::new() .nest_service(mcp.mcp_path(), mcp.mcp_router()) .merge(mcp.well_known_router()) @@ -94,5 +175,52 @@ pub fn setup_mcp( tasks.add("mcp", Arc::new(McpWrapper(mcp))); - Ok(router) + Ok(McpState { hostname, router }) +} + +#[cfg(test)] +mod test { + use super::*; + + #[test] + fn test_ii_type_from_str_prod() { + assert!(matches!(IiType::from_str("prod").unwrap(), IiType::Prod)); + } + + #[test] + fn test_ii_type_from_str_beta() { + assert!(matches!(IiType::from_str("beta").unwrap(), IiType::Beta)); + } + + #[test] + fn test_ii_type_from_str_custom() { + let s = "aaaaa-aa:https://example.com"; + let ii_type = IiType::from_str(s).unwrap(); + + let IiType::Custom(instance) = ii_type else { + panic!("expected IiType::Custom"); + }; + + assert_eq!(instance.name, "custom"); + assert_eq!( + instance.ii_canister, + Principal::from_str("aaaaa-aa").unwrap() + ); + assert_eq!(instance.ii_url, "https://example.com/"); + } + + #[test] + fn test_ii_type_from_str_custom_invalid_format() { + assert!(IiType::from_str("no-colon-here").is_err()); + } + + #[test] + fn test_ii_type_from_str_custom_invalid_canister_id() { + assert!(IiType::from_str("not-a-canister-id:https://example.com").is_err()); + } + + #[test] + fn test_ii_type_from_str_custom_invalid_url() { + assert!(IiType::from_str("aaaaa-aa:not a url").is_err()); + } } diff --git a/src/routing/mod.rs b/src/routing/mod.rs index 2b546e9..20436e2 100644 --- a/src/routing/mod.rs +++ b/src/routing/mod.rs @@ -26,7 +26,6 @@ use ic_bn_lib::{ http::{ Client, ClientHttp, cache::{CacheBuilder, KeyExtractorUriRange}, - extract_authority, extract_host, middleware::{request_meta, waf::WafLayer}, shed::{ ShardedOptions, ShedResponse, TypeExtractor, @@ -51,7 +50,6 @@ use crate::{ cli::Cli, metrics::{self}, routing::{ - error_cause::ClientError, ic::routing_table_manager::LooksUpSubnetType, middleware::{ canister_match, cors, headers, @@ -524,31 +522,7 @@ pub async fn setup_router( )) .layer(option_layer(prerender_mw)); - let api_hostname = cli.api.api_hostname.clone().map(|x| x.to_string()); - - #[cfg(feature = "mcp")] - let mcp = if let Some(v) = cli.mcp.mcp_ii_instance { - warn!( - "Starting MCP at {} (II {v})", - cli.mcp.mcp_public_url.as_ref().unwrap(), - ); - - let mcp_hostname = cli - .mcp - .mcp_public_url - .clone() - .unwrap() - .host_str() - .unwrap() - .to_string(); - - let router = crate::mcp::setup_mcp(&cli.mcp, ic_agent.clone(), registry, &mut *tasks) - .context("unable to set up MCP")?; - - Some((router, mcp_hostname)) - } else { - None - }; + let api_hostname = cli.api.api_hostname.clone(); let custom_domains_router = custom_domains_router.map(|x| { Router::new() @@ -570,23 +544,8 @@ pub async fn setup_router( .nest("/api/v4", router_api_v4) .fallback( |Extension(ctx): Extension>, request: Request| async move { - let Some(host) = extract_authority(&request) else { - return Ok(ErrorCause::Client(ClientError::NoAuthority).into_response()); - }; - - // Check if MCP is enabled & the request's host matches MCP hostname - #[cfg(feature = "mcp")] - if let (Some((mcp_router, mcp_hostname)), Some(host)) = (mcp, extract_host(host)) - && host.eq_ignore_ascii_case(&mcp_hostname) - { - return mcp_router.oneshot(request).await; - } - - // Check if the request's host matches API hostname - if api_hostname - .zip(extract_host(host)) - .is_some_and(|(a, b)| a == b) - { + // Check if API is enabled & the request's host matches API hostname + if api_hostname.is_some_and(|x| x == ctx.authority) { return router_api.oneshot(request).await; } @@ -626,6 +585,23 @@ pub async fn setup_router( ) .layer(common_layers); + #[cfg(feature = "mcp")] + if let Some(v) = &cli.mcp.mcp_ii_instance { + use crate::mcp; + + warn!( + "Starting MCP at {} (II {v})", + cli.mcp.mcp_public_url.as_ref().unwrap(), + ); + + let state = mcp::setup_mcp(&cli.mcp, ic_agent.clone(), registry, &mut *tasks) + .context("unable to set up MCP")?; + + // Inject MCP middleware to the top of the chain that will intercept the calls + // to the MCP hostname and route them to the MCP router directly + router = router.layer(from_fn_with_state(Arc::new(state), mcp::middleware)); + } + #[cfg(all(target_os = "linux", feature = "sev-snp"))] if cli.sev_snp.sev_snp_enable { let router_sev_snp = Router::new().route( @@ -839,10 +815,12 @@ mod test { ACCESS_CONTROL_ALLOW_METHODS, AUTHORIZATION, CACHE_CONTROL, LOCATION, WWW_AUTHENTICATE, }; - const MCP_HOST: &str = "mcp.ic0.app"; - const ISSUER: &str = "https://mcp.ic0.app/mcp"; + // Deliberately doesn't overlap with any base domain (e.g. ic0.app) to make sure MCP + // works on an arbitrary hostname rather than base domain or under it. + const MCP_HOST: &str = "mcp.example.com"; + const ISSUER: &str = "https://mcp.example.com/mcp"; const PROTECTED_RESOURCE_URL: &str = - "https://mcp.ic0.app/.well-known/oauth-protected-resource/mcp"; + "https://mcp.example.com/.well-known/oauth-protected-resource/mcp"; fn request(method: Method, host: &str, path_and_query: &str) -> Request { let mut req = Request::new(Body::from("")); @@ -871,7 +849,7 @@ mod test { "--mcp-ii-instance", "prod", "--mcp-public-url", - "https://mcp.ic0.app", + "https://mcp.example.com", "--mcp-state-dir", "/tmp/ic-gateway-test-mcp-state", "--mcp-root-redirect",