use async_trait::async_trait;
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use std::collections::BTreeSet;
use std::time::Instant;
use crate::types::mrtr::{InputRequestKind, InputRequests, InputResponses};
use crate::types::tasks::{Task, TaskStatus};
use crate::types::CallToolResult;
#[derive(Debug)]
pub enum TaskStoreError {
NotFound {
task_id: String,
},
InvalidTransition {
task_id: String,
from: TaskStatus,
to: TaskStatus,
},
Expired {
task_id: String,
},
Internal {
message: String,
},
}
impl std::fmt::Display for TaskStoreError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NotFound { task_id } => write!(f, "task not found: {task_id}"),
Self::InvalidTransition { task_id, from, to } => {
write!(f, "invalid transition for task {task_id}: {from} -> {to}")
},
Self::Expired { task_id } => write!(f, "task expired: {task_id}"),
Self::Internal { message } => write!(f, "internal error: {message}"),
}
}
}
impl std::error::Error for TaskStoreError {}
impl From<TaskStoreError> for crate::error::Error {
fn from(err: TaskStoreError) -> Self {
match &err {
TaskStoreError::NotFound { .. } => Self::not_found(err.to_string()),
TaskStoreError::InvalidTransition { .. } => Self::validation(err.to_string()),
TaskStoreError::Expired { .. } => Self::not_found(err.to_string()),
TaskStoreError::Internal { .. } => Self::internal(err.to_string()),
}
}
}
#[derive(Debug, Clone)]
pub struct StoreConfig {
pub default_ttl_ms: Option<u64>,
pub max_ttl_ms: Option<u64>,
pub default_poll_interval_ms: u64,
pub max_tasks_per_owner: usize,
}
impl Default for StoreConfig {
fn default() -> Self {
Self {
default_ttl_ms: Some(3_600_000), max_ttl_ms: Some(86_400_000), default_poll_interval_ms: 5000, max_tasks_per_owner: 100,
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TaskInputDelivery {
pub accepted: BTreeSet<String>,
pub ignored: BTreeSet<String>,
pub complete: bool,
}
pub fn partition_input_delivery(
outstanding: &BTreeSet<String>,
already_answered: impl Fn(&str) -> bool,
delivered: impl IntoIterator<Item = String>,
) -> TaskInputDelivery {
let mut delivery = TaskInputDelivery::default();
for key in delivered {
if outstanding.contains(&key) && !already_answered(&key) {
delivery.accepted.insert(key);
} else {
delivery.ignored.insert(key);
}
}
delivery.complete = !outstanding.is_empty()
&& outstanding
.iter()
.all(|key| delivery.accepted.contains(key) || already_answered(key));
delivery
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct TaskInputSnapshot {
pub input_requests: InputRequests,
pub input_responses: InputResponses,
pub status: TaskStatus,
}
impl TaskInputSnapshot {
pub fn outstanding(&self) -> BTreeSet<&str> {
self.input_requests
.keys()
.filter(|key| !self.input_responses.contains_key(*key))
.map(String::as_str)
.collect()
}
pub fn kind_of(&self, key: &str) -> Option<InputRequestKind> {
self.input_requests
.get(key)
.map(crate::types::mrtr::InputRequest::kind)
}
pub fn is_complete(&self) -> bool {
!self.input_requests.is_empty()
&& self
.input_requests
.keys()
.all(|key| self.input_responses.contains_key(key))
}
}
#[async_trait]
pub trait TaskStore: Send + Sync {
async fn create(&self, owner_id: &str, ttl: Option<u64>) -> Result<Task, TaskStoreError>;
async fn get(&self, task_id: &str, owner_id: &str) -> Result<Task, TaskStoreError>;
async fn update_status(
&self,
task_id: &str,
owner_id: &str,
status: TaskStatus,
message: Option<String>,
) -> Result<Task, TaskStoreError>;
async fn list(
&self,
owner_id: &str,
cursor: Option<&str>,
) -> Result<(Vec<Task>, Option<String>), TaskStoreError>;
async fn cancel(&self, task_id: &str, owner_id: &str) -> Result<Task, TaskStoreError>;
async fn cleanup_expired(&self) -> Result<usize, TaskStoreError>;
fn config(&self) -> &StoreConfig;
async fn set_result(
&self,
_task_id: &str,
_owner_id: &str,
_result: crate::types::CallToolResult,
) -> Result<(), TaskStoreError> {
Err(TaskStoreError::Internal {
message: "store does not support terminal results".to_string(),
})
}
async fn get_result(
&self,
task_id: &str,
_owner_id: &str,
) -> Result<crate::types::CallToolResult, TaskStoreError> {
Err(TaskStoreError::NotFound {
task_id: task_id.to_string(),
})
}
fn supports_results(&self) -> bool {
false
}
async fn deliver_task_inputs(
&self,
_task_id: &str,
_owner_id: &str,
_responses: InputResponses,
) -> Result<TaskInputDelivery, TaskStoreError> {
Err(TaskStoreError::Internal {
message: "store does not support task input delivery".to_string(),
})
}
async fn task_input_snapshot(
&self,
task_id: &str,
_owner_id: &str,
) -> Result<TaskInputSnapshot, TaskStoreError> {
Err(TaskStoreError::NotFound {
task_id: task_id.to_string(),
})
}
async fn record_input_requests(
&self,
_task_id: &str,
_owner_id: &str,
_requests: InputRequests,
) -> Result<Task, TaskStoreError> {
Err(TaskStoreError::Internal {
message: "store does not support recording task input requests".to_string(),
})
}
async fn set_error(
&self,
_task_id: &str,
_owner_id: &str,
_error: serde_json::Value,
) -> Result<(), TaskStoreError> {
Err(TaskStoreError::Internal {
message: "store does not support task errors".to_string(),
})
}
async fn get_error(
&self,
task_id: &str,
_owner_id: &str,
) -> Result<serde_json::Value, TaskStoreError> {
Err(TaskStoreError::NotFound {
task_id: task_id.to_string(),
})
}
fn supports_inputs(&self) -> bool {
false
}
}
#[derive(Debug)]
struct TaskRecord {
task: Task,
owner_id: String,
expires_at: Option<Instant>,
result: Option<CallToolResult>,
input_requests: Option<InputRequests>,
input_responses: Option<InputResponses>,
error: Option<serde_json::Value>,
}
#[derive(Debug)]
pub struct InMemoryTaskStore {
records: DashMap<String, TaskRecord>,
config: StoreConfig,
}
impl InMemoryTaskStore {
pub fn new() -> Self {
Self {
records: DashMap::new(),
config: StoreConfig::default(),
}
}
pub fn with_config(config: StoreConfig) -> Self {
Self {
records: DashMap::new(),
config,
}
}
fn validate_access(
record: &TaskRecord,
task_id: &str,
owner_id: &str,
) -> Result<(), TaskStoreError> {
if record.owner_id != owner_id {
return Err(TaskStoreError::NotFound {
task_id: task_id.to_string(),
});
}
if let Some(expires_at) = record.expires_at {
if Instant::now() > expires_at {
return Err(TaskStoreError::Expired {
task_id: task_id.to_string(),
});
}
}
Ok(())
}
}
impl Default for InMemoryTaskStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl TaskStore for InMemoryTaskStore {
async fn create(&self, owner_id: &str, ttl: Option<u64>) -> Result<Task, TaskStoreError> {
let now = Instant::now();
let owner_count = self
.records
.iter()
.filter(|entry| {
let v = entry.value();
v.owner_id == owner_id && v.expires_at.is_none_or(|e| now <= e)
})
.count();
if owner_count >= self.config.max_tasks_per_owner {
return Err(TaskStoreError::Internal {
message: format!(
"owner {owner_id} has reached the maximum of {} tasks",
self.config.max_tasks_per_owner
),
});
}
let task_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now();
let now_str = now.to_rfc3339();
let effective_ttl = ttl.or(self.config.default_ttl_ms);
let effective_ttl = match (effective_ttl, self.config.max_ttl_ms) {
(Some(t), Some(max)) if t > max => Some(max),
(t, _) => t,
};
let expires_at =
effective_ttl.map(|ms| Instant::now() + std::time::Duration::from_millis(ms));
let task = Task::new(&task_id, TaskStatus::Working)
.with_timestamps(&now_str, &now_str)
.with_poll_interval(self.config.default_poll_interval_ms);
let task = if let Some(ttl_val) = effective_ttl {
task.with_ttl(ttl_val)
} else {
task
};
let record = TaskRecord {
task: task.clone(),
owner_id: owner_id.to_string(),
expires_at,
result: None,
input_requests: None,
input_responses: None,
error: None,
};
self.records.insert(task_id, record);
Ok(task)
}
async fn get(&self, task_id: &str, owner_id: &str) -> Result<Task, TaskStoreError> {
let entry = self
.records
.get(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
Self::validate_access(entry.value(), task_id, owner_id)?;
Ok(entry.value().task.clone())
}
async fn update_status(
&self,
task_id: &str,
owner_id: &str,
status: TaskStatus,
message: Option<String>,
) -> Result<Task, TaskStoreError> {
let mut entry = self
.records
.get_mut(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
let record = entry.value_mut();
Self::validate_access(record, task_id, owner_id)?;
if !record.task.status.can_transition_to(&status) {
return Err(TaskStoreError::InvalidTransition {
task_id: task_id.to_string(),
from: record.task.status,
to: status,
});
}
let now_str = chrono::Utc::now().to_rfc3339();
record.task.status = status;
record.task.last_updated_at = now_str;
record.task.status_message = message;
Ok(record.task.clone())
}
async fn list(
&self,
owner_id: &str,
cursor: Option<&str>,
) -> Result<(Vec<Task>, Option<String>), TaskStoreError> {
const PAGE_SIZE: usize = 20;
let now = Instant::now();
let mut tasks: Vec<Task> = self
.records
.iter()
.filter(|entry| {
let v = entry.value();
v.owner_id == owner_id && v.expires_at.is_none_or(|e| now <= e)
})
.map(|entry| entry.value().task.clone())
.collect();
tasks.sort_by(|a, b| b.created_at.cmp(&a.created_at));
if let Some(cursor_id) = cursor {
if let Some(pos) = tasks.iter().position(|t| t.task_id == cursor_id) {
tasks = tasks.into_iter().skip(pos + 1).collect();
}
}
if tasks.len() > PAGE_SIZE {
let next_cursor = tasks[PAGE_SIZE - 1].task_id.clone();
tasks.truncate(PAGE_SIZE);
Ok((tasks, Some(next_cursor)))
} else {
Ok((tasks, None))
}
}
async fn cancel(&self, task_id: &str, owner_id: &str) -> Result<Task, TaskStoreError> {
self.update_status(task_id, owner_id, TaskStatus::Cancelled, None)
.await
}
async fn cleanup_expired(&self) -> Result<usize, TaskStoreError> {
let now = Instant::now();
let before = self.records.len();
self.records
.retain(|_, record| record.expires_at.is_none_or(|e| now <= e));
Ok(before - self.records.len())
}
fn config(&self) -> &StoreConfig {
&self.config
}
async fn set_result(
&self,
task_id: &str,
owner_id: &str,
result: CallToolResult,
) -> Result<(), TaskStoreError> {
let mut entry = self
.records
.get_mut(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
let record = entry.value_mut();
Self::validate_access(record, task_id, owner_id)?;
record.result = Some(result);
Ok(())
}
async fn get_result(
&self,
task_id: &str,
owner_id: &str,
) -> Result<CallToolResult, TaskStoreError> {
let entry = self
.records
.get(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
Self::validate_access(entry.value(), task_id, owner_id)?;
entry
.value()
.result
.clone()
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})
}
fn supports_results(&self) -> bool {
true
}
async fn deliver_task_inputs(
&self,
task_id: &str,
owner_id: &str,
responses: InputResponses,
) -> Result<TaskInputDelivery, TaskStoreError> {
let mut entry = self
.records
.get_mut(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
let record = entry.value_mut();
Self::validate_access(record, task_id, owner_id)?;
if !record.task.status.can_transition_to(&TaskStatus::Working) {
return Err(TaskStoreError::InvalidTransition {
task_id: task_id.to_string(),
from: record.task.status,
to: TaskStatus::Working,
});
}
let outstanding: BTreeSet<String> = record
.input_requests
.as_ref()
.map(|requests| requests.keys().cloned().collect())
.unwrap_or_default();
let delivery = partition_input_delivery(
&outstanding,
|key| {
record
.input_responses
.as_ref()
.is_some_and(|answered| answered.contains_key(key))
},
responses.keys().cloned(),
);
for (key, response) in responses {
if delivery.accepted.contains(&key) {
record
.input_responses
.get_or_insert_with(InputResponses::new)
.insert(key, response);
}
}
if delivery.complete && !delivery.accepted.is_empty() {
record.task.status = TaskStatus::Working;
record.task.last_updated_at = chrono::Utc::now().to_rfc3339();
}
Ok(delivery)
}
async fn task_input_snapshot(
&self,
task_id: &str,
owner_id: &str,
) -> Result<TaskInputSnapshot, TaskStoreError> {
let entry = self
.records
.get(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
Self::validate_access(entry.value(), task_id, owner_id)?;
let record = entry.value();
let input_requests =
record
.input_requests
.clone()
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
Ok(TaskInputSnapshot {
input_requests,
input_responses: record.input_responses.clone().unwrap_or_default(),
status: record.task.status,
})
}
async fn record_input_requests(
&self,
task_id: &str,
owner_id: &str,
requests: InputRequests,
) -> Result<Task, TaskStoreError> {
let mut entry = self
.records
.get_mut(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
let record = entry.value_mut();
Self::validate_access(record, task_id, owner_id)?;
if record
.input_requests
.as_ref()
.is_some_and(|recorded| !recorded.is_empty())
{
return Err(TaskStoreError::Internal {
message: format!("task {task_id} already has recorded input requests"),
});
}
if !record
.task
.status
.can_transition_to(&TaskStatus::InputRequired)
{
return Err(TaskStoreError::InvalidTransition {
task_id: task_id.to_string(),
from: record.task.status,
to: TaskStatus::InputRequired,
});
}
record.input_requests = Some(requests);
record.task.status = TaskStatus::InputRequired;
record.task.last_updated_at = chrono::Utc::now().to_rfc3339();
Ok(record.task.clone())
}
async fn set_error(
&self,
task_id: &str,
owner_id: &str,
error: serde_json::Value,
) -> Result<(), TaskStoreError> {
let mut entry = self
.records
.get_mut(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
let record = entry.value_mut();
Self::validate_access(record, task_id, owner_id)?;
record.error = Some(error);
Ok(())
}
async fn get_error(
&self,
task_id: &str,
owner_id: &str,
) -> Result<serde_json::Value, TaskStoreError> {
let entry = self
.records
.get(task_id)
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})?;
Self::validate_access(entry.value(), task_id, owner_id)?;
entry
.value()
.error
.clone()
.ok_or_else(|| TaskStoreError::NotFound {
task_id: task_id.to_string(),
})
}
fn supports_inputs(&self) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_creates_empty_store() {
let store = InMemoryTaskStore::new();
assert!(store.records.is_empty());
}
#[test]
fn default_creates_empty_store() {
let store = InMemoryTaskStore::default();
assert!(store.records.is_empty());
}
#[test]
fn with_config_applies_custom_config() {
let config = StoreConfig {
default_ttl_ms: Some(1_000),
max_ttl_ms: Some(2_000),
default_poll_interval_ms: 500,
max_tasks_per_owner: 10,
};
let store = InMemoryTaskStore::with_config(config);
assert_eq!(store.config().default_ttl_ms, Some(1_000));
assert_eq!(store.config().max_ttl_ms, Some(2_000));
assert_eq!(store.config().default_poll_interval_ms, 500);
assert_eq!(store.config().max_tasks_per_owner, 10);
}
#[test]
fn store_config_default_values() {
let config = StoreConfig::default();
assert_eq!(config.default_ttl_ms, Some(3_600_000));
assert_eq!(config.max_ttl_ms, Some(86_400_000));
assert_eq!(config.default_poll_interval_ms, 5000);
assert_eq!(config.max_tasks_per_owner, 100);
}
#[tokio::test]
async fn create_returns_working_task() {
let store = InMemoryTaskStore::new();
let task = store.create("owner-1", None).await.unwrap();
assert_eq!(task.status, TaskStatus::Working);
assert!(!task.task_id.is_empty());
assert!(!task.created_at.is_empty());
assert!(!task.last_updated_at.is_empty());
}
#[tokio::test]
async fn create_with_default_ttl() {
let store = InMemoryTaskStore::new();
let task = store.create("owner-1", None).await.unwrap();
assert_eq!(task.ttl, Some(3_600_000));
}
#[tokio::test]
async fn create_with_explicit_ttl() {
let store = InMemoryTaskStore::new();
let task = store.create("owner-1", Some(60_000)).await.unwrap();
assert_eq!(task.ttl, Some(60_000));
}
#[tokio::test]
async fn create_clamps_ttl_to_max() {
let store = InMemoryTaskStore::with_config(StoreConfig {
max_ttl_ms: Some(10_000),
..StoreConfig::default()
});
let task = store.create("owner-1", Some(999_999)).await.unwrap();
assert_eq!(task.ttl, Some(10_000));
}
#[tokio::test]
async fn create_sets_poll_interval() {
let store = InMemoryTaskStore::with_config(StoreConfig {
default_poll_interval_ms: 3000,
..StoreConfig::default()
});
let task = store.create("owner-1", None).await.unwrap();
assert_eq!(task.poll_interval, Some(3000));
}
#[tokio::test]
async fn get_returns_created_task() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let fetched = store.get(&created.task_id, "owner-1").await.unwrap();
assert_eq!(fetched.task_id, created.task_id);
assert_eq!(fetched.status, TaskStatus::Working);
}
#[tokio::test]
async fn get_owner_mismatch_returns_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let result = store.get(&created.task_id, "owner-2").await;
assert!(
matches!(&result, Err(TaskStoreError::NotFound { task_id }) if task_id == &created.task_id),
"expected NotFound, got: {result:?}"
);
}
#[tokio::test]
async fn get_nonexistent_returns_not_found() {
let store = InMemoryTaskStore::new();
let result = store.get("nonexistent", "owner-1").await;
assert!(matches!(result, Err(TaskStoreError::NotFound { .. })));
}
#[tokio::test]
async fn list_returns_owner_tasks_only() {
let store = InMemoryTaskStore::new();
store.create("owner-1", None).await.unwrap();
store.create("owner-1", None).await.unwrap();
store.create("owner-2", None).await.unwrap();
let (tasks, _) = store.list("owner-1", None).await.unwrap();
assert_eq!(tasks.len(), 2);
}
#[tokio::test]
async fn list_empty_for_unknown_owner() {
let store = InMemoryTaskStore::new();
store.create("owner-1", None).await.unwrap();
let (tasks, _) = store.list("owner-unknown", None).await.unwrap();
assert!(tasks.is_empty());
}
#[tokio::test]
async fn list_sorted_newest_first() {
let store = InMemoryTaskStore::new();
let first = store.create("owner-1", None).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(2)).await;
let second = store.create("owner-1", None).await.unwrap();
let (tasks, _) = store.list("owner-1", None).await.unwrap();
assert_eq!(tasks.len(), 2);
assert_eq!(tasks[0].task_id, second.task_id);
assert_eq!(tasks[1].task_id, first.task_id);
}
#[tokio::test]
async fn cancel_transitions_to_cancelled() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let cancelled = store.cancel(&created.task_id, "owner-1").await.unwrap();
assert_eq!(cancelled.status, TaskStatus::Cancelled);
}
#[tokio::test]
async fn cancel_completed_task_returns_invalid_transition() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.update_status(
&created.task_id,
"owner-1",
TaskStatus::Completed,
Some("Done".to_string()),
)
.await
.unwrap();
let result = store.cancel(&created.task_id, "owner-1").await;
assert!(
matches!(result, Err(TaskStoreError::InvalidTransition { .. })),
"expected InvalidTransition, got: {result:?}"
);
}
#[tokio::test]
async fn update_status_working_to_completed() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let updated = store
.update_status(
&created.task_id,
"owner-1",
TaskStatus::Completed,
Some("Done".to_string()),
)
.await
.unwrap();
assert_eq!(updated.status, TaskStatus::Completed);
assert_eq!(updated.status_message.as_deref(), Some("Done"));
}
#[tokio::test]
async fn update_status_from_terminal_returns_invalid_transition() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.update_status(&created.task_id, "owner-1", TaskStatus::Completed, None)
.await
.unwrap();
let result = store
.update_status(&created.task_id, "owner-1", TaskStatus::Working, None)
.await;
assert!(
matches!(result, Err(TaskStoreError::InvalidTransition { .. })),
"expected InvalidTransition, got: {result:?}"
);
}
#[tokio::test]
async fn update_status_self_transition_rejected() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let result = store
.update_status(&created.task_id, "owner-1", TaskStatus::Working, None)
.await;
assert!(
matches!(result, Err(TaskStoreError::InvalidTransition { .. })),
"expected InvalidTransition, got: {result:?}"
);
}
#[tokio::test]
async fn task_created_with_explicit_ttl_has_correct_field() {
let store = InMemoryTaskStore::new();
let task = store.create("owner-1", Some(60_000)).await.unwrap();
assert_eq!(task.ttl, Some(60_000));
}
#[tokio::test]
async fn task_created_with_none_ttl_gets_default() {
let config = StoreConfig {
default_ttl_ms: Some(120_000),
..StoreConfig::default()
};
let store = InMemoryTaskStore::with_config(config);
let task = store.create("owner-1", None).await.unwrap();
assert_eq!(task.ttl, Some(120_000));
}
#[tokio::test]
async fn cleanup_expired_removes_expired_tasks() {
let store = InMemoryTaskStore::with_config(StoreConfig {
default_ttl_ms: Some(1), ..StoreConfig::default()
});
store.create("owner-1", Some(1)).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let removed = store.cleanup_expired().await.unwrap();
assert_eq!(removed, 1);
assert!(store.records.is_empty());
}
#[tokio::test]
async fn cleanup_expired_keeps_non_expired() {
let store = InMemoryTaskStore::new();
store.create("owner-1", Some(3_600_000)).await.unwrap();
let removed = store.cleanup_expired().await.unwrap();
assert_eq!(removed, 0);
assert_eq!(store.records.len(), 1);
}
#[tokio::test]
async fn get_expired_task_returns_expired_error() {
let store = InMemoryTaskStore::with_config(StoreConfig {
default_ttl_ms: Some(1), ..StoreConfig::default()
});
let created = store.create("owner-1", Some(1)).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let result = store.get(&created.task_id, "owner-1").await;
assert!(
matches!(result, Err(TaskStoreError::Expired { .. })),
"expected Expired, got: {result:?}"
);
}
#[test]
fn task_store_error_display_not_found() {
let err = TaskStoreError::NotFound {
task_id: "t-123".to_string(),
};
assert_eq!(err.to_string(), "task not found: t-123");
}
#[test]
fn task_store_error_display_invalid_transition() {
let err = TaskStoreError::InvalidTransition {
task_id: "t-123".to_string(),
from: TaskStatus::Completed,
to: TaskStatus::Working,
};
let msg = err.to_string();
assert!(msg.contains("invalid transition"));
assert!(msg.contains("t-123"));
}
#[test]
fn task_store_error_display_expired() {
let err = TaskStoreError::Expired {
task_id: "t-123".to_string(),
};
assert_eq!(err.to_string(), "task expired: t-123");
}
#[test]
fn task_store_error_display_internal() {
let err = TaskStoreError::Internal {
message: "something broke".to_string(),
};
assert_eq!(err.to_string(), "internal error: something broke");
}
#[test]
fn task_store_error_converts_to_sdk_error() {
let err = TaskStoreError::NotFound {
task_id: "t-123".to_string(),
};
let sdk_err: crate::error::Error = err.into();
let msg = sdk_err.to_string();
assert!(msg.contains("task not found: t-123"));
}
#[tokio::test]
async fn max_tasks_per_owner_enforced() {
let store = InMemoryTaskStore::with_config(StoreConfig {
max_tasks_per_owner: 2,
..StoreConfig::default()
});
store.create("owner-1", None).await.unwrap();
store.create("owner-1", None).await.unwrap();
let result = store.create("owner-1", None).await;
assert!(
matches!(result, Err(TaskStoreError::Internal { .. })),
"expected Internal error for max tasks, got: {result:?}"
);
}
#[tokio::test]
async fn max_tasks_scoped_to_owner() {
let store = InMemoryTaskStore::with_config(StoreConfig {
max_tasks_per_owner: 2,
..StoreConfig::default()
});
store.create("owner-a", None).await.unwrap();
store.create("owner-a", None).await.unwrap();
let result = store.create("owner-b", None).await;
assert!(result.is_ok());
}
use crate::types::{CallToolResult, Content};
use serde_json::json;
fn sample_result(text: &str) -> CallToolResult {
CallToolResult::new(vec![Content::text(text)])
}
#[tokio::test]
async fn set_then_get_result_round_trips() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.set_result(&created.task_id, "owner-1", sample_result("hello"))
.await
.unwrap();
let fetched = store.get_result(&created.task_id, "owner-1").await.unwrap();
assert_eq!(fetched.content.len(), 1);
}
#[tokio::test]
async fn get_result_owner_mismatch_returns_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.set_result(&created.task_id, "owner-1", sample_result("secret"))
.await
.unwrap();
let result = store.get_result(&created.task_id, "owner-2").await;
assert!(
matches!(result, Err(TaskStoreError::NotFound { .. })),
"cross-owner read must be NotFound, got: {result:?}"
);
}
#[tokio::test]
async fn set_result_owner_mismatch_returns_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let result = store
.set_result(&created.task_id, "owner-2", sample_result("x"))
.await;
assert!(
matches!(result, Err(TaskStoreError::NotFound { .. })),
"cross-owner set must be NotFound, got: {result:?}"
);
}
#[tokio::test]
async fn get_result_existing_task_no_result_returns_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let result = store.get_result(&created.task_id, "owner-1").await;
assert!(
matches!(result, Err(TaskStoreError::NotFound { .. })),
"pending task (no result) must be NotFound, got: {result:?}"
);
}
#[tokio::test]
async fn in_memory_store_supports_results() {
let store = InMemoryTaskStore::new();
assert!(store.supports_results());
}
#[tokio::test]
async fn cleanup_expired_drops_result() {
let ttl_ms: u64 = 500;
let store = InMemoryTaskStore::with_config(StoreConfig {
default_ttl_ms: Some(ttl_ms),
..StoreConfig::default()
});
let created = store.create("owner-1", Some(ttl_ms)).await.unwrap();
store
.set_result(&created.task_id, "owner-1", sample_result("ephemeral"))
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(ttl_ms + 50)).await;
let removed = store.cleanup_expired().await.unwrap();
assert_eq!(removed, 1);
let result = store.get_result(&created.task_id, "owner-1").await;
assert!(matches!(result, Err(TaskStoreError::NotFound { .. })));
}
struct DefaultOnlyStore {
config: StoreConfig,
}
#[async_trait]
impl TaskStore for DefaultOnlyStore {
async fn create(&self, _owner_id: &str, _ttl: Option<u64>) -> Result<Task, TaskStoreError> {
Ok(Task::new("default-only", TaskStatus::Working))
}
async fn get(&self, task_id: &str, _owner_id: &str) -> Result<Task, TaskStoreError> {
Ok(Task::new(task_id, TaskStatus::Working))
}
async fn update_status(
&self,
task_id: &str,
_owner_id: &str,
status: TaskStatus,
_message: Option<String>,
) -> Result<Task, TaskStoreError> {
Ok(Task::new(task_id, status))
}
async fn list(
&self,
_owner_id: &str,
_cursor: Option<&str>,
) -> Result<(Vec<Task>, Option<String>), TaskStoreError> {
Ok((Vec::new(), None))
}
async fn cancel(&self, task_id: &str, _owner_id: &str) -> Result<Task, TaskStoreError> {
Ok(Task::new(task_id, TaskStatus::Cancelled))
}
async fn cleanup_expired(&self) -> Result<usize, TaskStoreError> {
Ok(0)
}
fn config(&self) -> &StoreConfig {
&self.config
}
}
#[tokio::test]
async fn default_impl_store_reports_unsupported() {
let store = DefaultOnlyStore {
config: StoreConfig::default(),
};
assert!(!store.supports_results());
let set = store.set_result("t", "owner-1", sample_result("x")).await;
assert!(
matches!(set, Err(TaskStoreError::Internal { .. })),
"default set_result must be an explicit unsupported error, got: {set:?}"
);
let get = store.get_result("t", "owner-1").await;
assert!(
matches!(get, Err(TaskStoreError::NotFound { .. })),
"default get_result must be NotFound, got: {get:?}"
);
}
#[tokio::test]
async fn default_impl_store_reports_inputs_unsupported() {
let store = DefaultOnlyStore {
config: StoreConfig::default(),
};
assert!(
!store.supports_inputs(),
"supports_inputs must default to false"
);
let delivered = store
.deliver_task_inputs("t", "owner-1", InputResponses::new())
.await;
assert!(
matches!(delivered, Err(TaskStoreError::Internal { .. })),
"default deliver_task_inputs must be an explicit unsupported error, got: {delivered:?}"
);
let recorded = store
.record_input_requests("t", "owner-1", InputRequests::new())
.await;
assert!(
matches!(recorded, Err(TaskStoreError::Internal { .. })),
"default record_input_requests must be an explicit unsupported error, got: {recorded:?}"
);
let set = store.set_error("t", "owner-1", json!({ "code": -1 })).await;
assert!(
matches!(set, Err(TaskStoreError::Internal { .. })),
"default set_error must be an explicit unsupported error, got: {set:?}"
);
let snapshot = store.task_input_snapshot("t", "owner-1").await;
assert!(
matches!(snapshot, Err(TaskStoreError::NotFound { .. })),
"default task_input_snapshot must be NotFound, got: {snapshot:?}"
);
let get = store.get_error("t", "owner-1").await;
assert!(
matches!(get, Err(TaskStoreError::NotFound { .. })),
"default get_error must be NotFound, got: {get:?}"
);
}
use crate::types::elicitation::{ElicitAction, ElicitRequestParams, ElicitResult};
use crate::types::mrtr::{InputRequest, InputResponse};
use crate::types::roots::ListRootsResult;
fn elicit_request(message: &str) -> InputRequest {
InputRequest::Elicitation(Box::new(ElicitRequestParams::Form {
message: message.to_string(),
requested_schema: json!({ "type": "object" }),
}))
}
fn elicit_response() -> InputResponse {
InputResponse::Elicitation(Box::new(ElicitResult {
action: ElicitAction::Accept,
content: None,
}))
}
fn roots_response() -> InputResponse {
InputResponse::Roots(Box::new(ListRootsResult { roots: Vec::new() }))
}
fn requests_of(keys: &[&str]) -> InputRequests {
let mut requests = InputRequests::new();
for key in keys {
requests.insert((*key).to_string(), elicit_request(key));
}
requests
}
fn responses_of(keys: &[&str]) -> InputResponses {
let mut responses = InputResponses::new();
for key in keys {
responses.insert((*key).to_string(), elicit_response());
}
responses
}
async fn paused_task(store: &InMemoryTaskStore, owner: &str, keys: &[&str]) -> Task {
let task = store.create(owner, None).await.unwrap();
store
.record_input_requests(&task.task_id, owner, requests_of(keys))
.await
.unwrap()
}
#[tokio::test]
async fn in_memory_store_supports_inputs() {
let store = InMemoryTaskStore::new();
assert!(store.supports_inputs());
}
#[tokio::test]
async fn deliver_inputs_completes_the_outstanding_set_and_transitions() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city"]).await;
let delivery = store
.deliver_task_inputs(&task.task_id, "owner-1", responses_of(&["city"]))
.await
.unwrap();
assert!(delivery.accepted.contains("city"), "got: {delivery:?}");
assert!(delivery.ignored.is_empty(), "got: {delivery:?}");
assert!(delivery.complete, "got: {delivery:?}");
let resumed = store.get(&task.task_id, "owner-1").await.unwrap();
assert_eq!(resumed.status, TaskStatus::Working);
assert_ne!(
resumed.last_updated_at, task.last_updated_at,
"a completing delivery must bump last_updated_at"
);
}
#[tokio::test]
async fn deliver_inputs_partial_set_persists_and_stays_input_required() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city", "units"]).await;
let delivery = store
.deliver_task_inputs(&task.task_id, "owner-1", responses_of(&["city"]))
.await
.unwrap();
assert!(delivery.accepted.contains("city"), "got: {delivery:?}");
assert!(
!delivery.complete,
"one of two outstanding keys is not a complete set: {delivery:?}"
);
let still = store.get(&task.task_id, "owner-1").await.unwrap();
assert_eq!(still.status, TaskStatus::InputRequired);
let snapshot = store
.task_input_snapshot(&task.task_id, "owner-1")
.await
.unwrap();
assert!(snapshot.input_responses.contains_key("city"));
assert_eq!(snapshot.outstanding(), BTreeSet::from(["units"]));
}
#[tokio::test]
async fn deliver_inputs_ignores_keys_that_are_not_outstanding() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city"]).await;
let delivery = store
.deliver_task_inputs(&task.task_id, "owner-1", responses_of(&["never-issued"]))
.await
.unwrap();
assert!(
delivery.accepted.is_empty(),
"a key the server never issued must not be accepted: {delivery:?}"
);
assert!(
delivery.ignored.contains("never-issued"),
"got: {delivery:?}"
);
assert!(!delivery.complete, "got: {delivery:?}");
let still = store.get(&task.task_id, "owner-1").await.unwrap();
assert_eq!(still.status, TaskStatus::InputRequired);
}
#[tokio::test]
async fn deliver_inputs_ignores_a_key_already_answered() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city", "units"]).await;
store
.deliver_task_inputs(&task.task_id, "owner-1", responses_of(&["city"]))
.await
.unwrap();
let replay = store
.deliver_task_inputs(&task.task_id, "owner-1", responses_of(&["city"]))
.await
.unwrap();
assert!(
replay.accepted.is_empty(),
"an already-answered key must not be re-accepted: {replay:?}"
);
assert!(replay.ignored.contains("city"), "got: {replay:?}");
assert!(!replay.complete, "got: {replay:?}");
}
#[tokio::test]
async fn deliver_inputs_on_a_completed_task_is_refused() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.update_status(&created.task_id, "owner-1", TaskStatus::Completed, None)
.await
.unwrap();
let result = store
.deliver_task_inputs(&created.task_id, "owner-1", responses_of(&["city"]))
.await;
assert!(
matches!(result, Err(TaskStoreError::InvalidTransition { .. })),
"a terminal task cannot be fed, got: {result:?}"
);
}
#[tokio::test]
async fn deliver_inputs_for_another_owner_is_not_found() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city"]).await;
let result = store
.deliver_task_inputs(&task.task_id, "owner-2", responses_of(&["city"]))
.await;
let Err(err) = result else {
panic!("cross-owner delivery must fail");
};
assert!(
matches!(err, TaskStoreError::NotFound { .. }),
"cross-owner delivery must be NotFound, got: {err:?}"
);
let rendered = err.to_string();
assert!(
!rendered.contains("owner"),
"refusal leaked the word `owner`: {rendered}"
);
assert!(
!rendered.contains("owner-1"),
"refusal leaked the other owner's id: {rendered}"
);
}
#[tokio::test]
async fn snapshot_returns_the_server_recorded_kinds_and_delivered_responses() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let mut requests = InputRequests::new();
requests.insert("city".to_string(), elicit_request("Which city?"));
requests.insert("roots".to_string(), InputRequest::ListRoots);
store
.record_input_requests(&created.task_id, "owner-1", requests)
.await
.unwrap();
let mut delivered = InputResponses::new();
delivered.insert("roots".to_string(), roots_response());
store
.deliver_task_inputs(&created.task_id, "owner-1", delivered)
.await
.unwrap();
let snapshot = store
.task_input_snapshot(&created.task_id, "owner-1")
.await
.unwrap();
assert_eq!(
snapshot.kind_of("city"),
Some(InputRequestKind::Elicitation)
);
assert_eq!(snapshot.kind_of("roots"), Some(InputRequestKind::Roots));
assert_eq!(
snapshot.kind_of("never-issued"),
None,
"a key the server never issued has no kind to decode against"
);
assert!(snapshot.input_responses.contains_key("roots"));
assert_eq!(snapshot.outstanding(), BTreeSet::from(["city"]));
assert!(!snapshot.is_complete());
assert_eq!(snapshot.status, TaskStatus::InputRequired);
}
#[tokio::test]
async fn snapshot_for_another_owner_is_not_found() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city"]).await;
let result = store.task_input_snapshot(&task.task_id, "owner-2").await;
let Err(err) = result else {
panic!("cross-owner snapshot must fail");
};
assert!(
matches!(err, TaskStoreError::NotFound { .. }),
"cross-owner snapshot must be NotFound, got: {err:?}"
);
let rendered = err.to_string();
assert!(
!rendered.contains("owner"),
"refusal leaked the word `owner`: {rendered}"
);
assert!(
!rendered.contains("owner-1"),
"refusal leaked the other owner's id: {rendered}"
);
}
#[tokio::test]
async fn snapshot_without_recorded_requests_is_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let result = store.task_input_snapshot(&created.task_id, "owner-1").await;
assert!(
matches!(result, Err(TaskStoreError::NotFound { .. })),
"a task with no recorded requests has no snapshot, got: {result:?}"
);
}
#[tokio::test]
async fn record_input_requests_transitions_to_input_required() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
assert_eq!(created.status, TaskStatus::Working);
let paused = store
.record_input_requests(&created.task_id, "owner-1", requests_of(&["city"]))
.await
.unwrap();
assert_eq!(paused.status, TaskStatus::InputRequired);
let fetched = store.get(&created.task_id, "owner-1").await.unwrap();
assert_eq!(fetched.status, TaskStatus::InputRequired);
}
#[tokio::test]
async fn record_input_requests_twice_is_refused_and_does_not_erase_answers() {
let store = InMemoryTaskStore::new();
let task = paused_task(&store, "owner-1", &["city", "units"]).await;
store
.deliver_task_inputs(&task.task_id, "owner-1", responses_of(&["city"]))
.await
.unwrap();
let second = store
.record_input_requests(&task.task_id, "owner-1", requests_of(&["something-else"]))
.await;
assert!(
second.is_err(),
"a second record_input_requests must be refused, got: {second:?}"
);
let snapshot = store
.task_input_snapshot(&task.task_id, "owner-1")
.await
.unwrap();
assert!(
snapshot.input_responses.contains_key("city"),
"the refusal erased a delivered answer: {snapshot:?}"
);
assert!(snapshot.input_requests.contains_key("city"));
assert!(snapshot.input_requests.contains_key("units"));
assert!(
!snapshot.input_requests.contains_key("something-else"),
"the refused write must not have landed: {snapshot:?}"
);
}
#[tokio::test]
async fn record_input_requests_on_a_terminal_task_is_refused() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.update_status(&created.task_id, "owner-1", TaskStatus::Cancelled, None)
.await
.unwrap();
let result = store
.record_input_requests(&created.task_id, "owner-1", requests_of(&["city"]))
.await;
assert!(
matches!(result, Err(TaskStoreError::InvalidTransition { .. })),
"a cancelled task cannot be paused for input, got: {result:?}"
);
}
#[tokio::test]
async fn set_error_then_get_error_round_trips_the_jsonrpc_error_value() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let error = json!({
"code": -32603,
"message": "upstream timed out",
"data": { "attempts": 3 }
});
store
.set_error(&created.task_id, "owner-1", error.clone())
.await
.unwrap();
let fetched = store.get_error(&created.task_id, "owner-1").await.unwrap();
assert_eq!(
fetched, error,
"the JSON-RPC error object must cross the Value seam unchanged"
);
}
#[tokio::test]
async fn get_error_for_another_owner_is_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
store
.set_error(
&created.task_id,
"owner-1",
json!({ "code": -32603, "message": "private" }),
)
.await
.unwrap();
let result = store.get_error(&created.task_id, "owner-2").await;
let Err(err) = result else {
panic!("cross-owner error read must fail");
};
assert!(
matches!(err, TaskStoreError::NotFound { .. }),
"cross-owner error read must be NotFound, got: {err:?}"
);
let rendered = err.to_string();
assert!(
!rendered.contains("owner"),
"refusal leaked the word `owner`: {rendered}"
);
assert!(
!rendered.contains("owner-1"),
"refusal leaked the other owner's id: {rendered}"
);
}
#[tokio::test]
async fn get_error_on_a_task_with_no_error_is_not_found() {
let store = InMemoryTaskStore::new();
let created = store.create("owner-1", None).await.unwrap();
let result = store.get_error(&created.task_id, "owner-1").await;
assert!(
matches!(result, Err(TaskStoreError::NotFound { .. })),
"a task that did not fail has no error, got: {result:?}"
);
}
#[tokio::test]
async fn cleanup_expired_drops_recorded_input_state() {
let ttl_ms: u64 = 500;
let store = InMemoryTaskStore::with_config(StoreConfig {
default_ttl_ms: Some(ttl_ms),
..StoreConfig::default()
});
let created = store.create("owner-1", Some(ttl_ms)).await.unwrap();
store
.record_input_requests(&created.task_id, "owner-1", requests_of(&["city"]))
.await
.unwrap();
store
.set_error(&created.task_id, "owner-1", json!({ "code": -32603 }))
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(ttl_ms + 50)).await;
assert_eq!(store.cleanup_expired().await.unwrap(), 1);
assert!(store.records.is_empty());
}
}