codoseo_web/routes/
mcp.rs1use std::net::SocketAddr;
14use std::sync::Arc;
15
16use axum::Router;
17use axum::extract::{ConnectInfo, Request, State};
18use axum::http::{HeaderValue, header};
19use axum::middleware::{self, Next};
20use axum::response::{IntoResponse, Response};
21use codoseo_mcp::cloud::{AnonBackend, Caller, CloudMcp};
22use rmcp::transport::streamable_http_server::session::never::NeverSessionManager;
23use rmcp::transport::streamable_http_server::{StreamableHttpServerConfig, StreamableHttpService};
24
25use crate::agent::anon::AnonCaller;
26use crate::agent::auth::{self, ApiCaller};
27use crate::agent::mcp::AgentBackend;
28use crate::config::{Config, Mode};
29use crate::metrics::{self, Surface, Tier};
30use crate::state::AppState;
31
32pub const PATH: &str = "/mcp";
34
35type McpCaller = Caller<ApiCaller, AnonCaller>;
37
38const MAX_BODY_BYTES: usize = 64 * 1024;
40
41pub fn routes(state: &AppState) -> Router<AppState> {
42 router_for(state, CloudMcp::new(AgentBackend::new(state.clone())))
43}
44
45pub fn router_for<B>(state: &AppState, handler: CloudMcp<B>) -> Router<AppState>
47where
48 B: AnonBackend<Keyed = ApiCaller, Anon = AnonCaller> + Clone,
49{
50 let config = StreamableHttpServerConfig::default()
51 .with_legacy_session_mode(false)
53 .with_json_response(true)
54 .with_max_request_body_bytes(MAX_BODY_BYTES)
55 .with_cancellation_token(state.shutdown.clone());
57 let config = match state.config.mode {
58 Mode::Cloud => config.with_allowed_hosts(allowed_hosts(&state.config)),
59 Mode::SelfHost => config.disable_allowed_hosts(),
64 };
65 let service = StreamableHttpService::new(
66 move || Ok(handler.clone()),
67 Arc::new(NeverSessionManager::default()),
68 config,
69 );
70 Router::new()
71 .route_service(PATH, service)
72 .layer(middleware::from_fn_with_state(
73 state.clone(),
74 resolve_caller,
75 ))
76}
77
78fn allowed_hosts(config: &Config) -> Vec<String> {
81 let mut hosts: Vec<String> = ["localhost", "127.0.0.1", "::1"].map(str::to_owned).into();
82 if let Some(host) = config.base_url.host_str() {
83 hosts.push(match config.base_url.port() {
84 Some(port) => format!("{host}:{port}"),
85 None => host.to_owned(),
86 });
87 }
88 hosts
89}
90
91async fn resolve_caller(State(state): State<AppState>, mut req: Request, next: Next) -> Response {
94 let headers = req.headers();
95 let caller: McpCaller = match auth::bearer_key(headers) {
96 Ok(None) if state.config.mode == Mode::Cloud => {
97 let peer = req
98 .extensions()
99 .get::<ConnectInfo<SocketAddr>>()
100 .map(|c| c.0.ip());
101 Caller::Anon(AnonCaller::from_request(&state, headers, peer))
102 }
103 _ => match auth::authenticate(&state, headers).await {
105 Ok(keyed) => Caller::Keyed(keyed),
106 Err(e) => {
107 metrics::api_request(Surface::Mcp, Tier::Key, e.code());
108 return e.into_response();
109 }
110 },
111 };
112 req.extensions_mut().insert(caller);
113 if let Some(value) = req.headers_mut().get_mut(header::AUTHORIZATION) {
115 value.set_sensitive(true);
116 }
117 let mut res = next.run(req).await;
118 res.headers_mut()
119 .insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
120 res
121}