#![allow(clippy::unwrap_used)]
#![allow(clippy::pedantic, clippy::nursery, missing_docs)]
use std::{
borrow::Cow,
net::SocketAddr,
sync::{Arc, Mutex},
time::Duration,
};
use axum::{
body::{Body, to_bytes},
extract::{Request, State},
middleware::{self, Next},
response::Response,
};
use polyc_agent::ToolExecutor;
use polyc_tools::{
AudienceBoundToken, CALLER_HEADER, ConnectOptions, ConnectionPool, McpToolSource,
};
use rmcp::{
ErrorData as McpError, ServerHandler,
handler::server::{
router::tool::ToolRouter,
tool::{ToolCallContext, ToolRoute},
},
model::{
CallToolRequestParams, CallToolResponse, CallToolResult, Implementation, InitializeResult,
ListToolsResult, PaginatedRequestParams, ServerCapabilities, Tool,
},
service::{RequestContext, RoleServer},
};
use serde_json::json;
use tokio_util::sync::CancellationToken;
#[derive(Debug, Clone)]
struct SeenRequest {
method: String,
authorization: Option<String>,
caller: Option<String>,
had_session_id: bool,
}
#[derive(Clone, Default)]
struct Instrument {
seen: Arc<Mutex<Vec<SeenRequest>>>,
}
#[derive(Clone)]
struct EchoServer {
router: Arc<ToolRouter<Self>>,
}
impl EchoServer {
fn new() -> Self {
let mut router: ToolRouter<Self> = ToolRouter::new();
let schema = json!({
"type": "object",
"properties": { "text": { "type": "string" } },
});
let mut echo = Tool::new(
Cow::Borrowed("echo"),
Cow::Borrowed("Echo the input text back."),
schema.as_object().cloned().unwrap_or_default(),
);
echo.annotations = Some(echo.annotations.unwrap_or_default().read_only(true));
router.add_route(ToolRoute::new_dyn(echo, |ctx: ToolCallContext<Self>| {
Box::pin(async move {
let text = ctx
.arguments
.as_ref()
.and_then(|o| o.get("text"))
.and_then(serde_json::Value::as_str)
.unwrap_or("")
.to_owned();
Ok(CallToolResult::structured(json!({ "echo": text })).into())
})
}));
Self {
router: Arc::new(router),
}
}
}
impl std::fmt::Debug for EchoServer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EchoServer").finish_non_exhaustive()
}
}
impl ServerHandler for EchoServer {
fn get_info(&self) -> rmcp::model::ServerInfo {
InitializeResult::new(ServerCapabilities::builder().enable_tools().build())
.with_server_info(Implementation::new("echo", env!("CARGO_PKG_VERSION")))
}
fn list_tools(
&self,
_request: Option<PaginatedRequestParams>,
_context: RequestContext<RoleServer>,
) -> impl Future<Output = Result<ListToolsResult, McpError>> + Send + '_ {
let tools = self.router.list_all();
async move { Ok(ListToolsResult::with_all_items(tools)) }
}
fn call_tool(
&self,
request: CallToolRequestParams,
context: RequestContext<RoleServer>,
) -> impl Future<Output = Result<CallToolResponse, McpError>> + Send + '_ {
let router = self.router.clone();
async move {
let ctx = ToolCallContext::new(self, request, context);
router.call(ctx).await
}
}
}
async fn instrument(State(inst): State<Instrument>, req: Request, next: Next) -> Response {
let (parts, body) = req.into_parts();
let authorization = parts
.headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
let caller = parts
.headers
.get(CALLER_HEADER)
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
let had_session_id = parts.headers.get("mcp-session-id").is_some();
let bytes = to_bytes(body, 1 << 20).await.unwrap_or_default();
let method = serde_json::from_slice::<serde_json::Value>(&bytes)
.ok()
.and_then(|v| {
v.get("method")
.and_then(serde_json::Value::as_str)
.map(str::to_owned)
})
.unwrap_or_default();
inst.seen.lock().unwrap().push(SeenRequest {
method,
authorization,
caller,
had_session_id,
});
next.run(Request::from_parts(parts, Body::from(bytes)))
.await
}
async fn spawn_server() -> (
SocketAddr,
Instrument,
CancellationToken,
tokio::task::JoinHandle<()>,
) {
let inst = Instrument::default();
let router = polyc_tools::mcp_server::build_router("/mcp", EchoServer::new())
.layer(middleware::from_fn_with_state(inst.clone(), instrument));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ct = CancellationToken::new();
let server_ct = ct.clone();
let handle = tokio::spawn(async move {
let _ = axum::serve(listener, router)
.with_graceful_shutdown(async move { server_ct.cancelled_owned().await })
.await;
});
tokio::time::sleep(Duration::from_millis(50)).await;
(addr, inst, ct, handle)
}
fn labeled(label: &str) -> ConnectOptions {
ConnectOptions {
label: Some(label.to_owned()),
..ConnectOptions::default()
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn two_turns_same_principal_reuse_one_session() {
let (addr, inst, ct, handle) = spawn_server().await;
let uri = format!("http://{addr}/mcp");
let pool = ConnectionPool::new();
let token = AudienceBoundToken::new("token-conv1", &uri).expect("valid resource");
let opts = || ConnectOptions {
bearer: Some(token.clone()),
caller: Some("persona-1".to_owned()),
..labeled("notes")
};
let turn1 = McpToolSource::pooled(pool.clone(), "conv-1", uri.clone(), opts())
.await
.expect("turn 1 composes");
let out = turn1.execute("notes__echo", r#"{"text":"one"}"#).await;
assert_eq!(
serde_json::from_str::<serde_json::Value>(&out).unwrap()["echo"],
"one"
);
drop(turn1);
let turn2 = McpToolSource::pooled(pool.clone(), "conv-1", uri, opts())
.await
.expect("turn 2 composes");
let out = turn2.execute("notes__echo", r#"{"text":"two"}"#).await;
assert_eq!(
serde_json::from_str::<serde_json::Value>(&out).unwrap()["echo"],
"two"
);
drop(turn2);
let seen = inst.seen.lock().unwrap().clone();
let calls: Vec<_> = seen.iter().filter(|r| r.method == "tools/call").collect();
assert_eq!(calls.len(), 2, "both turns' tool calls reached the wire");
for call in &calls {
assert_eq!(call.authorization.as_deref(), Some("Bearer token-conv1"));
assert_eq!(call.caller.as_deref(), Some("persona-1"));
assert!(
!call.had_session_id,
"a same-principal reuse must never carry a session id: {call:?}"
);
}
ct.cancel();
let _ = tokio::time::timeout(Duration::from_secs(5), handle).await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn two_principals_never_see_each_others_caller_or_bearer() {
let (addr, inst, ct, handle) = spawn_server().await;
let uri = format!("http://{addr}/mcp");
let pool = ConnectionPool::new();
let token_a = AudienceBoundToken::new("token-A", &uri).expect("valid resource");
let token_b = AudienceBoundToken::new("token-B", &uri).expect("valid resource");
let a = McpToolSource::pooled(
pool.clone(),
"conv-a",
uri.clone(),
ConnectOptions {
bearer: Some(token_a),
caller: Some("persona-a".to_owned()),
..labeled("notes")
},
)
.await
.expect("principal a composes");
let b = McpToolSource::pooled(
pool.clone(),
"conv-b",
uri,
ConnectOptions {
bearer: Some(token_b),
caller: Some("persona-b".to_owned()),
..labeled("notes")
},
)
.await
.expect("principal b composes");
let out_a = a.execute("notes__echo", r#"{"text":"a"}"#).await;
let out_b = b.execute("notes__echo", r#"{"text":"b"}"#).await;
assert_eq!(
serde_json::from_str::<serde_json::Value>(&out_a).unwrap()["echo"],
"a"
);
assert_eq!(
serde_json::from_str::<serde_json::Value>(&out_b).unwrap()["echo"],
"b"
);
let seen = inst.seen.lock().unwrap().clone();
let calls: Vec<_> = seen.iter().filter(|r| r.method == "tools/call").collect();
assert_eq!(
calls.len(),
2,
"each principal's tool call must reach the wire exactly once: {seen:?}"
);
for req in &seen {
assert!(
!req.had_session_id,
"no request may carry a session identifier under the modern-only dial: {req:?}"
);
}
for call in &calls {
let is_a = call.authorization.as_deref() == Some("Bearer token-A")
&& call.caller.as_deref() == Some("persona-a");
let is_b = call.authorization.as_deref() == Some("Bearer token-B")
&& call.caller.as_deref() == Some("persona-b");
assert!(
is_a || is_b,
"a tools/call must carry exactly one principal's caller header and \
bearer together, never a cross of the two: {call:?}"
);
}
drop(a);
drop(b);
ct.cancel();
let _ = tokio::time::timeout(Duration::from_secs(5), handle).await;
}