use crate::error::{JsonRpcError, McpError};
use crate::types::task::{
CancelTaskResult, GetTaskResult, ListTasksResult, Task, TaskId, TaskStatus,
};
use event_listener::Event;
use serde_json::Value;
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, RwLock};
use std::task::{Context as TaskContext, Poll};
use std::time::Instant;
pub const RELATED_TASK_META_KEY: &str = "io.modelcontextprotocol/related-task";
#[derive(Clone)]
pub struct CancellationToken {
cancelled: Arc<AtomicBool>,
event: Arc<Event>,
}
impl CancellationToken {
#[must_use]
pub fn new() -> Self {
Self {
cancelled: Arc::new(AtomicBool::new(false)),
event: Arc::new(Event::new()),
}
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::SeqCst)
}
pub fn cancel(&self) {
self.cancelled.store(true, Ordering::SeqCst);
self.event.notify(usize::MAX);
}
#[must_use]
pub fn cancelled(&self) -> CancelledFuture {
CancelledFuture::new(self.cancelled.clone(), self.event.clone())
}
}
impl std::fmt::Debug for CancellationToken {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CancellationToken")
.field("cancelled", &self.is_cancelled())
.finish()
}
}
impl Default for CancellationToken {
fn default() -> Self {
Self::new()
}
}
pub struct CancelledFuture {
inner: Pin<Box<dyn Future<Output = ()> + Send>>,
}
impl CancelledFuture {
fn new(cancelled: Arc<AtomicBool>, event: Arc<Event>) -> Self {
Self {
inner: Box::pin(async move {
loop {
if cancelled.load(Ordering::SeqCst) {
return;
}
let listener = event.listen();
if cancelled.load(Ordering::SeqCst) {
return;
}
listener.await;
}
}),
}
}
}
impl Future for CancelledFuture {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut TaskContext<'_>) -> Poll<Self::Output> {
self.inner.as_mut().poll(cx)
}
}
#[derive(Debug, Clone)]
pub enum TaskPayload {
Success(Value),
Error(JsonRpcError),
}
#[derive(Debug, Clone)]
pub struct TaskState {
pub task: Task,
pub payload: Option<TaskPayload>,
pub cancel_token: CancellationToken,
pub last_access: Instant,
pub created: Instant,
terminal: Arc<Event>,
}
impl TaskState {
fn new(task: Task) -> Self {
let now = Instant::now();
Self {
task,
payload: None,
cancel_token: CancellationToken::new(),
last_access: now,
created: now,
terminal: Arc::new(Event::new()),
}
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.cancel_token.is_cancelled()
}
}
pub struct TaskHandle {
task_id: TaskId,
manager: Arc<TaskManager>,
}
impl TaskHandle {
#[must_use]
pub const fn id(&self) -> &TaskId {
&self.task_id
}
#[must_use]
pub fn task(&self) -> Option<Task> {
self.manager.get(&self.task_id).map(|s| s.task)
}
#[must_use]
pub fn cancel_token(&self) -> Option<CancellationToken> {
self.manager.get(&self.task_id).map(|s| s.cancel_token)
}
pub fn mark_input_required(&self) -> Result<(), McpError> {
self.manager
.set_status(&self.task_id, TaskStatus::InputRequired, None)
}
pub fn complete(&self, payload: Value) -> Result<(), McpError> {
self.manager.finish(
&self.task_id,
TaskStatus::Completed,
Some(TaskPayload::Success(payload)),
None,
)
}
pub fn fail(&self, message: impl Into<String>) -> Result<(), McpError> {
let message = message.into();
self.manager.finish(
&self.task_id,
TaskStatus::Failed,
Some(TaskPayload::Error(JsonRpcError::internal_error(
message.clone(),
))),
Some(message),
)
}
pub fn fail_with_error(&self, error: JsonRpcError) -> Result<(), McpError> {
let message = error.message.clone();
self.manager.finish(
&self.task_id,
TaskStatus::Failed,
Some(TaskPayload::Error(error)),
Some(message),
)
}
pub fn fail_with_result(
&self,
payload: Value,
message: Option<String>,
) -> Result<(), McpError> {
self.manager.finish(
&self.task_id,
TaskStatus::Failed,
Some(TaskPayload::Success(payload)),
message,
)
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.manager
.get(&self.task_id)
.is_none_or(|s| s.is_cancelled())
}
pub async fn cancelled(&self) {
if let Some(state) = self.manager.get(&self.task_id) {
state.cancel_token.cancelled().await;
}
}
}
pub const DEFAULT_TASK_TTL_MS: u64 = 60 * 60 * 1000;
#[derive(Debug, Clone)]
pub struct TaskEvent {
pub task: Task,
pub previous_status: TaskStatus,
}
pub trait TaskObserver: Send + Sync + std::fmt::Debug {
fn on_task_event(&self, event: &TaskEvent);
}
#[derive(Debug)]
pub struct TaskManager {
tasks: RwLock<HashMap<TaskId, TaskState>>,
default_ttl_ms: Option<u64>,
observer: std::sync::OnceLock<Arc<dyn TaskObserver>>,
default_poll_interval_ms: Option<u64>,
}
impl Default for TaskManager {
fn default() -> Self {
Self::new()
}
}
impl TaskManager {
#[must_use]
pub fn new() -> Self {
Self::with_default_ttl(Some(DEFAULT_TASK_TTL_MS))
}
#[must_use]
pub fn with_default_ttl(default_ttl_ms: Option<u64>) -> Self {
Self {
tasks: RwLock::new(HashMap::new()),
default_ttl_ms,
observer: std::sync::OnceLock::new(),
default_poll_interval_ms: None,
}
}
#[must_use]
pub const fn with_poll_interval(mut self, poll_interval_ms: Option<u64>) -> Self {
self.default_poll_interval_ms = poll_interval_ms;
self
}
pub fn set_observer(&self, observer: Arc<dyn TaskObserver>) -> Result<(), McpError> {
self.observer
.set(observer)
.map_err(|_| McpError::internal("task observer already registered"))
}
fn emit(&self, event: &TaskEvent) {
if let Some(observer) = self.observer.get() {
observer.on_task_event(event);
}
}
pub fn create(self: &Arc<Self>, ttl: Option<u64>) -> TaskHandle {
self.cleanup_expired();
let mut task = Task::create();
task.ttl = ttl.or(self.default_ttl_ms);
task.poll_interval = self.default_poll_interval_ms;
let task_id = task.task_id.clone();
if let Ok(mut tasks) = self.tasks.write() {
tasks.insert(task_id.clone(), TaskState::new(task));
}
TaskHandle {
task_id,
manager: Arc::clone(self),
}
}
#[must_use]
pub fn get(&self, id: &TaskId) -> Option<TaskState> {
self.tasks.read().ok()?.get(id).cloned()
}
#[must_use]
pub fn list(&self) -> Vec<Task> {
self.tasks
.read()
.map(|tasks| tasks.values().map(|s| s.task.clone()).collect())
.unwrap_or_default()
}
#[must_use]
pub fn payload(&self, id: &TaskId) -> Option<TaskPayload> {
self.tasks.read().ok()?.get(id)?.payload.clone()
}
pub async fn wait_terminal(&self, id: &TaskId) -> Option<TaskState> {
loop {
let listener = {
let tasks = self.tasks.read().ok()?;
let state = tasks.get(id)?;
if state.task.status.is_terminal() {
return Some(state.clone());
}
state.terminal.listen()
};
listener.await;
}
}
pub fn cancel(&self, id: &TaskId) -> Result<(), McpError> {
let event = {
let mut tasks = self
.tasks
.write()
.map_err(|_| McpError::internal("Failed to acquire task lock"))?;
if let Some(state) = tasks.get_mut(id) {
if state.task.status.is_terminal() {
return Err(McpError::invalid_params(
"tasks/cancel",
format!(
"Cannot cancel task: already in terminal status '{}'",
state.task.status
),
));
}
let previous_status = state.task.status;
state.cancel_token.cancel();
state.task.set_status(TaskStatus::Cancelled);
state.last_access = Instant::now();
state.terminal.notify(usize::MAX);
TaskEvent {
task: state.task.clone(),
previous_status,
}
} else {
return Err(McpError::invalid_params(
"tasks/cancel",
format!("Unknown task: {}", id.as_str()),
));
}
};
self.emit(&event);
Ok(())
}
fn set_status(
&self,
id: &TaskId,
status: TaskStatus,
message: Option<String>,
) -> Result<(), McpError> {
let event = {
let mut tasks = self
.tasks
.write()
.map_err(|_| McpError::internal("Failed to acquire task lock"))?;
if let Some(state) = tasks.get_mut(id) {
if state.task.status.is_terminal() {
return Err(McpError::invalid_params(
"tasks/get",
format!(
"task {} is already terminal ('{}')",
id.as_str(),
state.task.status
),
));
}
let previous_status = state.task.status;
state.task.set_status(status);
if message.is_some() {
state.task.status_message = message;
}
state.last_access = Instant::now();
if status.is_terminal() {
state.terminal.notify(usize::MAX);
}
TaskEvent {
task: state.task.clone(),
previous_status,
}
} else {
return Err(McpError::invalid_params(
"tasks/get",
format!("Unknown task: {}", id.as_str()),
));
}
};
self.emit(&event);
Ok(())
}
fn finish(
&self,
id: &TaskId,
status: TaskStatus,
payload: Option<TaskPayload>,
message: Option<String>,
) -> Result<(), McpError> {
let event = {
let mut tasks = self
.tasks
.write()
.map_err(|_| McpError::internal("Failed to acquire task lock"))?;
if let Some(state) = tasks.get_mut(id) {
if state.task.status.is_terminal() {
return Err(McpError::invalid_params(
"tasks/result",
format!(
"task {} is already terminal ('{}')",
id.as_str(),
state.task.status
),
));
}
let previous_status = state.task.status;
state.task.set_status(status);
if message.is_some() {
state.task.status_message = message;
}
state.payload = payload;
state.last_access = Instant::now();
state.terminal.notify(usize::MAX);
TaskEvent {
task: state.task.clone(),
previous_status,
}
} else {
return Err(McpError::invalid_params(
"tasks/result",
format!("Unknown task: {}", id.as_str()),
));
}
};
self.emit(&event);
Ok(())
}
pub fn cleanup(&self, max_age: std::time::Duration) {
if let Ok(mut tasks) = self.tasks.write() {
tasks.retain(|_, state| {
let is_terminal = state.task.status.is_terminal();
!is_terminal || state.last_access.elapsed() < max_age
});
}
}
pub fn cleanup_expired(&self) {
if let Ok(mut tasks) = self.tasks.write() {
tasks.retain(|_, state| {
if !state.task.status.is_terminal() {
return true;
}
match state.task.ttl {
Some(ttl_ms) => {
state.created.elapsed() < std::time::Duration::from_millis(ttl_ms)
}
None => true,
}
});
}
}
}
#[must_use]
pub fn inject_related_task(mut payload: Value, id: &TaskId) -> Value {
if let Value::Object(map) = &mut payload {
if let Value::Object(meta) = map
.entry("_meta")
.or_insert_with(|| Value::Object(serde_json::Map::new()))
{
meta.insert(
RELATED_TASK_META_KEY.to_string(),
serde_json::json!({ "taskId": id.as_str() }),
);
}
}
payload
}
pub async fn route_task_store(
store: &TaskManager,
method: &str,
params: Option<&Value>,
) -> TaskRoute {
TaskRoute::from_parts(
route_task_store_inner(store, method, params).await,
method,
params,
)
}
#[derive(Debug)]
pub enum TaskRoute {
NotTaskMethod,
UnownedTask {
method: String,
task_id: String,
},
Handled(Result<Value, McpError>),
}
impl TaskRoute {
fn from_parts(
inner: Option<Result<Value, McpError>>,
method: &str,
params: Option<&Value>,
) -> Self {
match inner {
Some(result) => Self::Handled(result),
None if is_task_method(method) => Self::UnownedTask {
method: method.to_string(),
task_id: params
.and_then(|p| p.get("taskId"))
.and_then(|v| v.as_str())
.unwrap_or("<missing>")
.to_string(),
},
None => Self::NotTaskMethod,
}
}
#[must_use]
pub fn or_unknown_task(self) -> Option<Result<Value, McpError>> {
match self {
Self::NotTaskMethod => None,
Self::UnownedTask { method, task_id } => Some(Err(McpError::invalid_params(
method,
format!("Unknown task: {task_id}"),
))),
Self::Handled(result) => Some(result),
}
}
}
fn is_task_method(method: &str) -> bool {
matches!(method, "tasks/get" | "tasks/result" | "tasks/cancel")
}
async fn route_task_store_inner(
store: &TaskManager,
method: &str,
params: Option<&Value>,
) -> Option<Result<Value, McpError>> {
store.cleanup_expired();
let task_id = || {
params
.and_then(|p| p.get("taskId"))
.and_then(|v| v.as_str())
.map(TaskId::new)
};
match method {
"tasks/list" => {
let result = ListTasksResult::from(store.list());
Some(Ok(serde_json::to_value(result).unwrap_or_default()))
}
"tasks/get" => {
let Some(id) = task_id() else {
return Some(Err(McpError::invalid_params("tasks/get", "missing taskId")));
};
store.get(&id).map(|s| {
let result = GetTaskResult::from(s.task);
Ok(serde_json::to_value(result).unwrap_or_default())
})
}
"tasks/result" => {
let Some(id) = task_id() else {
return Some(Err(McpError::invalid_params(
"tasks/result",
"missing taskId",
)));
};
store.get(&id)?;
let Some(state) = store.wait_terminal(&id).await else {
return Some(Err(McpError::invalid_params(
"tasks/result",
format!("Task has expired: {}", id.as_str()),
)));
};
match state.payload {
Some(TaskPayload::Success(payload)) => Some(Ok(inject_related_task(payload, &id))),
Some(TaskPayload::Error(error)) => Some(Err(McpError::JsonRpc(error))),
None => Some(Err(McpError::invalid_params(
"tasks/result",
format!(
"task {} ended {} with no result",
id.as_str(),
state.task.status
),
))),
}
}
"tasks/cancel" => {
let Some(id) = task_id() else {
return Some(Err(McpError::invalid_params(
"tasks/cancel",
"missing taskId",
)));
};
if store.get(&id).is_some() {
if let Err(e) = store.cancel(&id) {
return Some(Err(e));
}
Some(Ok(store
.get(&id)
.map(|s| {
let result = CancelTaskResult::from(s.task);
serde_json::to_value(result).unwrap_or_default()
})
.unwrap_or_default()))
} else {
None
}
}
_ => None,
}
}
#[cfg(test)]
mod tests {
#[test]
fn poll_interval_is_absent_by_default_and_set_when_configured() {
let plain = Arc::new(TaskManager::new());
let a = plain.create(None);
assert_eq!(a.task().expect("task").poll_interval, None);
let suggesting = Arc::new(TaskManager::new().with_poll_interval(Some(250)));
let b = suggesting.create(None);
let task = b.task().expect("task");
assert_eq!(task.poll_interval, Some(250));
let wire = serde_json::to_value(&task).expect("serialize");
assert_eq!(wire["pollInterval"], 250);
}
use super::*;
#[derive(Debug, Default)]
struct Collector {
events: std::sync::Mutex<Vec<TaskEvent>>,
}
impl TaskObserver for Collector {
fn on_task_event(&self, event: &TaskEvent) {
if let Ok(mut events) = self.events.lock() {
events.push(event.clone());
}
}
}
impl Collector {
fn transitions(&self) -> Vec<(TaskStatus, TaskStatus)> {
self.events
.lock()
.map(|e| {
e.iter()
.map(|e| (e.previous_status, e.task.status))
.collect()
})
.unwrap_or_default()
}
}
#[test]
fn test_observer_sees_every_transition() -> Result<(), Box<dyn std::error::Error>> {
let manager = Arc::new(TaskManager::new());
let collector = Arc::new(Collector::default());
manager.set_observer(collector.clone())?;
let a = manager.create(None);
a.mark_input_required()?;
a.complete(serde_json::json!({"ok": true}))?;
let b = manager.create(None);
manager.cancel(b.id())?;
assert_eq!(
collector.transitions(),
vec![
(TaskStatus::Working, TaskStatus::InputRequired),
(TaskStatus::InputRequired, TaskStatus::Completed),
(TaskStatus::Working, TaskStatus::Cancelled),
]
);
Ok(())
}
#[test]
fn test_observer_may_reenter_the_manager() -> Result<(), Box<dyn std::error::Error>> {
#[derive(Debug)]
struct Reentrant(std::sync::Weak<TaskManager>);
impl TaskObserver for Reentrant {
fn on_task_event(&self, event: &TaskEvent) {
if let Some(manager) = self.0.upgrade() {
assert!(manager.get(&event.task.task_id).is_some());
let _ = manager.list();
}
}
}
let manager = Arc::new(TaskManager::new());
manager.set_observer(Arc::new(Reentrant(Arc::downgrade(&manager))))?;
let handle = manager.create(None);
handle.complete(serde_json::json!({}))?;
assert_eq!(
manager.get(handle.id()).ok_or("not found")?.task.status,
TaskStatus::Completed
);
Ok(())
}
#[test]
fn test_observer_registers_at_most_once() {
let manager = Arc::new(TaskManager::new());
assert!(manager.set_observer(Arc::new(Collector::default())).is_ok());
assert!(
manager
.set_observer(Arc::new(Collector::default()))
.is_err()
);
}
#[test]
fn test_task_manager_create_and_list() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
assert!(!handle.is_cancelled());
let tasks = manager.list();
assert_eq!(tasks.len(), 1);
assert_eq!(tasks[0].status, TaskStatus::Working);
}
#[test]
fn test_task_complete_stores_payload() -> Result<(), Box<dyn std::error::Error>> {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let task_id = handle.id().clone();
handle.complete(serde_json::json!({"result": "ok"}))?;
let state = manager.get(&task_id).ok_or("Task not found")?;
assert_eq!(state.task.status, TaskStatus::Completed);
match manager.payload(&task_id) {
Some(TaskPayload::Success(v)) => {
assert_eq!(v, serde_json::json!({"result": "ok"}));
}
other => panic!("expected success payload, got {other:?}"),
}
Ok(())
}
#[test]
fn test_task_input_required_and_fail() -> Result<(), Box<dyn std::error::Error>> {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let task_id = handle.id().clone();
handle.mark_input_required()?;
assert_eq!(
manager.get(&task_id).ok_or("not found")?.task.status,
TaskStatus::InputRequired
);
handle.fail("boom")?;
let state = manager.get(&task_id).ok_or("not found")?;
assert_eq!(state.task.status, TaskStatus::Failed);
assert_eq!(state.task.status_message.as_deref(), Some("boom"));
Ok(())
}
#[test]
fn test_task_cancellation() -> Result<(), Box<dyn std::error::Error>> {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let task_id = handle.id().clone();
assert!(!handle.is_cancelled());
manager.cancel(&task_id)?;
assert!(handle.is_cancelled());
assert_eq!(
manager.get(&task_id).ok_or("not found")?.task.status,
TaskStatus::Cancelled
);
Ok(())
}
#[test]
fn omitted_ttl_is_materialized_to_default() {
let manager = Arc::new(TaskManager::with_default_ttl(Some(5000)));
assert_eq!(manager.create(None).task().unwrap().ttl, Some(5000));
assert_eq!(manager.create(Some(1234)).task().unwrap().ttl, Some(1234));
}
#[test]
fn cleanup_expired_evicts_old_terminal_task() {
let manager = Arc::new(TaskManager::with_default_ttl(Some(1)));
let handle = manager.create(None); let id = handle.id().clone();
handle.complete(serde_json::json!({})).unwrap();
std::thread::sleep(std::time::Duration::from_millis(20));
manager.cleanup_expired();
assert!(
manager.get(&id).is_none(),
"expired terminal task not evicted"
);
}
#[test]
fn cleanup_expired_keeps_fresh_terminal_task() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(Some(60_000));
let id = handle.id().clone();
handle.complete(serde_json::json!({})).unwrap();
manager.cleanup_expired();
assert!(
manager.get(&id).is_some(),
"fresh terminal task wrongly evicted"
);
}
#[test]
fn cleanup_expired_keeps_non_terminal_task() {
let manager = Arc::new(TaskManager::with_default_ttl(Some(1)));
let handle = manager.create(None); let id = handle.id().clone();
std::thread::sleep(std::time::Duration::from_millis(20));
manager.cleanup_expired();
assert!(
manager.get(&id).is_some(),
"non-terminal task must never be evicted"
);
}
#[test]
fn unlimited_ttl_is_never_evicted() {
let manager = Arc::new(TaskManager::with_default_ttl(None));
let handle = manager.create(None); let id = handle.id().clone();
assert_eq!(handle.task().unwrap().ttl, None);
handle.complete(serde_json::json!({})).unwrap();
std::thread::sleep(std::time::Duration::from_millis(20));
manager.cleanup_expired();
assert!(
manager.get(&id).is_some(),
"unlimited-ttl task wrongly evicted"
);
}
#[tokio::test]
async fn route_task_store_access_triggers_cleanup() {
let manager = Arc::new(TaskManager::with_default_ttl(Some(1)));
let handle = manager.create(None);
let id = handle.id().clone();
handle.complete(serde_json::json!({})).unwrap();
std::thread::sleep(std::time::Duration::from_millis(20));
let _ = route_task_store(&manager, "tasks/list", None).await;
assert!(manager.get(&id).is_none(), "access did not trigger cleanup");
}
fn result_params(id: &TaskId) -> Value {
serde_json::json!({ "taskId": id.as_str() })
}
#[tokio::test]
async fn tasks_result_returns_immediately_for_terminal_task() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
handle
.complete(serde_json::json!({ "answer": 42 }))
.unwrap();
let params = result_params(&id);
let result = route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
.expect("owned task")
.expect("success");
assert_eq!(result["answer"], 42);
}
#[tokio::test]
async fn tasks_result_blocks_until_terminal() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
let completer = {
let manager = Arc::clone(&manager);
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let handle = TaskHandle {
task_id: id,
manager,
};
handle
.complete(serde_json::json!({ "late": true }))
.unwrap();
})
};
let params = result_params(handle.id());
let started = std::time::Instant::now();
let result = tokio::time::timeout(std::time::Duration::from_secs(5), async {
route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
})
.await
.expect("must not hang")
.expect("owned task")
.expect("success");
assert_eq!(result["late"], true);
assert!(
started.elapsed() >= std::time::Duration::from_millis(40),
"must have blocked until completion"
);
completer.await.unwrap();
}
#[tokio::test]
async fn tasks_result_reproduces_stored_jsonrpc_error() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
let stored = JsonRpcError {
code: -32001,
message: "downstream exploded".to_string(),
data: Some(serde_json::json!({ "detail": "xyz" })),
};
handle.fail_with_error(stored.clone()).unwrap();
let params = result_params(&id);
let err = route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
.expect("owned task")
.expect_err("stored error");
let wire: JsonRpcError = (&err).into();
assert_eq!(wire.code, stored.code);
assert_eq!(wire.message, stored.message);
assert_eq!(wire.data, stored.data);
let state = manager.get(&id).unwrap();
assert_eq!(state.task.status, TaskStatus::Failed);
assert_eq!(
state.task.status_message.as_deref(),
Some("downstream exploded")
);
}
#[tokio::test]
async fn tasks_result_success_carries_related_task_meta() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
handle
.complete(serde_json::json!({ "ok": true, "_meta": { "keep": 1 } }))
.unwrap();
let params = result_params(&id);
let result = route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
.expect("owned task")
.expect("success");
assert_eq!(result["_meta"]["keep"], 1);
assert_eq!(
result["_meta"][RELATED_TASK_META_KEY]["taskId"],
id.as_str()
);
}
#[tokio::test]
async fn failed_task_with_success_payload_returns_it() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
handle
.fail_with_result(
serde_json::json!({ "isError": true, "content": [] }),
Some("tool reported an error".to_string()),
)
.unwrap();
assert_eq!(manager.get(&id).unwrap().task.status, TaskStatus::Failed);
let params = result_params(&id);
let result = route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
.expect("owned task")
.expect("isError result is still a successful JSON-RPC response");
assert_eq!(result["isError"], true);
assert_eq!(
result["_meta"][RELATED_TASK_META_KEY]["taskId"],
id.as_str()
);
}
#[tokio::test]
async fn tasks_result_for_cancelled_task_is_an_error() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
manager.cancel(&id).unwrap();
let params = result_params(&id);
let err = route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
.expect("owned task")
.expect_err("cancelled task has no result");
assert_eq!(err.code(), -32602);
}
#[tokio::test]
async fn cancel_unblocks_tasks_result_waiter() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
let canceller = {
let manager = Arc::clone(&manager);
let id = id.clone();
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
manager.cancel(&id).unwrap();
})
};
let params = result_params(&id);
let outcome = tokio::time::timeout(std::time::Duration::from_secs(5), async {
route_task_store(&manager, "tasks/result", Some(¶ms))
.await
.or_unknown_task()
})
.await
.expect("cancel must unblock the waiter")
.expect("owned task");
assert!(outcome.is_err(), "cancelled task has no result");
canceller.await.unwrap();
}
#[test]
fn cancelled_task_stays_cancelled_when_execution_completes() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
manager.cancel(&id).unwrap();
assert!(handle.complete(serde_json::json!({ "late": 1 })).is_err());
assert!(handle.fail("late failure").is_err());
let state = manager.get(&id).unwrap();
assert_eq!(state.task.status, TaskStatus::Cancelled);
assert!(state.payload.is_none(), "late outcome must be discarded");
}
#[tokio::test]
async fn cancel_on_terminal_task_is_invalid_params() {
let manager = Arc::new(TaskManager::new());
let handle = manager.create(None);
let id = handle.id().clone();
handle.complete(serde_json::json!({})).unwrap();
let err = manager.cancel(&id).expect_err("terminal cancel rejected");
assert_eq!(err.code(), -32602);
let params = result_params(&id);
let err = route_task_store(&manager, "tasks/cancel", Some(¶ms))
.await
.or_unknown_task()
.expect("owned task")
.expect_err("terminal cancel rejected");
assert_eq!(err.code(), -32602);
}
#[test]
fn test_cancellation_token() {
let token = CancellationToken::new();
assert!(!token.is_cancelled());
token.cancel();
assert!(token.is_cancelled());
}
#[test]
fn cancelled_future_parks_and_wakes_on_cancel() {
use std::sync::atomic::AtomicUsize;
use std::task::{Wake, Waker};
struct CountingWaker(AtomicUsize);
impl Wake for CountingWaker {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let mut cx = TaskContext::from_waker(&waker);
let token = CancellationToken::new();
let mut fut = Box::pin(token.cancelled());
assert_eq!(fut.as_mut().poll(&mut cx), Poll::Pending);
assert_eq!(
counter.0.load(Ordering::SeqCst),
0,
"cancelled future must park, not busy-spin (no self-wake)"
);
token.cancel();
assert!(
counter.0.load(Ordering::SeqCst) >= 1,
"cancel() must wake the parked waiter"
);
assert_eq!(fut.as_mut().poll(&mut cx), Poll::Ready(()));
}
#[test]
fn cancelled_future_ready_when_already_cancelled() {
let waker = std::task::Waker::noop();
let mut cx = TaskContext::from_waker(waker);
let token = CancellationToken::new();
token.cancel();
let mut fut = Box::pin(token.cancelled());
assert_eq!(fut.as_mut().poll(&mut cx), Poll::Ready(()));
}
}