#![cfg(not(target_arch = "wasm32"))]
#[path = "common/duplex.rs"]
mod duplex;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use async_trait::async_trait;
use serde_json::{json, Value};
use pmcp::server::core::ProtocolHandler;
use pmcp::shared::{Transport, TransportMessage};
use pmcp::testing::{META_CLIENT_CAPABILITIES, META_PROTOCOL_VERSION};
use pmcp::types::jsonrpc::{JSONRPCResponse, RequestId, ResponsePayload};
use pmcp::types::protocol::{
ProtocolVersion, LATEST_PROTOCOL_VERSION, PROTOCOL_VERSION_2026_07_28,
};
use pmcp::types::{
ClientCapabilities, ClientRequest, Content, CreateMessageResult, GetPromptResult,
ListResourcesResult, ReadResourceResult, Request, RootsCapabilities, SamplingCapabilities,
};
use pmcp::{
PromptHandler, ResourceHandler, Result, SamplingHandler, Server, ServerBuilder,
ServerCapabilities, ToolHandler,
};
use duplex::DuplexTransport;
const DEADLINE: Duration = Duration::from_secs(5);
fn distinctive_capabilities() -> ClientCapabilities {
let mut caps = ClientCapabilities::default();
caps.roots = Some(RootsCapabilities { list_changed: true });
caps.sampling = Some(SamplingCapabilities {
models: Some(vec!["g9-model".to_string()]),
context: None,
tools: None,
});
caps.experimental = Some(HashMap::from([(
"io.pmcp.test/g9-marker".to_string(),
json!("v1-handshake"),
)]));
caps
}
fn v2_meta_capabilities() -> ClientCapabilities {
let mut caps = ClientCapabilities::default();
caps.experimental = Some(HashMap::from([(
"io.pmcp.test/g9-marker".to_string(),
json!("v2-meta"),
)]));
caps
}
fn canonical(caps: &ClientCapabilities) -> Value {
serde_json::to_value(caps).expect("ClientCapabilities serialises")
}
#[derive(Clone, Debug)]
struct Observed {
capabilities: Option<Value>,
era: Option<String>,
}
impl Observed {
fn of(extra: &pmcp::RequestHandlerExtra) -> Self {
Self {
capabilities: extra.client_capabilities().map(canonical),
era: extra.era().map(|era| format!("{era:?}")),
}
}
}
type Slot = Arc<Mutex<Option<Observed>>>;
fn slot() -> Slot {
Arc::new(Mutex::new(None))
}
struct Slots {
tool: Slot,
prompt: Slot,
read: Slot,
list: Slot,
sampling: Slot,
}
impl Slots {
fn new() -> Self {
Self {
tool: slot(),
prompt: slot(),
read: slot(),
list: slot(),
sampling: slot(),
}
}
}
fn observed(slot: &Slot, what: &str) -> Observed {
slot.lock()
.expect("observation slot is not poisoned")
.clone()
.unwrap_or_else(|| panic!("{what} handler never ran"))
}
fn record(slot: &Slot, extra: &pmcp::RequestHandlerExtra) {
*slot.lock().expect("observation slot is not poisoned") = Some(Observed::of(extra));
}
struct CapturingTool(Slot);
#[async_trait]
impl ToolHandler for CapturingTool {
async fn handle(&self, _args: Value, extra: pmcp::RequestHandlerExtra) -> Result<Value> {
record(&self.0, &extra);
Ok(json!("ok"))
}
}
struct CapturingPrompt(Slot);
#[async_trait]
impl PromptHandler for CapturingPrompt {
async fn handle(
&self,
_args: HashMap<String, String>,
extra: pmcp::RequestHandlerExtra,
) -> Result<GetPromptResult> {
record(&self.0, &extra);
Ok(GetPromptResult::new(vec![], None))
}
}
struct CapturingResource {
read: Slot,
list: Slot,
}
#[async_trait]
impl ResourceHandler for CapturingResource {
async fn read(
&self,
_uri: &str,
extra: pmcp::RequestHandlerExtra,
) -> Result<ReadResourceResult> {
record(&self.read, &extra);
Ok(ReadResourceResult::new(vec![Content::text("ok")]))
}
async fn list(
&self,
_cursor: Option<String>,
extra: pmcp::RequestHandlerExtra,
) -> Result<ListResourcesResult> {
record(&self.list, &extra);
Ok(ListResourcesResult::new(vec![]))
}
}
struct CapturingSampling(Slot);
#[async_trait]
impl SamplingHandler for CapturingSampling {
async fn create_message(
&self,
_params: pmcp::types::CreateMessageParams,
extra: pmcp::RequestHandlerExtra,
) -> Result<CreateMessageResult> {
record(&self.0, &extra);
Ok(CreateMessageResult::new(Content::text("ok"), "g9-model"))
}
}
fn client_request(method: &str, params: Value) -> Request {
let mut envelope = serde_json::Map::new();
envelope.insert("method".to_string(), Value::String(method.to_string()));
envelope.insert("params".to_string(), params);
let parsed: ClientRequest = serde_json::from_value(Value::Object(envelope))
.unwrap_or_else(|e| panic!("`{method}` deserializes into ClientRequest ({e})"));
Request::Client(Box::new(parsed))
}
fn initialize_params(caps: &ClientCapabilities) -> Value {
json!({
"protocolVersion": LATEST_PROTOCOL_VERSION,
"capabilities": canonical(caps),
"clientInfo": { "name": "g9-fence", "version": "0.0.0" },
})
}
fn v2_meta(caps: &ClientCapabilities) -> Value {
json!({
META_PROTOCOL_VERSION: PROTOCOL_VERSION_2026_07_28,
META_CLIENT_CAPABILITIES: canonical(caps),
})
}
fn assert_ok(response: &JSONRPCResponse, what: &str) {
if let ResponsePayload::Error(error) = &response.payload {
panic!("{what} must succeed, got JSON-RPC error: {error:?}");
}
}
fn probes(meta: Option<&Value>) -> [(&'static str, Value); 3] {
let mut probes = [
("tools/call", json!({ "name": "probe", "arguments": {} })),
(
"prompts/get",
json!({ "name": "greeting", "arguments": {} }),
),
("resources/read", json!({ "uri": "mem://greeting" })),
];
if let Some(meta) = meta {
for (_, params) in &mut probes {
params["_meta"] = meta.clone();
}
}
probes
}
fn assert_carries(cell: &Slot, what: &str, expected: &Value) {
let seen = observed(cell, what);
assert_eq!(
seen.capabilities.as_ref(),
Some(expected),
"{what}: extra.client_capabilities() must carry the v1 handshake's advertised set \
(era seen: {:?})",
seen.era
);
}
struct ServerDriver {
transport: DuplexTransport,
handle: tokio::task::JoinHandle<()>,
next_id: i64,
}
impl ServerDriver {
fn spawn(server: Server) -> Self {
let (client_t, server_t) = DuplexTransport::pair();
let handle = tokio::spawn(async move {
let _ = server.run(server_t).await;
});
Self {
transport: client_t,
handle,
next_id: 1,
}
}
async fn call(&mut self, method: &str, params: Value) -> JSONRPCResponse {
let id = RequestId::from(self.next_id);
self.next_id += 1;
tokio::time::timeout(
DEADLINE,
self.transport.send(TransportMessage::Request {
id: id.clone(),
request: client_request(method, params),
}),
)
.await
.unwrap_or_else(|_| panic!("`{method}` send deadline"))
.unwrap_or_else(|e| panic!("`{method}` send failed ({e})"));
loop {
let message = tokio::time::timeout(DEADLINE, self.transport.receive())
.await
.unwrap_or_else(|_| panic!("`{method}` response deadline"))
.unwrap_or_else(|e| panic!("`{method}` receive failed ({e})"));
if let TransportMessage::Response(response) = message {
if response.id == id {
return response;
}
}
}
}
async fn probe_all(&mut self, meta: Option<&Value>, context: &str) {
for (method, params) in probes(meta) {
let response = self.call(method, params).await;
assert_ok(&response, &format!("{method} [{context}]"));
}
}
async fn shutdown(self) {
let Self {
transport, handle, ..
} = self;
drop(transport);
handle.abort();
let _ = tokio::time::timeout(DEADLINE, handle).await;
}
}
fn build_server(slots: &Slots, accept_list: Option<Vec<ProtocolVersion>>) -> Server {
let builder: ServerBuilder = Server::builder()
.name("g9-fence-server")
.version("1.0.0")
.capabilities(ServerCapabilities::default())
.tool("probe", CapturingTool(slots.tool.clone()))
.prompt("greeting", CapturingPrompt(slots.prompt.clone()))
.resources(CapturingResource {
read: slots.read.clone(),
list: slots.list.clone(),
})
.sampling(CapturingSampling(slots.sampling.clone()));
let builder = match accept_list {
Some(versions) => builder.with_supported_protocol_versions(versions),
None => builder,
};
builder.build().expect("server builds")
}
fn build_core(slots: &Slots) -> Arc<dyn ProtocolHandler> {
Arc::new(
pmcp::server::builder::ServerCoreBuilder::new()
.name("g9-fence-core")
.version("1.0.0")
.tool("probe", CapturingTool(slots.tool.clone()))
.prompt("greeting", CapturingPrompt(slots.prompt.clone()))
.resources(CapturingResource {
read: slots.read.clone(),
list: slots.list.clone(),
})
.stateless_mode(false)
.build()
.expect("core builds"),
)
}
async fn core_call(
core: &Arc<dyn ProtocolHandler>,
id: i64,
method: &str,
params: Value,
) -> JSONRPCResponse {
tokio::time::timeout(
DEADLINE,
core.handle_request(RequestId::from(id), client_request(method, params), None),
)
.await
.unwrap_or_else(|_| panic!("`{method}` deadline against ServerCore"))
}
async fn probe_all_via_core(core: &Arc<dyn ProtocolHandler>, first_id: i64, context: &str) {
for (offset, (method, params)) in probes(None).into_iter().enumerate() {
let id = first_id + i64::try_from(offset).expect("probe count fits in i64");
let response = core_call(core, id, method, params).await;
assert_ok(&response, &format!("{method} [{context}]"));
}
}
#[tokio::test]
async fn v1_handshake_capabilities_reach_tool_prompt_and_resource_handlers() {
let slots = Slots::new();
let server = build_server(
&slots,
None,
);
let mut driver = ServerDriver::spawn(server);
let advertised = distinctive_capabilities();
let init = driver
.call("initialize", initialize_params(&advertised))
.await;
assert_ok(&init, "v1 initialize");
driver.probe_all(None, "Server").await;
driver.shutdown().await;
let expected = canonical(&advertised);
assert_carries(&slots.tool, "tools/call", &expected);
assert_carries(&slots.prompt, "prompts/get", &expected);
assert_carries(&slots.read, "resources/read", &expected);
}
#[tokio::test]
async fn v1_handshake_capabilities_reach_the_server_core_dispatcher() {
let slots = Slots::new();
let core = build_core(&slots);
let advertised = distinctive_capabilities();
let init = core_call(&core, 0, "initialize", initialize_params(&advertised)).await;
assert_ok(&init, "v1 initialize against ServerCore");
probe_all_via_core(&core, 1, "ServerCore").await;
let expected = canonical(&advertised);
assert_carries(&slots.tool, "ServerCore tools/call", &expected);
assert_carries(&slots.prompt, "ServerCore prompts/get", &expected);
assert_carries(&slots.read, "ServerCore resources/read", &expected);
}
#[tokio::test]
async fn v1_handshake_capabilities_reach_the_thread_then_fold_sites() {
let advertised = distinctive_capabilities();
let expected = canonical(&advertised);
let slots = Slots::new();
let server = build_server(&slots, None);
let mut driver = ServerDriver::spawn(server);
let init = driver
.call("initialize", initialize_params(&advertised))
.await;
assert_ok(&init, "v1 initialize");
let list_resources = driver.call("resources/list", json!({})).await;
assert_ok(&list_resources, "resources/list");
let create_message = driver
.call(
"sampling/createMessage",
json!({ "messages": [], "maxTokens": 16 }),
)
.await;
assert_ok(&create_message, "sampling/createMessage");
driver.shutdown().await;
assert_carries(&slots.list, "Server resources/list", &expected);
assert_carries(&slots.sampling, "Server sampling/createMessage", &expected);
let core_slots = Slots::new();
let core = build_core(&core_slots);
let init = core_call(&core, 0, "initialize", initialize_params(&advertised)).await;
assert_ok(&init, "v1 initialize against ServerCore");
let list_resources = core_call(&core, 1, "resources/list", json!({})).await;
assert_ok(&list_resources, "resources/list against ServerCore");
assert_carries(&core_slots.list, "ServerCore resources/list", &expected);
}
#[tokio::test]
async fn v2_meta_client_capabilities_still_win_over_a_v1_handshake() {
let slots = Slots::new();
let server = build_server(&slots, Some(duplex::v2_accept_list()));
let mut driver = ServerDriver::spawn(server);
let handshake = distinctive_capabilities();
let init = driver
.call("initialize", initialize_params(&handshake))
.await;
assert_ok(&init, "v1 initialize on a v2-opted-in server");
let per_request = v2_meta_capabilities();
driver.probe_all(Some(&v2_meta(&per_request)), "v2").await;
driver.shutdown().await;
let expected = canonical(&per_request);
let handshake_json = canonical(&handshake);
for (what, cell) in [
("tools/call", &slots.tool),
("prompts/get", &slots.prompt),
("resources/read", &slots.read),
] {
let seen = observed(cell, what);
assert_eq!(
seen.era.as_deref(),
Some("V2"),
"v2 {what}: the dispatcher must have resolved Era::V2, otherwise this case asserts \
nothing about the v2 path"
);
assert_eq!(
seen.capabilities.as_ref(),
Some(&expected),
"v2 {what}: the per-request `_meta` clientCapabilities must win over the v1 \
handshake value ({handshake_json})"
);
}
}
#[tokio::test]
async fn no_handshake_and_no_meta_yields_none() {
let slots = Slots::new();
let server = build_server(&slots, None);
let mut driver = ServerDriver::spawn(server);
driver.probe_all(None, "no handshake").await;
let list_resources = driver.call("resources/list", json!({})).await;
assert_ok(&list_resources, "resources/list with no handshake");
driver.shutdown().await;
for (what, cell) in [
("tools/call", &slots.tool),
("prompts/get", &slots.prompt),
("resources/read", &slots.read),
("resources/list", &slots.list),
] {
let seen = observed(cell, what);
assert_eq!(
seen.capabilities, None,
"{what}: with no handshake and no `_meta`, client_capabilities() must be None, not a \
fabricated default"
);
assert_eq!(
seen.era, None,
"{what}: a server that never saw a handshake must not gain a synthesised era either"
);
}
}