use serde::{Deserialize, Serialize};
use super::{InputRequests, JsonObject, MetaObject, ResultType};
pub const TASKS_EXTENSION_ID: &str = "io.modelcontextprotocol/tasks";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub enum TaskStatus {
#[default]
Working,
InputRequired,
Completed,
Failed,
Cancelled,
}
impl TaskStatus {
pub fn is_terminal(&self) -> bool {
matches!(self, Self::Completed | Self::Failed | Self::Cancelled)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
#[serde(rename_all = "camelCase")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub struct Task {
pub task_id: String,
pub status: TaskStatus,
#[serde(skip_serializing_if = "Option::is_none")]
pub status_message: Option<String>,
pub created_at: String,
pub last_updated_at: String,
pub ttl_ms: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub poll_interval_ms: Option<u64>,
}
impl Task {
pub fn new(
task_id: impl Into<String>,
status: TaskStatus,
created_at: impl Into<String>,
last_updated_at: impl Into<String>,
) -> Self {
Self {
task_id: task_id.into(),
status,
status_message: None,
created_at: created_at.into(),
last_updated_at: last_updated_at.into(),
ttl_ms: None,
poll_interval_ms: None,
}
}
pub fn with_status_message(mut self, status_message: impl Into<String>) -> Self {
self.status_message = Some(status_message.into());
self
}
pub fn with_ttl_ms(mut self, ttl_ms: u64) -> Self {
self.ttl_ms = Some(ttl_ms);
self
}
pub fn with_poll_interval_ms(mut self, poll_interval_ms: u64) -> Self {
self.poll_interval_ms = Some(poll_interval_ms);
self
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub enum TaskPayload {
Working,
InputRequired {
input_requests: InputRequests,
},
Completed {
result: JsonObject,
},
Failed {
error: JsonObject,
},
Cancelled,
}
impl TaskPayload {
pub fn status(&self) -> TaskStatus {
match self {
Self::Working => TaskStatus::Working,
Self::InputRequired { .. } => TaskStatus::InputRequired,
Self::Completed { .. } => TaskStatus::Completed,
Self::Failed { .. } => TaskStatus::Failed,
Self::Cancelled => TaskStatus::Cancelled,
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct DetailedTask {
pub task: Task,
pub payload: TaskPayload,
}
impl DetailedTask {
pub fn new(mut task: Task, payload: TaskPayload) -> Self {
task.status = payload.status();
Self { task, payload }
}
pub fn status(&self) -> TaskStatus {
self.task.status
}
}
#[derive(Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
struct DetailedTaskWire {
#[serde(flatten)]
task: Task,
#[serde(skip_serializing_if = "Option::is_none")]
input_requests: Option<InputRequests>,
#[serde(skip_serializing_if = "Option::is_none")]
result: Option<JsonObject>,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<JsonObject>,
}
impl From<DetailedTask> for DetailedTaskWire {
fn from(value: DetailedTask) -> Self {
let DetailedTask { task, payload } = value;
let (input_requests, result, error) = match payload {
TaskPayload::Working | TaskPayload::Cancelled => (None, None, None),
TaskPayload::InputRequired { input_requests } => (Some(input_requests), None, None),
TaskPayload::Completed { result } => (None, Some(result), None),
TaskPayload::Failed { error } => (None, None, Some(error)),
};
Self {
task,
input_requests,
result,
error,
}
}
}
impl TryFrom<DetailedTaskWire> for DetailedTask {
type Error = String;
fn try_from(wire: DetailedTaskWire) -> Result<Self, String> {
let payload = match wire.task.status {
TaskStatus::Working => TaskPayload::Working,
TaskStatus::Cancelled => TaskPayload::Cancelled,
TaskStatus::InputRequired => TaskPayload::InputRequired {
input_requests: wire.input_requests.ok_or_else(|| {
"task with status \"input_required\" is missing `inputRequests`".to_owned()
})?,
},
TaskStatus::Completed => TaskPayload::Completed {
result: wire.result.ok_or_else(|| {
"task with status \"completed\" is missing `result`".to_owned()
})?,
},
TaskStatus::Failed => TaskPayload::Failed {
error: wire
.error
.ok_or_else(|| "task with status \"failed\" is missing `error`".to_owned())?,
},
};
Ok(DetailedTask {
task: wire.task,
payload,
})
}
}
impl Serialize for DetailedTask {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
DetailedTaskWire::from(self.clone()).serialize(serializer)
}
}
impl<'de> Deserialize<'de> for DetailedTask {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let wire = DetailedTaskWire::deserialize(deserializer)?;
Self::try_from(wire).map_err(serde::de::Error::custom)
}
}
#[cfg(feature = "schemars")]
impl schemars::JsonSchema for DetailedTask {
fn schema_name() -> std::borrow::Cow<'static, str> {
"DetailedTask".into()
}
fn json_schema(generator: &mut schemars::SchemaGenerator) -> schemars::Schema {
<DetailedTaskWire as schemars::JsonSchema>::json_schema(generator)
}
}
#[derive(Debug, Clone, PartialEq, Serialize)]
#[serde(rename_all = "camelCase")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub struct CreateTaskResult {
pub result_type: ResultType,
#[serde(flatten)]
pub task: Task,
#[serde(rename = "_meta", default, skip_serializing_if = "Option::is_none")]
pub meta: Option<MetaObject>,
}
impl<'de> Deserialize<'de> for CreateTaskResult {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(rename_all = "camelCase")]
struct Helper {
result_type: ResultType,
#[serde(flatten)]
task: Task,
#[serde(rename = "_meta", default)]
meta: Option<MetaObject>,
}
let helper = Helper::deserialize(deserializer)?;
if !helper.result_type.is_task() {
return Err(serde::de::Error::custom(
"CreateTaskResult requires resultType to be \"task\"",
));
}
Ok(CreateTaskResult {
result_type: helper.result_type,
task: helper.task,
meta: helper.meta,
})
}
}
impl CreateTaskResult {
pub fn new(task: Task) -> Self {
Self {
result_type: ResultType::TASK,
task,
meta: None,
}
}
pub fn with_meta(mut self, meta: MetaObject) -> Self {
self.meta = Some(meta);
self
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub struct GetTaskResult {
#[serde(default)]
pub result_type: ResultType,
#[serde(rename = "_meta", default, skip_serializing_if = "Option::is_none")]
pub meta: Option<MetaObject>,
#[serde(flatten)]
pub task: DetailedTask,
}
impl GetTaskResult {
pub fn new(task: DetailedTask) -> Self {
Self {
result_type: ResultType::COMPLETE,
meta: None,
task,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize)]
#[serde(rename_all = "camelCase")]
#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
#[non_exhaustive]
pub struct TaskAckResult {
pub result_type: ResultType,
#[serde(rename = "_meta", default, skip_serializing_if = "Option::is_none")]
pub meta: Option<MetaObject>,
}
impl<'de> Deserialize<'de> for TaskAckResult {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(rename_all = "camelCase", deny_unknown_fields)]
struct Helper {
result_type: ResultType,
#[serde(rename = "_meta", default)]
meta: Option<MetaObject>,
}
let helper = Helper::deserialize(deserializer)?;
if !helper.result_type.is_complete() {
return Err(serde::de::Error::custom(
"TaskAckResult requires resultType to be \"complete\"",
));
}
Ok(TaskAckResult {
result_type: helper.result_type,
meta: helper.meta,
})
}
}
impl Default for TaskAckResult {
fn default() -> Self {
Self {
result_type: ResultType::COMPLETE,
meta: None,
}
}
}
impl TaskAckResult {
pub fn new() -> Self {
Self::default()
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
fn base_task(status: TaskStatus) -> Task {
Task::new(
"task-1",
status,
"2025-11-25T10:30:00Z",
"2025-11-25T10:40:00Z",
)
.with_ttl_ms(60000)
.with_poll_interval_ms(5000)
}
#[test]
fn create_task_result_wire_shape() {
let result = CreateTaskResult::new(base_task(TaskStatus::Working));
let value = serde_json::to_value(&result).unwrap();
assert_eq!(
value,
json!({
"resultType": "task",
"taskId": "task-1",
"status": "working",
"createdAt": "2025-11-25T10:30:00Z",
"lastUpdatedAt": "2025-11-25T10:40:00Z",
"ttlMs": 60000,
"pollIntervalMs": 5000
})
);
let roundtrip: CreateTaskResult = serde_json::from_value(value).unwrap();
assert_eq!(roundtrip, result);
}
#[test]
fn ttl_ms_null_means_unlimited() {
let mut task = base_task(TaskStatus::Working);
task.ttl_ms = None;
let value = serde_json::to_value(&task).unwrap();
assert_eq!(value["ttlMs"], serde_json::Value::Null);
let roundtrip: Task = serde_json::from_value(value).unwrap();
assert_eq!(roundtrip.ttl_ms, None);
}
#[test]
fn detailed_task_completed_roundtrip() {
let detailed = DetailedTask::new(
base_task(TaskStatus::Working),
TaskPayload::Completed {
result: serde_json::from_value(json!({
"content": [{"type": "text", "text": "ok"}],
"isError": false
}))
.unwrap(),
},
);
assert_eq!(detailed.status(), TaskStatus::Completed);
let value = serde_json::to_value(&detailed).unwrap();
assert_eq!(value["status"], "completed");
assert_eq!(value["result"]["isError"], false);
let roundtrip: DetailedTask = serde_json::from_value(value).unwrap();
assert_eq!(roundtrip, detailed);
}
#[test]
fn detailed_task_input_required_requires_input_requests() {
let err = serde_json::from_value::<DetailedTask>(json!({
"taskId": "task-1",
"status": "input_required",
"createdAt": "2025-11-25T10:30:00Z",
"lastUpdatedAt": "2025-11-25T10:40:00Z",
"ttlMs": null
}))
.unwrap_err();
assert!(err.to_string().contains("inputRequests"));
}
#[test]
fn detailed_task_failed_roundtrip() {
let detailed = DetailedTask::new(
base_task(TaskStatus::Failed),
TaskPayload::Failed {
error: serde_json::from_value(json!({
"code": -32603,
"message": "boom"
}))
.unwrap(),
},
);
let value = serde_json::to_value(&detailed).unwrap();
assert_eq!(value["status"], "failed");
assert_eq!(value["error"]["code"], -32603);
let roundtrip: DetailedTask = serde_json::from_value(value).unwrap();
assert_eq!(roundtrip, detailed);
}
#[test]
fn get_task_result_flattens_detailed_task() {
let result = GetTaskResult::new(DetailedTask::new(
base_task(TaskStatus::Working),
TaskPayload::Working,
));
let value = serde_json::to_value(&result).unwrap();
assert_eq!(value["taskId"], "task-1");
assert_eq!(value["status"], "working");
let roundtrip: GetTaskResult = serde_json::from_value(value).unwrap();
assert_eq!(roundtrip, result);
}
}