Skip to main content

prism_mcp_rs/protocol/
tasks.rs

1//! `io.modelcontextprotocol/tasks` extension types (SEP-2663).
2
3use std::collections::HashMap;
4
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7
8/// Official extension identifier.
9pub const TASKS_EXTENSION_ID: &str = "io.modelcontextprotocol/tasks";
10
11/// Lifecycle state of a durable MCP task.
12#[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/// Complete task state returned by `tasks/get` and `notifications/tasks`.
29#[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    /// Verify status-specific fields before a task is put on the wire.
58    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/// Task-shaped result returned in place of an immediate tool result.
82#[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/// Result of `tasks/get`.
93#[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/// Empty acknowledgement returned by task updates and cancellation.
130#[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}