use std::net::SocketAddr;
use std::sync::Arc;
use axum::Router;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::{HeaderValue, header};
use axum::middleware::{self, Next};
use axum::response::{IntoResponse, Response};
use codoseo_mcp::cloud::{AnonBackend, Caller, CloudMcp};
use rmcp::transport::streamable_http_server::session::never::NeverSessionManager;
use rmcp::transport::streamable_http_server::{StreamableHttpServerConfig, StreamableHttpService};
use crate::agent::anon::AnonCaller;
use crate::agent::auth::{self, ApiCaller};
use crate::agent::mcp::AgentBackend;
use crate::config::{Config, Mode};
use crate::metrics::{self, Surface, Tier};
use crate::state::AppState;
pub const PATH: &str = "/mcp";
type McpCaller = Caller<ApiCaller, AnonCaller>;
const MAX_BODY_BYTES: usize = 64 * 1024;
pub fn routes(state: &AppState) -> Router<AppState> {
router_for(state, CloudMcp::new(AgentBackend::new(state.clone())))
}
pub fn router_for<B>(state: &AppState, handler: CloudMcp<B>) -> Router<AppState>
where
B: AnonBackend<Keyed = ApiCaller, Anon = AnonCaller> + Clone,
{
let config = StreamableHttpServerConfig::default()
.with_legacy_session_mode(false)
.with_json_response(true)
.with_max_request_body_bytes(MAX_BODY_BYTES)
.with_cancellation_token(state.shutdown.clone());
let config = match state.config.mode {
Mode::Cloud => config.with_allowed_hosts(allowed_hosts(&state.config)),
Mode::SelfHost => config.disable_allowed_hosts(),
};
let service = StreamableHttpService::new(
move || Ok(handler.clone()),
Arc::new(NeverSessionManager::default()),
config,
);
Router::new()
.route_service(PATH, service)
.layer(middleware::from_fn_with_state(
state.clone(),
resolve_caller,
))
}
fn allowed_hosts(config: &Config) -> Vec<String> {
let mut hosts: Vec<String> = ["localhost", "127.0.0.1", "::1"].map(str::to_owned).into();
if let Some(host) = config.base_url.host_str() {
hosts.push(match config.base_url.port() {
Some(port) => format!("{host}:{port}"),
None => host.to_owned(),
});
}
hosts
}
async fn resolve_caller(State(state): State<AppState>, mut req: Request, next: Next) -> Response {
let headers = req.headers();
let caller: McpCaller = match auth::bearer_key(headers) {
Ok(None) if state.config.mode == Mode::Cloud => {
let peer = req
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|c| c.0.ip());
Caller::Anon(AnonCaller::from_request(&state, headers, peer))
}
_ => match auth::authenticate(&state, headers).await {
Ok(keyed) => Caller::Keyed(keyed),
Err(e) => {
metrics::api_request(Surface::Mcp, Tier::Key, e.code());
return e.into_response();
}
},
};
req.extensions_mut().insert(caller);
if let Some(value) = req.headers_mut().get_mut(header::AUTHORIZATION) {
value.set_sensitive(true);
}
let mut res = next.run(req).await;
res.headers_mut()
.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-store"));
res
}