use std::collections::{BTreeSet, HashMap};
use std::fmt::Write as _;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use crate::error::JsonRpcError;
use crate::protocol::{CallToolResult, InputRequests, InputResponses, TaskObject, TaskStatus};
const DEFAULT_TTL_MS: u64 = 300_000;
const DEFAULT_POLL_INTERVAL_MS: u64 = 2_000;
#[derive(Debug)]
pub struct Task {
pub id: String,
pub tool_name: String,
pub arguments: serde_json::Value,
pub status: TaskStatus,
pub created_at: Instant,
pub created_at_str: String,
pub last_updated_at_str: String,
pub ttl: u64,
pub poll_interval: u64,
pub status_message: Option<String>,
pub result: Option<CallToolResult>,
pub error: Option<JsonRpcError>,
pub owner: TaskOwner,
pub input_requests: InputRequests,
pub answered_input_keys: BTreeSet<String>,
pub superseded_input_keys: BTreeSet<String>,
pub cancellation_token: CancellationToken,
pub completed_at: Option<Instant>,
pub completion_notify: Arc<tokio::sync::Notify>,
}
impl Task {
fn new(
id: String,
tool_name: String,
arguments: serde_json::Value,
ttl: Option<u64>,
owner: TaskOwner,
) -> Self {
let cancelled = Arc::new(AtomicBool::new(false));
let now_str = chrono_now_iso8601();
Self {
id,
tool_name,
arguments,
status: TaskStatus::Working,
created_at: Instant::now(),
created_at_str: now_str.clone(),
last_updated_at_str: now_str,
ttl: ttl.unwrap_or(DEFAULT_TTL_MS),
poll_interval: DEFAULT_POLL_INTERVAL_MS,
status_message: Some("Task started".to_string()),
result: None,
error: None,
owner,
input_requests: InputRequests::new(),
answered_input_keys: BTreeSet::new(),
superseded_input_keys: BTreeSet::new(),
cancellation_token: CancellationToken { cancelled },
completed_at: None,
completion_notify: Arc::new(tokio::sync::Notify::new()),
}
}
pub fn to_task_object(&self) -> TaskObject {
TaskObject {
task_id: self.id.clone(),
status: self.status,
status_message: self.status_message.clone(),
created_at: self.created_at_str.clone(),
last_updated_at: self.last_updated_at_str.clone(),
ttl: Some(self.ttl),
poll_interval: Some(self.poll_interval),
result: None,
error: None,
meta: None,
}
}
pub fn is_expired(&self) -> bool {
self.created_at.elapsed() > Duration::from_millis(self.ttl)
}
pub fn outstanding_input_requests(&self) -> &InputRequests {
&self.input_requests
}
pub fn is_cancelled(&self) -> bool {
self.cancellation_token.is_cancelled()
}
}
pub fn generate_task_id() -> String {
let mut bytes = [0u8; 16];
getrandom::fill(&mut bytes).expect("system entropy source unavailable for task ID generation");
let mut id = String::with_capacity(2 * bytes.len());
for byte in bytes {
let _ = write!(id, "{byte:02x}");
}
id
}
pub type TaskOwner = Option<String>;
pub fn owner_matches(owner: &TaskOwner, principal: Option<&str>) -> bool {
owner.as_deref() == principal
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct AppliedInputResponses {
pub accepted: BTreeSet<String>,
pub ignored: BTreeSet<String>,
pub still_outstanding: BTreeSet<String>,
}
impl AppliedInputResponses {
pub fn is_complete(&self) -> bool {
self.still_outstanding.is_empty()
}
}
#[derive(Debug, Clone)]
pub struct CancellationToken {
cancelled: Arc<AtomicBool>,
}
impl CancellationToken {
pub fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::Relaxed)
}
pub fn cancel(&self) {
self.cancelled.store(true, Ordering::Relaxed);
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum TaskStoreError {
#[error("encode error: {0}")]
Encode(String),
#[error("decode error: {0}")]
Decode(String),
#[error("backend error: {0}")]
Backend(String),
}
pub type Result<T> = std::result::Result<T, TaskStoreError>;
pub type TaskSnapshot = (TaskObject, Option<CallToolResult>, Option<JsonRpcError>);
#[async_trait]
pub trait TaskStore: Send + Sync + 'static {
async fn create_task(
&self,
tool_name: &str,
arguments: serde_json::Value,
ttl: Option<u64>,
owner: TaskOwner,
) -> Result<(String, CancellationToken)>;
async fn task_owner(&self, task_id: &str) -> Result<Option<TaskOwner>>;
async fn get_task(&self, task_id: &str) -> Result<Option<TaskObject>>;
async fn get_task_result(&self, task_id: &str) -> Result<Option<TaskSnapshot>>;
async fn wait_for_completion(&self, task_id: &str) -> Result<Option<TaskSnapshot>>;
async fn list_tasks(&self, status_filter: Option<TaskStatus>) -> Result<Vec<TaskObject>>;
async fn require_input(
&self,
task_id: &str,
requests: InputRequests,
message: Option<&str>,
) -> Result<bool>;
async fn outstanding_input_requests(&self, task_id: &str) -> Result<Option<InputRequests>>;
async fn apply_input_responses(
&self,
task_id: &str,
responses: InputResponses,
) -> Result<Option<AppliedInputResponses>>;
async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> Result<bool>;
async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool>;
async fn fail_task(&self, task_id: &str, error: JsonRpcError) -> Result<bool>;
async fn cancel_task(&self, task_id: &str, reason: Option<&str>) -> Result<Option<TaskObject>>;
}
#[derive(Debug, Clone)]
pub struct MemoryTaskStore {
tasks: Arc<RwLock<HashMap<String, Task>>>,
}
impl Default for MemoryTaskStore {
fn default() -> Self {
Self::new()
}
}
impl MemoryTaskStore {
pub fn new() -> Self {
Self {
tasks: Arc::new(RwLock::new(HashMap::new())),
}
}
pub fn cleanup_expired(&self) -> usize {
if let Ok(mut tasks) = self.tasks.write() {
let before = tasks.len();
tasks.retain(|_, t| !t.is_expired());
before - tasks.len()
} else {
0
}
}
#[cfg(test)]
pub fn len(&self) -> usize {
if let Ok(tasks) = self.tasks.read() {
tasks.len()
} else {
0
}
}
#[cfg(test)]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[async_trait]
impl TaskStore for MemoryTaskStore {
async fn create_task(
&self,
tool_name: &str,
arguments: serde_json::Value,
ttl: Option<u64>,
owner: TaskOwner,
) -> Result<(String, CancellationToken)> {
let id = generate_task_id();
let task = Task::new(id.clone(), tool_name.to_string(), arguments, ttl, owner);
let token = task.cancellation_token.clone();
if let Ok(mut tasks) = self.tasks.write() {
tasks.insert(id.clone(), task);
}
Ok((id, token))
}
async fn get_task(&self, task_id: &str) -> Result<Option<TaskObject>> {
Ok(if let Ok(tasks) = self.tasks.read() {
tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| t.to_task_object())
} else {
None
})
}
async fn task_owner(&self, task_id: &str) -> Result<Option<TaskOwner>> {
Ok(if let Ok(tasks) = self.tasks.read() {
tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| t.owner.clone())
} else {
None
})
}
async fn get_task_result(&self, task_id: &str) -> Result<Option<TaskSnapshot>> {
Ok(if let Ok(tasks) = self.tasks.read() {
tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| (t.to_task_object(), t.result.clone(), t.error.clone()))
} else {
None
})
}
async fn wait_for_completion(&self, task_id: &str) -> Result<Option<TaskSnapshot>> {
let notify = {
let Ok(tasks) = self.tasks.read() else {
return Ok(None);
};
let Some(task) = tasks.get(task_id).filter(|t| !t.is_expired()) else {
return Ok(None);
};
if task.status.is_terminal() {
return Ok(Some((
task.to_task_object(),
task.result.clone(),
task.error.clone(),
)));
}
task.completion_notify.clone()
};
notify.notified().await;
self.get_task_result(task_id).await
}
async fn list_tasks(&self, status_filter: Option<TaskStatus>) -> Result<Vec<TaskObject>> {
Ok(if let Ok(tasks) = self.tasks.read() {
tasks
.values()
.filter(|t| !t.is_expired())
.filter(|t| status_filter.is_none() || status_filter == Some(t.status))
.map(|t| t.to_task_object())
.collect()
} else {
vec![]
})
}
async fn require_input(
&self,
task_id: &str,
requests: InputRequests,
message: Option<&str>,
) -> Result<bool> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(false);
};
let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
return Ok(false);
};
if task.status.is_terminal() {
return Ok(false);
}
for key in std::mem::take(&mut task.input_requests).into_keys() {
if !requests.contains_key(&key) {
task.superseded_input_keys.insert(key);
}
}
for key in requests.keys() {
task.answered_input_keys.remove(key);
task.superseded_input_keys.remove(key);
}
task.input_requests = requests;
task.status = TaskStatus::InputRequired;
task.status_message = Some(
message
.map(str::to_string)
.unwrap_or_else(|| "Awaiting client input".to_string()),
);
task.last_updated_at_str = chrono_now_iso8601();
Ok(true)
}
async fn outstanding_input_requests(&self, task_id: &str) -> Result<Option<InputRequests>> {
Ok(if let Ok(tasks) = self.tasks.read() {
tasks
.get(task_id)
.filter(|t| !t.is_expired())
.map(|t| t.input_requests.clone())
} else {
None
})
}
async fn apply_input_responses(
&self,
task_id: &str,
responses: InputResponses,
) -> Result<Option<AppliedInputResponses>> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(None);
};
let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
return Ok(None);
};
if task.status.is_terminal() {
return Ok(None);
}
let mut applied = AppliedInputResponses::default();
for key in responses.into_keys() {
if task.input_requests.remove(&key).is_some() {
task.answered_input_keys.insert(key.clone());
applied.accepted.insert(key);
} else {
applied.ignored.insert(key);
}
}
applied.still_outstanding = task.input_requests.keys().cloned().collect();
if !applied.accepted.is_empty() {
task.last_updated_at_str = chrono_now_iso8601();
}
if applied.is_complete() && task.status == TaskStatus::InputRequired {
task.status = TaskStatus::Working;
task.status_message = Some("Task resumed".to_string());
}
Ok(Some(applied))
}
async fn set_ttl(&self, task_id: &str, ttl_ms: u64) -> Result<bool> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(false);
};
let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
return Ok(false);
};
task.ttl = ttl_ms;
task.last_updated_at_str = chrono_now_iso8601();
Ok(true)
}
async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(false);
};
let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
return Ok(false);
};
if task.status.is_terminal() {
return Ok(false);
}
task.status = TaskStatus::Completed;
task.status_message = Some("Task completed".to_string());
task.result = Some(result);
task.input_requests.clear();
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
task.completion_notify.notify_waiters();
Ok(true)
}
async fn fail_task(&self, task_id: &str, error: JsonRpcError) -> Result<bool> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(false);
};
let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
return Ok(false);
};
if task.status.is_terminal() {
return Ok(false);
}
task.status = TaskStatus::Failed;
task.status_message = Some(format!("Task failed: {}", error.message));
task.error = Some(error);
task.input_requests.clear();
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
task.completion_notify.notify_waiters();
Ok(true)
}
async fn cancel_task(&self, task_id: &str, reason: Option<&str>) -> Result<Option<TaskObject>> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(None);
};
let Some(task) = tasks.get_mut(task_id).filter(|t| !t.is_expired()) else {
return Ok(None);
};
task.cancellation_token.cancel();
if !task.status.is_terminal() {
task.input_requests.clear();
task.status = TaskStatus::Cancelled;
task.status_message = Some(
reason
.map(|r| format!("Cancelled: {}", r))
.unwrap_or_else(|| "Task cancelled".to_string()),
);
task.completed_at = Some(Instant::now());
task.last_updated_at_str = chrono_now_iso8601();
task.completion_notify.notify_waiters();
}
Ok(Some(task.to_task_object()))
}
}
pub fn tasks_extension() -> crate::ExtensionDeclaration {
crate::ExtensionDeclaration::empty(crate::protocol::TASKS_EXTENSION_ID)
.expect("the built-in Tasks extension declaration is valid")
}
impl crate::McpRouter {
pub fn with_tasks(self) -> Self {
self.with_protocol_extension(tasks_extension())
}
}
impl crate::McpClientBuilder {
pub fn with_tasks(self) -> Self {
self.with_protocol_extension(tasks_extension())
}
}
impl crate::RequestContext {
pub fn supports_tasks(&self) -> bool {
self.negotiated_extensions()
.is_some_and(|extensions| extensions.contains(crate::protocol::TASKS_EXTENSION_ID))
}
}
fn chrono_now_iso8601() -> String {
use std::time::SystemTime;
let now = SystemTime::now();
let duration = now
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default();
let secs = duration.as_secs();
let millis = duration.subsec_millis();
let days = secs / 86400;
let remaining = secs % 86400;
let hours = remaining / 3600;
let remaining = remaining % 3600;
let minutes = remaining / 60;
let seconds = remaining % 60;
let mut year = 1970i32;
let mut remaining_days = days as i32;
loop {
let days_in_year = if is_leap_year(year) { 366 } else { 365 };
if remaining_days < days_in_year {
break;
}
remaining_days -= days_in_year;
year += 1;
}
let days_in_months: [i32; 12] = if is_leap_year(year) {
[31, 29, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
} else {
[31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]
};
let mut month = 1;
for days_in_month in days_in_months.iter() {
if remaining_days < *days_in_month {
break;
}
remaining_days -= days_in_month;
month += 1;
}
let day = remaining_days + 1;
format!(
"{:04}-{:02}-{:02}T{:02}:{:02}:{:02}.{:03}Z",
year, month, day, hours, minutes, seconds, millis
)
}
fn is_leap_year(year: i32) -> bool {
(year % 4 == 0 && year % 100 != 0) || (year % 400 == 0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::{
ElicitAction, ElicitResult, InputRequest, InputResponse, ListRootsParams,
};
#[tokio::test]
async fn test_create_task() {
let store = MemoryTaskStore::new();
let (id, token) = store
.create_task("test-tool", serde_json::json!({"a": 1}), None, None)
.await
.unwrap();
assert!(!id.is_empty());
assert!(!token.is_cancelled());
let info = store
.get_task(&id)
.await
.unwrap()
.expect("task should exist");
assert_eq!(info.task_id, id);
assert_eq!(info.status, TaskStatus::Working);
}
#[tokio::test]
async fn test_task_lifecycle() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap()
);
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Completed);
}
#[tokio::test]
async fn test_task_cancellation() {
let store = MemoryTaskStore::new();
let (id, token) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(!token.is_cancelled());
let task_obj = store
.cancel_task(&id, Some("User requested"))
.await
.unwrap();
assert!(task_obj.is_some());
assert_eq!(task_obj.unwrap().status, TaskStatus::Cancelled);
assert!(token.is_cancelled());
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Cancelled);
}
#[tokio::test]
async fn test_task_failure() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(
store
.fail_task(&id, JsonRpcError::internal_error("Something went wrong"))
.await
.unwrap()
);
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Failed);
assert!(info.status_message.as_ref().unwrap().contains("failed"));
}
#[tokio::test]
async fn test_list_tasks() {
let store = MemoryTaskStore::new();
store
.create_task("tool1", serde_json::json!({}), None, None)
.await
.unwrap();
store
.create_task("tool2", serde_json::json!({}), None, None)
.await
.unwrap();
let (id3, _) = store
.create_task("tool3", serde_json::json!({}), None, None)
.await
.unwrap();
store
.complete_task(&id3, CallToolResult::text("Done"))
.await
.unwrap();
let all = store.list_tasks(None).await.unwrap();
assert_eq!(all.len(), 3);
let working = store.list_tasks(Some(TaskStatus::Working)).await.unwrap();
assert_eq!(working.len(), 2);
let completed = store.list_tasks(Some(TaskStatus::Completed)).await.unwrap();
assert_eq!(completed.len(), 1);
}
#[tokio::test]
async fn test_terminal_state_immutable() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap();
assert!(
!store
.fail_task(&id, JsonRpcError::internal_error("Error"))
.await
.unwrap()
);
let info = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(info.status, TaskStatus::Completed);
}
#[tokio::test]
async fn test_task_ids_unique() {
let store = MemoryTaskStore::new();
let (id1, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
let (id2, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
let (id3, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert_ne!(id1, id2);
assert_ne!(id2, id3);
assert_ne!(id1, id3);
}
#[tokio::test]
async fn test_get_task_result() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let result = CallToolResult::text("The result");
store.complete_task(&id, result).await.unwrap();
let (task_obj, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
assert_eq!(task_obj.status, TaskStatus::Completed);
assert!(result.is_some());
assert!(error.is_none());
}
#[tokio::test]
async fn test_wait_for_completion_returns_terminal_snapshot() {
let store = MemoryTaskStore::new();
let (id, _) = store
.create_task("test-tool", serde_json::json!({}), None, None)
.await
.unwrap();
let waiter_store = store.clone();
let waiter_id = id.clone();
let waiter =
tokio::spawn(async move { waiter_store.wait_for_completion(&waiter_id).await });
tokio::time::sleep(Duration::from_millis(10)).await;
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap();
let (task_obj, result, error) = waiter.await.unwrap().unwrap().unwrap();
assert_eq!(task_obj.status, TaskStatus::Completed);
assert!(result.is_some());
assert!(error.is_none());
}
#[tokio::test]
async fn dyn_task_store_object_safe() {
let store: Arc<dyn TaskStore> = Arc::new(MemoryTaskStore::new());
let (id, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert!(store.get_task(&id).await.unwrap().is_some());
}
#[test]
fn test_iso8601_timestamp() {
let ts = chrono_now_iso8601();
assert!(ts.ends_with('Z'));
assert!(ts.contains('T'));
assert_eq!(ts.len(), 24); }
#[test]
fn test_task_status_display() {
assert_eq!(TaskStatus::Working.to_string(), "working");
assert_eq!(TaskStatus::InputRequired.to_string(), "input_required");
assert_eq!(TaskStatus::Completed.to_string(), "completed");
assert_eq!(TaskStatus::Failed.to_string(), "failed");
assert_eq!(TaskStatus::Cancelled.to_string(), "cancelled");
}
#[test]
fn test_task_status_is_terminal() {
assert!(!TaskStatus::Working.is_terminal());
assert!(!TaskStatus::InputRequired.is_terminal());
assert!(TaskStatus::Completed.is_terminal());
assert!(TaskStatus::Failed.is_terminal());
assert!(TaskStatus::Cancelled.is_terminal());
}
fn requests(keys: &[&str]) -> InputRequests {
keys.iter()
.map(|k| {
(
k.to_string(),
InputRequest::ListRoots(ListRootsParams { meta: None }),
)
})
.collect()
}
fn accept(key: &str) -> (String, InputResponse) {
(
key.to_string(),
InputResponse::Elicit(ElicitResult {
action: ElicitAction::Accept,
content: None,
meta: None,
}),
)
}
async fn working_task(store: &MemoryTaskStore, ttl: Option<u64>) -> String {
store
.create_task("tool", serde_json::json!({}), ttl, None)
.await
.unwrap()
.0
}
#[tokio::test]
async fn task_ids_are_unguessable_not_sequential() {
let store = MemoryTaskStore::new();
let mut ids = BTreeSet::new();
for _ in 0..64 {
ids.insert(working_task(&store, None).await);
}
assert_eq!(ids.len(), 64, "task IDs collided");
for id in &ids {
assert_eq!(id.len(), 32, "expected 128 bits of hex: {id}");
assert!(id.chars().all(|c| c.is_ascii_hexdigit()), "{id}");
assert!(!id.starts_with("task-"), "sequential-looking ID: {id}");
}
let leading: BTreeSet<&str> = ids.iter().map(|id| &id[..2]).collect();
assert!(
leading.len() > 32,
"only {} distinct leading bytes across 64 IDs",
leading.len()
);
}
#[tokio::test]
async fn ttl_runs_from_creation_and_expired_tasks_read_as_absent() {
let store = MemoryTaskStore::new();
let id = working_task(&store, Some(0)).await;
tokio::time::sleep(Duration::from_millis(5)).await;
assert!(store.get_task(&id).await.unwrap().is_none());
assert!(store.get_task_result(&id).await.unwrap().is_none());
assert!(store.list_tasks(None).await.unwrap().is_empty());
assert!(
store
.outstanding_input_requests(&id)
.await
.unwrap()
.is_none()
);
assert!(store.cancel_task(&id, None).await.unwrap().is_none());
assert!(!store.set_ttl(&id, 60_000).await.unwrap());
assert!(
!store
.complete_task(&id, CallToolResult::text("late"))
.await
.unwrap()
);
}
#[tokio::test]
async fn ttl_is_mutable_over_the_task_lifetime() {
let store = MemoryTaskStore::new();
let id = working_task(&store, Some(60_000)).await;
assert!(store.set_ttl(&id, 120_000).await.unwrap());
let task = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(task.ttl, Some(120_000));
assert!(store.set_ttl(&id, 0).await.unwrap());
tokio::time::sleep(Duration::from_millis(5)).await;
assert!(store.get_task(&id).await.unwrap().is_none());
}
#[tokio::test]
async fn require_input_records_requests_and_exposes_them() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
assert!(
store
.require_input(&id, requests(&["approval", "region"]), Some("need input"))
.await
.unwrap()
);
let task = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(task.status, TaskStatus::InputRequired);
assert_eq!(task.status_message.as_deref(), Some("need input"));
let outstanding = store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap();
assert_eq!(
outstanding.keys().collect::<Vec<_>>(),
vec!["approval", "region"],
"every outstanding request must be exposed, not just the newest"
);
}
#[tokio::test]
async fn partial_input_responses_leave_the_rest_outstanding() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval", "region"]), None)
.await
.unwrap();
let applied = store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert_eq!(applied.accepted, ["approval".to_string()].into());
assert!(applied.ignored.is_empty());
assert_eq!(applied.still_outstanding, ["region".to_string()].into());
assert!(!applied.is_complete());
let task = store.get_task(&id).await.unwrap().unwrap();
assert_eq!(task.status, TaskStatus::InputRequired);
assert_eq!(
store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap()
.keys()
.collect::<Vec<_>>(),
vec!["region"]
);
let applied = store
.apply_input_responses(&id, [accept("region")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert!(applied.is_complete());
assert_eq!(
store.get_task(&id).await.unwrap().unwrap().status,
TaskStatus::Working
);
}
#[tokio::test]
async fn unknown_answered_and_superseded_response_keys_are_ignored() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval", "stale"]), None)
.await
.unwrap();
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
store
.require_input(&id, requests(&["region"]), None)
.await
.unwrap();
let applied = store
.apply_input_responses(
&id,
[accept("never-issued"), accept("approval"), accept("stale")]
.into_iter()
.collect(),
)
.await
.unwrap()
.unwrap();
assert!(
applied.accepted.is_empty(),
"none of these keys are outstanding"
);
assert_eq!(
applied.ignored,
[
"never-issued".to_string(),
"approval".to_string(),
"stale".to_string()
]
.into(),
"unknown, already-answered, and superseded keys are all ignored"
);
assert_eq!(applied.still_outstanding, ["region".to_string()].into());
assert_eq!(
store.get_task(&id).await.unwrap().unwrap().status,
TaskStatus::InputRequired,
"ignoring a stale update must not resume or fail the task"
);
}
#[tokio::test]
async fn reissued_key_becomes_a_fresh_question() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
let applied = store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.unwrap();
assert_eq!(applied.accepted, ["approval".to_string()].into());
assert!(applied.is_complete());
}
#[tokio::test]
async fn failed_tasks_preserve_the_structured_error() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
let mut error = JsonRpcError::invalid_params("bad region");
error.data = Some(serde_json::json!({"field": "region"}));
assert!(store.fail_task(&id, error).await.unwrap());
let (_, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
assert!(result.is_none());
let error = error.expect("structured error must survive the store");
assert_eq!(
error.code, -32602,
"the original code must not be flattened"
);
assert_eq!(error.message, "bad region");
assert_eq!(error.data.unwrap()["field"], "region");
}
#[tokio::test]
async fn tool_error_results_complete_the_task() {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
let mut result = CallToolResult::text("domain failure");
result.is_error = true;
assert!(store.complete_task(&id, result).await.unwrap());
let (task, result, error) = store.get_task_result(&id).await.unwrap().unwrap();
assert_eq!(
task.status,
TaskStatus::Completed,
"isError is a domain error, not an execution failure"
);
assert!(result.unwrap().is_error);
assert!(error.is_none(), "no JSON-RPC error accompanies isError");
}
#[tokio::test]
async fn tasks_record_their_creating_principal() {
let store = MemoryTaskStore::new();
let (owned, _) = store
.create_task("tool", serde_json::json!({}), None, Some("alice".into()))
.await
.unwrap();
let (unowned, _) = store
.create_task("tool", serde_json::json!({}), None, None)
.await
.unwrap();
assert_eq!(
store.task_owner(&owned).await.unwrap(),
Some(Some("alice".to_string()))
);
assert_eq!(store.task_owner(&unowned).await.unwrap(), Some(None));
assert_eq!(
store.task_owner("does-not-exist").await.unwrap(),
None,
"an unknown task has no owner record at all"
);
let wire = serde_json::to_value(store.get_task(&owned).await.unwrap().unwrap()).unwrap();
assert!(
wire.get("owner").is_none(),
"owner leaked to the wire: {wire}"
);
assert!(!wire.to_string().contains("alice"));
}
#[test]
fn owner_matching_is_equality_not_leniency() {
assert!(owner_matches(&None, None), "no auth configured");
assert!(owner_matches(&Some("alice".into()), Some("alice")));
assert!(
!owner_matches(&Some("alice".into()), Some("bob")),
"a different principal must not inherit the task"
);
assert!(
!owner_matches(&Some("alice".into()), None),
"dropping the token must not grant access"
);
assert!(
!owner_matches(&None, Some("alice")),
"an unowned task belongs to a different security context"
);
}
#[tokio::test]
async fn terminal_states_clear_outstanding_requests() {
for (label, terminate) in [("completed", true), ("cancelled", false)] {
let store = MemoryTaskStore::new();
let id = working_task(&store, None).await;
store
.require_input(&id, requests(&["approval"]), None)
.await
.unwrap();
if terminate {
store
.complete_task(&id, CallToolResult::text("done"))
.await
.unwrap();
} else {
store.cancel_task(&id, None).await.unwrap();
}
assert!(
store
.outstanding_input_requests(&id)
.await
.unwrap()
.unwrap()
.is_empty(),
"{label} task still advertises outstanding input requests"
);
assert!(
store
.apply_input_responses(&id, [accept("approval")].into_iter().collect())
.await
.unwrap()
.is_none(),
"{label} task accepted a late input response"
);
}
}
}