pub mod core;
pub mod initialization;
pub mod logging;
pub mod prompts;
pub mod resources;
pub mod sampling;
pub mod tools;
pub use core::*;
pub use initialization::*;
pub use logging::{
LogLevel, LoggingNotification, ProgressNotification,
PromptListChangedNotification as LoggingPromptListChangedNotification,
ResourceListChangedNotification as LoggingResourceListChangedNotification,
ResourceUpdatedNotification as LoggingResourceUpdatedNotification, SetLevelRequest,
ToolListChangedNotification as LoggingToolListChangedNotification,
};
pub use prompts::{
GetPromptRequest, GetPromptResponse, ListPromptsRequest, ListPromptsResponse,
MessageRole as PromptMessageRole, Prompt, PromptContent, PromptListChangedNotification,
PromptMessage, ResourceReference as PromptResourceReference,
};
pub use resources::{
ListResourcesRequest, ListResourcesResponse, ReadResourceRequest, ReadResourceResponse,
Resource, ResourceContent, ResourceListChangedNotification, ResourceUpdatedNotification,
SubscribeRequest, UnsubscribeRequest,
};
pub use sampling::{
CompleteRequest, CompleteResponse, CompletionArgument, CompletionResult, CostPriority,
IntelligencePriority, MessageRole, ModelPreferences, SamplingContent, SamplingMessage,
SpeedPriority, StopReason,
};
pub use tools::{
CallToolRequest, CallToolResponse, ListToolsRequest, ListToolsResponse,
ResourceReference as ToolResourceReference, Tool, ToolListChangedNotification, ToolResult,
};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum ProtocolVersion {
#[serde(rename = "2024-11-05")]
V2024_11_05,
#[serde(rename = "2025-03-26")]
V2025_03_26,
#[serde(untagged)]
Custom(String),
}
impl ProtocolVersion {
pub fn as_str(&self) -> &str {
match self {
Self::V2024_11_05 => "2024-11-05",
Self::V2025_03_26 => "2025-03-26",
Self::Custom(version) => version,
}
}
pub fn is_supported(&self) -> bool {
matches!(self, Self::V2024_11_05 | Self::V2025_03_26)
}
pub fn supported_versions() -> Vec<Self> {
vec![Self::V2024_11_05, Self::V2025_03_26]
}
}
impl Default for ProtocolVersion {
fn default() -> Self {
Self::V2025_03_26
}
}
impl std::fmt::Display for ProtocolVersion {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct Capabilities {
#[serde(flatten)]
pub standard: StandardCapabilities,
#[serde(flatten)]
pub custom: HashMap<String, serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct StandardCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<ToolCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub resources: Option<ResourceCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompts: Option<PromptCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub sampling: Option<SamplingCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logging: Option<LoggingCapabilities>,
#[serde(skip_serializing_if = "Option::is_none")]
pub roots: Option<RootsCapabilities>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct ToolCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
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, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct PromptCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct SamplingCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub enabled: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct LoggingCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub level: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct RootsCapabilities {
#[serde(skip_serializing_if = "Option::is_none")]
pub list_changed: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Implementation {
pub name: String,
pub version: String,
#[serde(flatten)]
pub metadata: HashMap<String, serde_json::Value>,
}
impl Implementation {
pub fn new(name: impl Into<String>, version: impl Into<String>) -> Self {
Self {
name: name.into(),
version: version.into(),
metadata: HashMap::new(),
}
}
pub fn with_metadata(mut self, key: impl Into<String>, value: serde_json::Value) -> Self {
self.metadata.insert(key.into(), value);
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ProgressToken {
String(String),
Number(i64),
}
impl From<String> for ProgressToken {
fn from(s: String) -> Self {
Self::String(s)
}
}
impl From<&str> for ProgressToken {
fn from(s: &str) -> Self {
Self::String(s.to_string())
}
}
impl From<i64> for ProgressToken {
fn from(n: i64) -> Self {
Self::Number(n)
}
}
impl std::fmt::Display for ProgressToken {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::String(s) => write!(f, "{}", s),
Self::Number(n) => write!(f, "{}", n),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct PaginationCursor {
pub cursor: String,
}
impl PaginationCursor {
pub fn new(cursor: impl Into<String>) -> Self {
Self {
cursor: cursor.into(),
}
}
}
impl From<String> for PaginationCursor {
fn from(cursor: String) -> Self {
Self::new(cursor)
}
}
impl From<&str> for PaginationCursor {
fn from(cursor: &str) -> Self {
Self::new(cursor)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json;
#[test]
fn test_protocol_version_serialization() {
let version = ProtocolVersion::V2024_11_05;
let json = serde_json::to_string(&version).unwrap();
assert_eq!(json, "\"2024-11-05\"");
let deserialized: ProtocolVersion = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized, version);
}
#[test]
fn test_protocol_version_custom() {
let custom = ProtocolVersion::Custom("2025-01-01".to_string());
assert_eq!(custom.as_str(), "2025-01-01");
assert!(!custom.is_supported());
}
#[test]
fn test_capabilities_serialization() {
let capabilities = Capabilities {
standard: StandardCapabilities {
tools: Some(ToolCapabilities {
list_changed: Some(true),
}),
resources: Some(ResourceCapabilities {
subscribe: Some(true),
list_changed: Some(false),
}),
..Default::default()
},
custom: {
let mut custom = HashMap::new();
custom.insert("experimental".to_string(), serde_json::json!(true));
custom
},
};
let json = serde_json::to_value(&capabilities).unwrap();
let deserialized: Capabilities = serde_json::from_value(json).unwrap();
assert_eq!(deserialized, capabilities);
}
#[test]
fn test_implementation_creation() {
let impl_info = Implementation::new("mcp-probe", "0.1.0")
.with_metadata("platform", serde_json::json!("rust"));
assert_eq!(impl_info.name, "mcp-probe");
assert_eq!(impl_info.version, "0.1.0");
assert_eq!(
impl_info.metadata.get("platform").unwrap(),
&serde_json::json!("rust")
);
}
#[test]
fn test_progress_token_variants() {
let string_token = ProgressToken::from("progress-1");
let number_token = ProgressToken::from(42i64);
assert_eq!(string_token.to_string(), "progress-1");
assert_eq!(number_token.to_string(), "42");
let json_string = serde_json::to_string(&string_token).unwrap();
let json_number = serde_json::to_string(&number_token).unwrap();
assert_eq!(json_string, "\"progress-1\"");
assert_eq!(json_number, "42");
}
}