Skip to content
Merged
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
152 changes: 140 additions & 12 deletions src/mcp.rs
Original file line number Diff line number Diff line change
@@ -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<Self, Self::Err> {
Ok(match s {
"prod" => Self::Prod,
"beta" => Self::Beta,
_ => {
let (canister_id, url) = s.splitn(2, ':').collect_tuple().ok_or_else(|| {
Comment thread
blind-oracle marked this conversation as resolved.
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);

Expand All @@ -36,18 +89,38 @@ impl Run for McpWrapper {
}
}

pub struct McpState {
hostname: String,
router: Router,
}

pub async fn middleware(
State(state): State<Arc<McpState>>,
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<Router, Error> {
) -> Result<McpState, Error> {
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
Expand Down Expand Up @@ -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())
Expand All @@ -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());
}
}
74 changes: 26 additions & 48 deletions src/routing/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -51,7 +50,6 @@ use crate::{
cli::Cli,
metrics::{self},
routing::{
error_cause::ClientError,
ic::routing_table_manager::LooksUpSubnetType,
middleware::{
canister_match, cors, headers,
Expand Down Expand Up @@ -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()
Expand All @@ -570,23 +544,8 @@ pub async fn setup_router(
.nest("/api/v4", router_api_v4)
.fallback(
|Extension(ctx): Extension<Arc<RequestCtx>>, 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;
}

Expand Down Expand Up @@ -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!(
Comment thread
blind-oracle marked this conversation as resolved.
"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));
Comment thread
blind-oracle marked this conversation as resolved.
}

#[cfg(all(target_os = "linux", feature = "sev-snp"))]
if cli.sev_snp.sev_snp_enable {
let router_sev_snp = Router::new().route(
Expand Down Expand Up @@ -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(""));
Expand Down Expand Up @@ -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",
Expand Down
Loading