use std::{
collections::HashMap,
pin::Pin,
sync::{Arc, Mutex},
time::Instant,
};
use futures::Future;
use tokio::sync::oneshot;
use crate::{
error::ErrorData as McpError,
model::{
CallToolResult, DetailedTask, InputRequest, InputRequests, JsonObject, Task, TaskPayload,
TaskStatus,
},
};
pub const DEFAULT_TASK_TTL_MS: u64 = 300_000;
pub const DEFAULT_POLL_INTERVAL_MS: u64 = 1_000;
pub fn current_timestamp() -> String {
chrono::Utc::now().to_rfc3339()
}
#[derive(Clone)]
pub struct TaskContext {
task_id: String,
inner: Arc<Mutex<TaskManagerInner>>,
}
impl TaskContext {
pub fn task_id(&self) -> &str {
&self.task_id
}
pub async fn request_input(
&self,
key: impl Into<String>,
request: InputRequest,
) -> Result<serde_json::Value, TaskExit> {
let key = key.into();
let (tx, rx) = oneshot::channel();
{
let mut inner = self.inner.lock().expect("task manager lock poisoned");
let entry = inner.tasks.get_mut(&self.task_id).ok_or_else(|| {
TaskExit::Error(McpError::internal_error(
"task no longer exists".to_string(),
None,
))
})?;
if !entry.used_input_keys.insert(key.clone()) {
return Err(TaskExit::Error(McpError::internal_error(
format!("inputRequests key {key:?} was already used for this task"),
None,
)));
}
entry.pending_inputs.insert(key.clone(), (request, tx));
entry.touch();
}
rx.await.map_err(|_| TaskExit::Cancelled)
}
pub fn set_status_message(&self, message: impl Into<String>) {
let mut inner = self.inner.lock().expect("task manager lock poisoned");
if let Some(entry) = inner.tasks.get_mut(&self.task_id) {
entry.task.status_message = Some(message.into());
entry.touch();
}
}
pub fn is_cancel_requested(&self) -> bool {
let inner = self.inner.lock().expect("task manager lock poisoned");
inner
.tasks
.get(&self.task_id)
.is_some_and(|e| e.cancel_requested)
}
pub async fn cancelled(&self) {
let mut rx = {
let inner = self.inner.lock().expect("task manager lock poisoned");
let Some(entry) = inner.tasks.get(&self.task_id) else {
return;
};
if entry.cancel_requested {
return;
}
entry.cancel_signal.subscribe()
};
while !*rx.borrow_and_update() {
if rx.changed().await.is_err() {
return;
}
}
}
}
#[expect(
clippy::exhaustive_enums,
reason = "error variant for task exit may only be due to error or cancellation"
)]
#[derive(Debug)]
pub enum TaskExit {
Cancelled,
Error(McpError),
}
impl From<McpError> for TaskExit {
fn from(error: McpError) -> Self {
TaskExit::Error(error)
}
}
pub type TaskFuture = Pin<Box<dyn Future<Output = Result<CallToolResult, TaskExit>> + Send>>;
struct TaskEntry {
task: Task,
terminal: Option<TaskPayload>,
terminal_at: Option<Instant>,
pending_inputs: HashMap<String, (InputRequest, oneshot::Sender<serde_json::Value>)>,
used_input_keys: std::collections::HashSet<String>,
cancel_requested: bool,
cancel_signal: tokio::sync::watch::Sender<bool>,
created: Instant,
join_handle: Option<tokio::task::JoinHandle<()>>,
}
impl TaskEntry {
fn touch(&mut self) {
self.task.last_updated_at = current_timestamp();
}
fn current_status(&self) -> TaskStatus {
match &self.terminal {
Some(payload) => payload.status(),
None if !self.pending_inputs.is_empty() => TaskStatus::InputRequired,
None => TaskStatus::Working,
}
}
fn detailed(&self) -> DetailedTask {
let payload = match &self.terminal {
Some(p) => p.clone(),
None if !self.pending_inputs.is_empty() => TaskPayload::InputRequired {
input_requests: self
.pending_inputs
.iter()
.map(|(k, (req, _))| (k.clone(), req.clone()))
.collect::<InputRequests>(),
},
None => TaskPayload::Working,
};
DetailedTask::new(self.task.clone(), payload)
}
}
#[derive(Default)]
struct TaskManagerInner {
tasks: HashMap<String, TaskEntry>,
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct TaskOptions {
pub ttl_ms: Option<u64>,
pub poll_interval_ms: Option<u64>,
pub status_message: Option<String>,
}
impl Default for TaskOptions {
fn default() -> Self {
Self {
ttl_ms: Some(DEFAULT_TASK_TTL_MS),
poll_interval_ms: Some(DEFAULT_POLL_INTERVAL_MS),
status_message: None,
}
}
}
impl TaskOptions {
pub fn new() -> Self {
Self::default()
}
pub fn with_ttl_ms(mut self, ttl_ms: impl Into<Option<u64>>) -> Self {
self.ttl_ms = ttl_ms.into();
self
}
pub fn with_poll_interval_ms(mut self, poll_interval_ms: u64) -> Self {
self.poll_interval_ms = Some(poll_interval_ms);
self
}
pub fn with_status_message(mut self, message: impl Into<String>) -> Self {
self.status_message = Some(message.into());
self
}
}
#[derive(Clone, Default)]
pub struct TaskManager {
inner: Arc<Mutex<TaskManagerInner>>,
}
impl TaskManager {
pub fn new() -> Self {
Self::default()
}
pub fn spawn<F>(&self, options: TaskOptions, make_future: F) -> Task
where
F: FnOnce(TaskContext) -> TaskFuture,
{
let task_id = uuid::Uuid::new_v4().to_string();
let now = current_timestamp();
let mut task = Task::new(task_id.clone(), TaskStatus::Working, now.clone(), now);
task.ttl_ms = options.ttl_ms;
task.poll_interval_ms = options.poll_interval_ms;
task.status_message = options.status_message;
let entry = TaskEntry {
task: task.clone(),
terminal: None,
terminal_at: None,
pending_inputs: HashMap::new(),
used_input_keys: std::collections::HashSet::new(),
cancel_requested: false,
cancel_signal: tokio::sync::watch::channel(false).0,
created: Instant::now(),
join_handle: None,
};
{
let mut inner = self.inner.lock().expect("task manager lock poisoned");
Self::sweep_expired(&mut inner);
inner.tasks.insert(task_id.clone(), entry);
}
let context = TaskContext {
task_id: task_id.clone(),
inner: self.inner.clone(),
};
let future = make_future(context);
let inner = self.inner.clone();
let id_for_task = task_id.clone();
let originating_request = crate::service::ORIGINATING_REQUEST
.try_with(|id| id.clone())
.ok();
let handle = tokio::spawn(async move {
let result = run_task_operation(originating_request, future).await;
let mut inner = inner.lock().expect("task manager lock poisoned");
if let Some(entry) = inner.tasks.get_mut(&id_for_task) {
if entry.terminal.is_none() {
entry.terminal = Some(match result {
Ok(result) => TaskPayload::Completed {
result: result_to_object(&result),
},
Err(TaskExit::Cancelled) => TaskPayload::Cancelled,
Err(TaskExit::Error(error)) => TaskPayload::Failed {
error: error_to_object(&error),
},
});
entry.terminal_at = Some(Instant::now());
entry.pending_inputs.clear();
entry.touch();
entry.task.status = entry.current_status();
}
entry.join_handle = None;
}
});
match self
.inner
.lock()
.expect("task manager lock poisoned")
.tasks
.get_mut(&task_id)
{
Some(entry) => {
if entry.terminal.is_none() {
entry.join_handle = Some(handle);
}
}
None => handle.abort(),
}
task
}
pub fn get_task(&self, task_id: &str) -> Result<DetailedTask, McpError> {
let mut inner = self.inner.lock().expect("task manager lock poisoned");
Self::sweep_expired(&mut inner);
let entry = inner
.tasks
.get_mut(task_id)
.ok_or_else(|| unknown_task(task_id))?;
entry.task.status = entry.current_status();
Ok(entry.detailed())
}
pub fn update_task(
&self,
task_id: &str,
input_responses: impl IntoIterator<Item = (String, serde_json::Value)>,
) -> Result<(), McpError> {
let mut inner = self.inner.lock().expect("task manager lock poisoned");
Self::sweep_expired(&mut inner);
let entry = inner
.tasks
.get_mut(task_id)
.ok_or_else(|| unknown_task(task_id))?;
for (key, value) in input_responses {
if let Some((_, tx)) = entry.pending_inputs.remove(&key) {
let _ = tx.send(value);
}
}
entry.touch();
entry.task.status = entry.current_status();
Ok(())
}
pub fn cancel_task(&self, task_id: &str) -> Result<(), McpError> {
let mut inner = self.inner.lock().expect("task manager lock poisoned");
Self::sweep_expired(&mut inner);
let entry = inner
.tasks
.get_mut(task_id)
.ok_or_else(|| unknown_task(task_id))?;
entry.cancel_requested = true;
let _ = entry.cancel_signal.send(true);
if entry.terminal.is_none() {
entry.pending_inputs.clear();
entry.touch();
entry.task.status = entry.current_status();
}
Ok(())
}
pub fn running_task_count(&self) -> usize {
let mut inner = self.inner.lock().expect("task manager lock poisoned");
Self::sweep_expired(&mut inner);
inner
.tasks
.values()
.filter(|e| e.terminal.is_none())
.count()
}
pub fn shutdown(&self) {
let mut inner = self.inner.lock().expect("task manager lock poisoned");
for (_, mut entry) in inner.tasks.drain() {
if let Some(handle) = entry.join_handle.take() {
handle.abort();
}
}
}
fn sweep_expired(inner: &mut TaskManagerInner) {
for entry in inner.tasks.values_mut() {
if entry.terminal.is_none()
&& let Some(ttl_ms) = entry.task.ttl_ms
&& entry.created.elapsed().as_millis() >= u128::from(ttl_ms)
{
if let Some(handle) = entry.join_handle.take() {
handle.abort();
}
entry.terminal = Some(TaskPayload::Failed {
error: error_to_object(&McpError::internal_error(
"task expired: TTL elapsed before completion".to_string(),
None,
)),
});
entry.terminal_at = Some(Instant::now());
entry.pending_inputs.clear();
entry.touch();
entry.task.status = TaskStatus::Failed;
}
}
inner.tasks.retain(|_, entry| {
let (Some(ttl_ms), Some(terminal_at)) = (entry.task.ttl_ms, entry.terminal_at) else {
return true;
};
terminal_at.elapsed().as_millis() < u128::from(ttl_ms)
});
}
}
fn unknown_task(task_id: &str) -> McpError {
McpError::invalid_params(format!("unknown task: {task_id}"), None)
}
async fn run_task_operation(
originating_request: Option<crate::model::RequestId>,
future: TaskFuture,
) -> Result<CallToolResult, TaskExit> {
match originating_request {
Some(id) => crate::service::ORIGINATING_REQUEST.scope(id, future).await,
None => future.await,
}
}
fn result_to_object(result: &CallToolResult) -> JsonObject {
match serde_json::to_value(result) {
Ok(serde_json::Value::Object(map)) => map,
_ => JsonObject::new(),
}
}
fn error_to_object(error: &McpError) -> JsonObject {
match serde_json::to_value(error) {
Ok(serde_json::Value::Object(map)) => map,
_ => JsonObject::new(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::ContentBlock;
fn ok_result(text: &str) -> CallToolResult {
CallToolResult::success(vec![ContentBlock::text(text.to_string())])
}
#[tokio::test]
async fn task_completes_and_result_is_inlined() {
let manager = TaskManager::new();
let task = manager.spawn(TaskOptions::default(), |_ctx| {
Box::pin(async { Ok(ok_result("42")) })
});
assert_eq!(task.status, TaskStatus::Working);
manager.get_task(&task.task_id).unwrap();
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
if detailed.status() == TaskStatus::Completed {
match detailed.payload {
TaskPayload::Completed { result } => {
assert!(result.contains_key("content"));
return;
}
other => panic!("unexpected payload: {other:?}"),
}
}
}
panic!("task did not complete");
}
#[tokio::test]
async fn cancel_settles_to_cancelled_when_operation_honors_it() {
let manager = TaskManager::new();
let task = manager.spawn(TaskOptions::default(), |ctx| {
Box::pin(async move {
tokio::select! {
_ = ctx.cancelled() => Err(TaskExit::Cancelled),
_ = tokio::time::sleep(std::time::Duration::from_secs(60)) => {
Ok(ok_result("never"))
}
}
})
});
manager.cancel_task(&task.task_id).unwrap();
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
if detailed.status().is_terminal() {
assert_eq!(detailed.status(), TaskStatus::Cancelled);
return;
}
}
panic!("task did not settle after cancel");
}
#[tokio::test]
async fn post_cancel_unrelated_error_settles_as_failed() {
let manager = TaskManager::new();
let task = manager.spawn(TaskOptions::default(), |ctx| {
Box::pin(async move {
ctx.cancelled().await;
Err(TaskExit::Error(McpError::internal_error(
"database write failed",
None,
)))
})
});
manager.cancel_task(&task.task_id).unwrap();
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
if detailed.status().is_terminal() {
assert_eq!(detailed.status(), TaskStatus::Failed);
match detailed.payload {
TaskPayload::Failed { error } => {
assert!(
error.get("message").is_some_and(|m| m
.as_str()
.is_some_and(|s| s.contains("database write failed"))),
"error payload should be preserved: {error:?}"
);
}
other => panic!("unexpected payload: {other:?}"),
}
return;
}
}
panic!("task did not settle after cancel");
}
#[tokio::test]
async fn cancel_is_cooperative_and_lets_the_operation_clean_up() {
let manager = TaskManager::new();
let (cleanup_tx, cleanup_rx) = oneshot::channel::<&'static str>();
let task = manager.spawn(TaskOptions::default(), |ctx| {
Box::pin(async move {
ctx.cancelled().await;
let _ = cleanup_tx.send("cleaned up");
Ok(ok_result("finished despite cancel"))
})
});
manager.cancel_task(&task.task_id).unwrap();
let detailed = manager.get_task(&task.task_id).unwrap();
assert!(
!detailed.status().is_terminal(),
"cancel must not force terminal state"
);
let cleanup = tokio::time::timeout(std::time::Duration::from_secs(5), cleanup_rx)
.await
.expect("cleanup should not time out")
.expect("cleanup channel should not be dropped");
assert_eq!(cleanup, "cleaned up");
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
if detailed.status().is_terminal() {
assert_eq!(detailed.status(), TaskStatus::Completed);
return;
}
}
panic!("task did not settle after cancel");
}
#[tokio::test]
async fn cancel_wakes_parked_input_requests() {
let manager = TaskManager::new();
let (exit_tx, exit_rx) = oneshot::channel::<&'static str>();
let task = manager.spawn(TaskOptions::default(), |ctx| {
Box::pin(async move {
let request: InputRequest = serde_json::from_value(serde_json::json!({
"method": "elicitation/create",
"params": {
"message": "Waiting forever",
"requestedSchema": {"type": "object", "properties": {}}
}
}))
.map_err(|e| McpError::internal_error(e.to_string(), None))?;
let err = ctx.request_input("k1", request).await.unwrap_err();
let _ = exit_tx.send("woken");
Err(err)
})
});
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
if manager.get_task(&task.task_id).unwrap().status() == TaskStatus::InputRequired {
break;
}
}
manager.cancel_task(&task.task_id).unwrap();
let woken = tokio::time::timeout(std::time::Duration::from_secs(5), exit_rx)
.await
.expect("parked operation should be woken by cancel")
.expect("exit channel should not be dropped");
assert_eq!(woken, "woken");
assert_eq!(
manager.get_task(&task.task_id).unwrap().status(),
TaskStatus::Cancelled
);
}
#[tokio::test]
async fn unknown_task_is_invalid_params() {
let manager = TaskManager::new();
for err in [
manager.get_task("nope").unwrap_err(),
manager.cancel_task("nope").unwrap_err(),
manager.update_task("nope", []).unwrap_err(),
] {
assert_eq!(err.code, crate::model::ErrorCode::INVALID_PARAMS);
}
}
#[tokio::test]
async fn terminal_tasks_are_evicted_after_retention_window() {
let manager = TaskManager::new();
let task = manager.spawn(TaskOptions::new().with_ttl_ms(50), |_ctx| {
Box::pin(async { Ok(ok_result("fast")) })
});
let mut completed = false;
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(5)).await;
if manager.get_task(&task.task_id).unwrap().status() == TaskStatus::Completed {
completed = true;
break;
}
}
assert!(completed, "task should have completed");
tokio::time::sleep(std::time::Duration::from_millis(120)).await;
let err = manager.get_task(&task.task_id).unwrap_err();
assert_eq!(err.code, crate::model::ErrorCode::INVALID_PARAMS);
assert_eq!(manager.running_task_count(), 0);
}
#[tokio::test]
async fn abandoned_tasks_are_swept_by_other_entry_points() {
let manager = TaskManager::new();
let abandoned = manager.spawn(TaskOptions::new().with_ttl_ms(10), |_ctx| {
Box::pin(async {
tokio::time::sleep(std::time::Duration::from_secs(60)).await;
Ok(ok_result("never"))
})
});
tokio::time::sleep(std::time::Duration::from_millis(40)).await;
let _ = manager.spawn(TaskOptions::default(), |_ctx| {
Box::pin(async { Ok(ok_result("other")) })
});
tokio::time::sleep(std::time::Duration::from_millis(40)).await;
let _ = manager.spawn(TaskOptions::default(), |_ctx| {
Box::pin(async { Ok(ok_result("other2")) })
});
let err = manager.get_task(&abandoned.task_id).unwrap_err();
assert_eq!(err.code, crate::model::ErrorCode::INVALID_PARAMS);
}
#[tokio::test]
async fn running_task_count_sweeps_expired_tasks() {
let manager = TaskManager::new();
let _task = manager.spawn(TaskOptions::new().with_ttl_ms(10), |_ctx| {
Box::pin(async {
tokio::time::sleep(std::time::Duration::from_secs(60)).await;
Ok(ok_result("never"))
})
});
assert_eq!(manager.running_task_count(), 1);
tokio::time::sleep(std::time::Duration::from_millis(30)).await;
assert_eq!(manager.running_task_count(), 0);
}
#[tokio::test]
async fn unlimited_ttl_tasks_are_retained() {
let manager = TaskManager::new();
let task = manager.spawn(TaskOptions::new().with_ttl_ms(None), |_ctx| {
Box::pin(async { Ok(ok_result("kept")) })
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let _ = manager.spawn(TaskOptions::default(), |_ctx| {
Box::pin(async { Ok(ok_result("other")) })
});
assert_eq!(
manager.get_task(&task.task_id).unwrap().status(),
TaskStatus::Completed
);
}
#[tokio::test]
async fn ttl_expiry_fails_task() {
let manager = TaskManager::new();
let task = manager.spawn(
TaskOptions {
ttl_ms: Some(10),
..Default::default()
},
|_ctx| {
Box::pin(async {
tokio::time::sleep(std::time::Duration::from_secs(60)).await;
Ok(ok_result("never"))
})
},
);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
assert_eq!(detailed.status(), TaskStatus::Failed);
}
#[tokio::test]
async fn input_required_roundtrip() {
let manager = TaskManager::new();
let task = manager.spawn(TaskOptions::default(), |ctx| {
Box::pin(async move {
let request: InputRequest = serde_json::from_value(serde_json::json!({
"method": "elicitation/create",
"params": {
"message": "What is your name?",
"requestedSchema": {"type": "object", "properties": {}}
}
}))
.map_err(|e| McpError::internal_error(e.to_string(), None))?;
let response = ctx.request_input("name-1", request).await?;
let name = response
.get("content")
.and_then(|c| c.get("name"))
.and_then(|v| v.as_str())
.unwrap_or("unknown")
.to_string();
Ok(ok_result(&format!("hello {name}")))
})
});
let mut saw_input_required = false;
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
if let TaskPayload::InputRequired { input_requests } = &detailed.payload {
assert!(input_requests.contains_key("name-1"));
saw_input_required = true;
break;
}
}
assert!(saw_input_required, "task never reached input_required");
manager
.update_task(
&task.task_id,
[(
"name-1".to_string(),
serde_json::json!({"action": "accept", "content": {"name": "Ada"}}),
)],
)
.unwrap();
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let detailed = manager.get_task(&task.task_id).unwrap();
if detailed.status() == TaskStatus::Completed {
return;
}
}
panic!("task did not complete after input response");
}
#[tokio::test]
async fn task_operation_reestablishes_request_association_scope() {
use crate::{
model::RequestId,
service::{ORIGINATING_REQUEST, in_request_handler_scope},
};
let manager = TaskManager::new();
let observed = Arc::new(Mutex::new(None::<bool>));
let observed_in_task = observed.clone();
ORIGINATING_REQUEST
.scope(RequestId::Number(7), async {
manager.spawn(TaskOptions::default(), move |_ctx| {
let observed_in_task = observed_in_task.clone();
Box::pin(async move {
*observed_in_task.lock().unwrap() = Some(in_request_handler_scope());
Ok(ok_result("done"))
})
})
})
.await;
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
if let Some(scoped) = *observed.lock().unwrap() {
assert!(
scoped,
"task operation must run inside the originating request's association scope"
);
return;
}
}
panic!("task operation did not run");
}
#[tokio::test]
async fn task_operation_without_originating_request_is_unscoped() {
use crate::service::in_request_handler_scope;
let manager = TaskManager::new();
let observed = Arc::new(Mutex::new(None::<bool>));
let observed_in_task = observed.clone();
manager.spawn(TaskOptions::default(), move |_ctx| {
let observed_in_task = observed_in_task.clone();
Box::pin(async move {
*observed_in_task.lock().unwrap() = Some(in_request_handler_scope());
Ok(ok_result("done"))
})
});
for _ in 0..100 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
if let Some(scoped) = *observed.lock().unwrap() {
assert!(
!scoped,
"task operation started without an originating request must remain unscoped"
);
return;
}
}
panic!("task operation did not run");
}
}