1use std::collections::HashMap;
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7
8pub const TASKS_EXTENSION_ID: &str = "io.modelcontextprotocol/tasks";
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
13#[serde(rename_all = "snake_case")]
14pub enum TaskStatus {
15 Working,
16 InputRequired,
17 Completed,
18 Failed,
19 Cancelled,
20}
21
22impl TaskStatus {
23 pub fn is_terminal(self) -> bool {
24 matches!(self, Self::Completed | Self::Failed | Self::Cancelled)
25 }
26}
27
28#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
30pub struct Task {
31 #[serde(rename = "taskId")]
32 pub task_id: String,
33 pub status: TaskStatus,
34 #[serde(rename = "statusMessage", skip_serializing_if = "Option::is_none")]
35 pub status_message: Option<String>,
36 #[serde(rename = "createdAt")]
37 pub created_at: String,
38 #[serde(rename = "lastUpdatedAt")]
39 pub last_updated_at: String,
40 #[serde(rename = "ttlMs")]
41 pub ttl_ms: Option<u64>,
42 #[serde(rename = "pollIntervalMs", skip_serializing_if = "Option::is_none")]
43 pub poll_interval_ms: Option<u64>,
44 #[serde(
45 rename = "inputRequests",
46 default,
47 skip_serializing_if = "HashMap::is_empty"
48 )]
49 pub input_requests: HashMap<String, Value>,
50 #[serde(skip_serializing_if = "Option::is_none")]
51 pub result: Option<Value>,
52 #[serde(skip_serializing_if = "Option::is_none")]
53 pub error: Option<Value>,
54}
55
56impl Task {
57 pub fn validate(&self) -> Result<(), String> {
59 match self.status {
60 TaskStatus::InputRequired if self.input_requests.is_empty() => {
61 Err("input_required task must include inputRequests".to_string())
62 }
63 TaskStatus::Completed if self.result.is_none() => {
64 Err("completed task must include result".to_string())
65 }
66 TaskStatus::Failed if self.error.is_none() => {
67 Err("failed task must include error".to_string())
68 }
69 TaskStatus::Working | TaskStatus::Cancelled
70 if !self.input_requests.is_empty()
71 || self.result.is_some()
72 || self.error.is_some() =>
73 {
74 Err("task contains fields that do not match its status".to_string())
75 }
76 _ => Ok(()),
77 }
78 }
79}
80
81#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
83pub struct CreateTaskResult {
84 #[serde(rename = "resultType")]
85 pub result_type: String,
86 #[serde(flatten)]
87 pub task: Task,
88 #[serde(rename = "_meta", default, skip_serializing_if = "HashMap::is_empty")]
89 pub meta: HashMap<String, Value>,
90}
91
92#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
94pub struct GetTaskResult {
95 #[serde(rename = "resultType")]
96 pub result_type: String,
97 #[serde(flatten)]
98 pub task: Task,
99 #[serde(rename = "_meta", default, skip_serializing_if = "HashMap::is_empty")]
100 pub meta: HashMap<String, Value>,
101}
102
103#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
104pub struct GetTaskParams {
105 #[serde(rename = "taskId")]
106 pub task_id: String,
107 #[serde(rename = "_meta", default, skip_serializing_if = "HashMap::is_empty")]
108 pub meta: HashMap<String, Value>,
109}
110
111#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
112pub struct UpdateTaskParams {
113 #[serde(rename = "taskId")]
114 pub task_id: String,
115 #[serde(rename = "inputResponses")]
116 pub input_responses: HashMap<String, Value>,
117 #[serde(rename = "_meta", default, skip_serializing_if = "HashMap::is_empty")]
118 pub meta: HashMap<String, Value>,
119}
120
121#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
122pub struct CancelTaskParams {
123 #[serde(rename = "taskId")]
124 pub task_id: String,
125 #[serde(rename = "_meta", default, skip_serializing_if = "HashMap::is_empty")]
126 pub meta: HashMap<String, Value>,
127}
128
129#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
131pub struct TaskAcknowledgement {
132 #[serde(rename = "resultType")]
133 pub result_type: String,
134 #[serde(rename = "_meta", default, skip_serializing_if = "HashMap::is_empty")]
135 pub meta: HashMap<String, Value>,
136}
137
138pub fn has_tasks_extension(capabilities: &super::ClientCapabilities) -> bool {
139 capabilities
140 .extensions
141 .as_ref()
142 .is_some_and(|extensions| extensions.contains_key(TASKS_EXTENSION_ID))
143}
144
145#[cfg(test)]
146mod tests {
147 use super::*;
148
149 #[test]
150 fn task_wire_fields_are_flat() {
151 let task = Task {
152 task_id: "t".to_string(),
153 status: TaskStatus::Completed,
154 status_message: None,
155 created_at: "2026-01-01T00:00:00Z".to_string(),
156 last_updated_at: "2026-01-01T00:00:01Z".to_string(),
157 ttl_ms: Some(60_000),
158 poll_interval_ms: Some(100),
159 input_requests: HashMap::new(),
160 result: Some(serde_json::json!({"content":[]})),
161 error: None,
162 };
163 task.validate().unwrap();
164 let value = serde_json::to_value(GetTaskResult {
165 result_type: "complete".to_string(),
166 task,
167 meta: HashMap::new(),
168 })
169 .unwrap();
170 assert_eq!(value["taskId"], "t");
171 assert_eq!(value["resultType"], "complete");
172 assert!(value.get("task").is_none());
173 }
174}