use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use crate::protocol::{CallToolResult, 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<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>) -> 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,
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 {
if let Some(completed_at) = self.completed_at {
completed_at.elapsed() > Duration::from_millis(self.ttl)
} else {
false
}
}
pub fn is_cancelled(&self) -> bool {
self.cancellation_token.is_cancelled()
}
}
#[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<String>);
#[async_trait]
pub trait TaskStore: Send + Sync + 'static {
async fn create_task(
&self,
tool_name: &str,
arguments: serde_json::Value,
ttl: Option<u64>,
) -> Result<(String, CancellationToken)>;
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, message: &str) -> Result<bool>;
async fn complete_task(&self, task_id: &str, result: CallToolResult) -> Result<bool>;
async fn fail_task(&self, task_id: &str, error: &str) -> 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>>>,
next_id: Arc<AtomicU64>,
}
impl Default for MemoryTaskStore {
fn default() -> Self {
Self::new()
}
}
impl MemoryTaskStore {
pub fn new() -> Self {
Self {
tasks: Arc::new(RwLock::new(HashMap::new())),
next_id: Arc::new(AtomicU64::new(1)),
}
}
fn generate_id(&self) -> String {
let id = self.next_id.fetch_add(1, Ordering::Relaxed);
format!("task-{}", id)
}
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>,
) -> Result<(String, CancellationToken)> {
let id = self.generate_id();
let task = Task::new(id.clone(), tool_name.to_string(), arguments, ttl);
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).map(|t| t.to_task_object())
} 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)
.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) 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| 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, message: &str) -> Result<bool> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(false);
};
let Some(task) = tasks.get_mut(task_id) else {
return Ok(false);
};
if task.status.is_terminal() {
return Ok(false);
}
task.status = TaskStatus::InputRequired;
task.status_message = Some(message.to_string());
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) 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.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: &str) -> Result<bool> {
let Ok(mut tasks) = self.tasks.write() else {
return Ok(false);
};
let Some(task) = tasks.get_mut(task_id) else {
return Ok(false);
};
if task.status.is_terminal() {
return Ok(false);
}
task.status = TaskStatus::Failed;
task.status_message = Some(format!("Task failed: {}", error));
task.error = Some(error.to_string());
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) else {
return Ok(None);
};
task.cancellation_token.cancel();
if !task.status.is_terminal() {
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()))
}
}
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::*;
#[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)
.await
.unwrap();
assert!(id.starts_with("task-"));
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)
.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)
.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)
.await
.unwrap();
assert!(store.fail_task(&id, "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)
.await
.unwrap();
store
.create_task("tool2", serde_json::json!({}), None)
.await
.unwrap();
let (id3, _) = store
.create_task("tool3", serde_json::json!({}), 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)
.await
.unwrap();
store
.complete_task(&id, CallToolResult::text("Done"))
.await
.unwrap();
assert!(!store.fail_task(&id, "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)
.await
.unwrap();
let (id2, _) = store
.create_task("tool", serde_json::json!({}), None)
.await
.unwrap();
let (id3, _) = store
.create_task("tool", serde_json::json!({}), 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)
.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)
.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)
.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());
}
}