use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[non_exhaustive]
#[serde(rename_all = "camelCase")]
pub struct ClientCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub sampling: Option<SamplingCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub elicitation: Option<ElicitationCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub roots: Option<RootsCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tasks: Option<ClientTasksCapability>,
#[serde(skip_serializing_if = "Option::is_none")]
pub experimental: Option<HashMap<String, serde_json::Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extensions: Option<HashMap<String, serde_json::Value>>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[non_exhaustive]
#[serde(rename_all = "camelCase")]
pub struct ServerCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<ToolCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompts: Option<PromptCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub resources: Option<ResourceCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logging: Option<LoggingCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub completions: Option<CompletionCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub sampling: Option<SamplingCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tasks: Option<ServerTasksCapability>,
#[serde(skip_serializing_if = "Option::is_none")]
pub experimental: Option<HashMap<String, serde_json::Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extensions: Option<HashMap<String, serde_json::Value>>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ToolCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PromptCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ResourceCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub subscribe: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct LoggingCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub levels: Option<Vec<String>>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct SamplingCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub models: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub context: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ElicitationCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub form: Option<FormElicitationCapability>,
#[serde(skip_serializing_if = "Option::is_none")]
pub url: Option<Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct FormElicitationCapability {
#[serde(skip_serializing_if = "Option::is_none")]
pub apply_defaults: Option<bool>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ServerTasksCapability {
#[serde(skip_serializing_if = "Option::is_none")]
pub list: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cancel: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub requests: Option<ServerTasksRequestCapability>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ServerTasksRequestCapability {
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<ServerTasksToolsCapability>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ServerTasksToolsCapability {
#[serde(skip_serializing_if = "Option::is_none")]
pub call: Option<Value>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ClientTasksCapability {
#[serde(skip_serializing_if = "Option::is_none")]
pub list: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cancel: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub requests: Option<ClientTasksRequestCapability>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ClientTasksRequestCapability {
#[serde(skip_serializing_if = "Option::is_none")]
pub sampling: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub elicitation: Option<Value>,
}
pub const TASKS_EXTENSION_KEY: &str = "io.modelcontextprotocol/tasks";
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
#[serde(rename_all = "camelCase")]
pub struct TasksExtensionCapability {}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RootsCapabilities {
#[serde(default)]
pub list_changed: bool,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CompletionCapabilities {
#[serde(skip)]
_reserved: (),
}
impl ClientCapabilities {
pub fn minimal() -> Self {
Self::default()
}
pub fn full() -> Self {
Self {
sampling: Some(SamplingCapabilities::default()),
elicitation: Some(ElicitationCapabilities::default()),
roots: Some(RootsCapabilities { list_changed: true }),
tasks: Some(ClientTasksCapability::default()),
experimental: None,
extensions: None,
}
}
pub fn supports_elicitation(&self) -> bool {
self.elicitation.is_some()
}
pub fn supports_sampling(&self) -> bool {
self.sampling.is_some()
}
}
impl ServerCapabilities {
pub fn minimal() -> Self {
Self::default()
}
pub fn tools_only() -> Self {
Self {
tools: Some(ToolCapabilities {
list_changed: Some(true),
}),
..Default::default()
}
}
pub fn prompts_only() -> Self {
Self {
prompts: Some(PromptCapabilities {
list_changed: Some(true),
}),
..Default::default()
}
}
pub fn resources_only() -> Self {
Self {
resources: Some(ResourceCapabilities {
subscribe: Some(true),
list_changed: Some(true),
}),
..Default::default()
}
}
pub fn provides_tools(&self) -> bool {
self.tools.is_some()
}
pub fn provides_prompts(&self) -> bool {
self.prompts.is_some()
}
pub fn provides_resources(&self) -> bool {
self.resources.is_some()
}
pub fn provides_tasks(&self) -> bool {
self.tasks.is_some()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn client_capabilities_helpers() {
let minimal = ClientCapabilities::minimal();
assert!(!minimal.supports_sampling());
assert!(!minimal.supports_elicitation());
let full = ClientCapabilities::full();
assert!(full.supports_sampling());
assert!(full.supports_elicitation());
}
#[test]
fn server_capabilities_helpers() {
let tools_only = ServerCapabilities::tools_only();
assert!(tools_only.provides_tools());
assert!(!tools_only.provides_prompts());
assert!(!tools_only.provides_resources());
let prompts_only = ServerCapabilities::prompts_only();
assert!(!prompts_only.provides_tools());
assert!(prompts_only.provides_prompts());
assert!(!prompts_only.provides_resources());
}
#[test]
fn capabilities_serialization() {
let caps = ClientCapabilities {
sampling: Some(SamplingCapabilities::default()),
elicitation: Some(ElicitationCapabilities::default()),
roots: Some(RootsCapabilities { list_changed: true }),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
assert!(json.get("sampling").is_some());
assert!(json.get("elicitation").is_some());
assert_eq!(json["roots"]["listChanged"], true);
assert!(json.get("tools").is_none());
assert!(json.get("prompts").is_none());
assert!(json.get("resources").is_none());
}
#[test]
fn server_capabilities_auto_set_serialization() {
let caps = ServerCapabilities {
tools: Some(ToolCapabilities {
list_changed: Some(false),
}),
prompts: Some(PromptCapabilities {
list_changed: Some(false),
}),
resources: Some(ResourceCapabilities {
subscribe: Some(false),
list_changed: Some(false),
}),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
println!(
"Serialized capabilities: {}",
serde_json::to_string_pretty(&json).unwrap()
);
assert!(json.get("tools").is_some(), "tools should be present");
assert!(json.get("prompts").is_some(), "prompts should be present");
assert!(
json.get("resources").is_some(),
"resources should be present"
);
assert_eq!(json["tools"]["listChanged"], false);
assert_eq!(json["prompts"]["listChanged"], false);
assert_eq!(json["resources"]["listChanged"], false);
assert_eq!(json["resources"]["subscribe"], false);
}
#[test]
fn server_tasks_capability_serialization() {
let caps = ServerCapabilities {
tasks: Some(ServerTasksCapability {
list: Some(serde_json::json!({})),
cancel: Some(serde_json::json!({})),
requests: Some(ServerTasksRequestCapability {
tools: Some(ServerTasksToolsCapability {
call: Some(serde_json::json!({})),
}),
}),
}),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
assert!(json.get("tasks").is_some());
assert!(json["tasks"].get("list").is_some());
assert!(json["tasks"].get("cancel").is_some());
assert!(json["tasks"]["requests"]["tools"]["call"].is_object());
assert!(caps.provides_tasks());
}
#[test]
fn client_tasks_capability_serialization() {
let caps = ClientCapabilities::full();
let json = serde_json::to_value(&caps).unwrap();
assert!(json.get("tasks").is_some());
}
#[test]
fn default_serializes_without_extensions_key() {
let caps = ServerCapabilities::default();
let json = serde_json::to_value(&caps).unwrap();
assert!(
json.get("extensions").is_none(),
"default ServerCapabilities should not serialize an `extensions` key, \
got: {json}"
);
}
#[test]
fn extensions_round_trip_byte_equal() {
let mut ext = HashMap::new();
ext.insert(
"io.modelcontextprotocol/skills".to_string(),
serde_json::json!({}),
);
let caps = ServerCapabilities {
extensions: Some(ext),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
let round: ServerCapabilities = serde_json::from_value(json).unwrap();
assert_eq!(
round
.extensions
.as_ref()
.unwrap()
.get("io.modelcontextprotocol/skills"),
Some(&serde_json::json!({})),
"round-tripped extensions value must equal the original"
);
}
#[test]
fn extensions_and_experimental_coexist() {
let mut exp = HashMap::new();
exp.insert("old-thing".to_string(), serde_json::json!({"v": 1}));
let mut ext = HashMap::new();
ext.insert(
"io.modelcontextprotocol/skills".to_string(),
serde_json::json!({}),
);
let caps = ServerCapabilities {
experimental: Some(exp),
extensions: Some(ext),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
assert!(json.get("experimental").is_some(), "experimental missing");
assert!(json.get("extensions").is_some(), "extensions missing");
assert!(
json["experimental"].get("extensions").is_none(),
"extensions must NOT be nested inside experimental"
);
assert!(
json["extensions"].get("experimental").is_none(),
"experimental must NOT be nested inside extensions"
);
let round: ServerCapabilities = serde_json::from_value(json).unwrap();
assert!(round.experimental.is_some());
assert!(round.extensions.is_some());
assert_eq!(
round
.extensions
.as_ref()
.unwrap()
.get("io.modelcontextprotocol/skills"),
Some(&serde_json::json!({}))
);
assert_eq!(
round.experimental.as_ref().unwrap().get("old-thing"),
Some(&serde_json::json!({"v": 1}))
);
}
#[test]
fn extensions_camelcase_serde() {
let mut ext = HashMap::new();
ext.insert("k".to_string(), serde_json::json!(1));
let caps = ServerCapabilities {
extensions: Some(ext),
..Default::default()
};
let s = serde_json::to_string(&caps).unwrap();
assert!(
s.contains("\"extensions\""),
"wire form must contain exactly `\"extensions\"`, got: {s}"
);
assert!(!s.contains("\"Extensions\""));
assert!(!s.contains("\"extension\""));
}
#[test]
fn client_default_serializes_without_extensions_key() {
let serialized = serde_json::to_string(&ClientCapabilities::default()).unwrap();
assert!(
!serialized.contains("extensions"),
"default ClientCapabilities must not serialize an `extensions` key \
at all (not even as null or {{}}); got: {serialized}"
);
assert_eq!(
serialized, "{}",
"default ClientCapabilities must serialize as exactly `{{}}`; got: {serialized}"
);
}
#[test]
fn client_extensions_round_trip_byte_equal() {
let mut ext = HashMap::new();
ext.insert(TASKS_EXTENSION_KEY.to_string(), serde_json::json!({}));
let caps = ClientCapabilities {
extensions: Some(ext),
..Default::default()
};
let first = serde_json::to_string(&caps).unwrap();
let round: ClientCapabilities = serde_json::from_str(&first).unwrap();
let second = serde_json::to_string(&round).unwrap();
assert_eq!(
first, second,
"ClientCapabilities.extensions must survive a round-trip byte-for-byte"
);
assert_eq!(
first, r#"{"extensions":{"io.modelcontextprotocol/tasks":{}}}"#,
"the wire form of a tasks-declaring client must match the spec's own \
example bytes; re-verify against schema/vendored/ext-tasks/schema.ts"
);
assert_eq!(
round.extensions.as_ref().unwrap().get(TASKS_EXTENSION_KEY),
Some(&serde_json::json!({})),
"round-tripped extensions value must equal the original"
);
}
#[test]
fn client_extensions_and_experimental_coexist() {
let mut exp = HashMap::new();
exp.insert("old-thing".to_string(), serde_json::json!({"v": 1}));
let mut ext = HashMap::new();
ext.insert(TASKS_EXTENSION_KEY.to_string(), serde_json::json!({}));
let caps = ClientCapabilities {
experimental: Some(exp),
extensions: Some(ext),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
assert!(json.get("experimental").is_some(), "experimental missing");
assert!(json.get("extensions").is_some(), "extensions missing");
assert!(
json["experimental"].get("extensions").is_none(),
"extensions must NOT be nested inside experimental"
);
assert!(
json["extensions"].get("experimental").is_none(),
"experimental must NOT be nested inside extensions"
);
let round: ClientCapabilities = serde_json::from_value(json).unwrap();
assert_eq!(
round.extensions.as_ref().unwrap().get(TASKS_EXTENSION_KEY),
Some(&serde_json::json!({}))
);
assert_eq!(
round.experimental.as_ref().unwrap().get("old-thing"),
Some(&serde_json::json!({"v": 1}))
);
}
#[test]
fn tasks_extension_capability_serializes_as_empty_object() {
let serialized = serde_json::to_string(&TasksExtensionCapability::default()).unwrap();
assert_eq!(
serialized, "{}",
"TasksExtensionCapability must serialize as exactly `{{}}` \
(schema/vendored/ext-tasks/schema.ts declares it \
`Record<string, never>`); got: {serialized}"
);
let parsed: TasksExtensionCapability = serde_json::from_str("{}").unwrap();
assert_eq!(parsed, TasksExtensionCapability::default());
let with_future_field: TasksExtensionCapability =
serde_json::from_str(r#"{"someFutureSetting":true}"#).expect(
"TasksExtensionCapability must tolerate an unknown key so a future \
upstream field cannot break an older client",
);
assert_eq!(with_future_field, TasksExtensionCapability::default());
}
#[test]
fn tasks_extension_key_is_the_reverse_dns_spelling() {
assert_eq!(
TASKS_EXTENSION_KEY, "io.modelcontextprotocol/tasks",
"TASKS_EXTENSION_KEY changed. This value is PRE-FINAL and held under \
Phase 114's D-18 hold: re-verify against \
schema/vendored/ext-tasks/schema.ts (see PROVENANCE.md for the \
pinned commit) and 114-SPEC-RECHECK.md. Do not `fix` this test to \
match the code — a mismatch with the published schema is a \
phase-reopening event."
);
let mut ext = HashMap::new();
ext.insert(
TASKS_EXTENSION_KEY.to_string(),
serde_json::to_value(TasksExtensionCapability::default()).unwrap(),
);
let caps = ServerCapabilities {
extensions: Some(ext),
..Default::default()
};
assert_eq!(
serde_json::to_string(&caps).unwrap(),
r#"{"extensions":{"io.modelcontextprotocol/tasks":{}}}"#,
"must match schema/draft/examples/ServerCapabilities/extensions-tasks.json \
byte-for-byte"
);
}
#[test]
fn server_capabilities_with_none_fields_serialization() {
let caps = ServerCapabilities {
tools: Some(ToolCapabilities { list_changed: None }),
prompts: Some(PromptCapabilities { list_changed: None }),
resources: Some(ResourceCapabilities {
subscribe: None,
list_changed: None,
}),
..Default::default()
};
let json = serde_json::to_value(&caps).unwrap();
println!(
"Serialized capabilities with None: {}",
serde_json::to_string_pretty(&json).unwrap()
);
assert!(
json.get("tools").is_some(),
"tools should be present even with None fields"
);
assert!(
json.get("prompts").is_some(),
"prompts should be present even with None fields"
);
assert!(
json.get("resources").is_some(),
"resources should be present even with None fields"
);
}
}