use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use async_trait::async_trait;
use tokio::sync::RwLock;
use crate::protocol::{A2ATask, A2ATaskResult, TaskFilter};
#[derive(Debug, Clone)]
pub struct StoredTask {
pub task: A2ATask,
pub result: Option<A2ATaskResult>,
pub error: Option<String>,
pub trace_id: Option<String>,
pub created_at: Instant,
pub updated_at: Instant,
}
impl StoredTask {
pub fn new(task: A2ATask) -> Self {
let now = Instant::now();
Self {
task,
result: None,
error: None,
trace_id: None,
created_at: now,
updated_at: now,
}
}
pub fn with_trace_id(mut self, trace_id: impl Into<String>) -> Self {
self.trace_id = Some(trace_id.into());
self
}
pub fn touch(&mut self) {
self.updated_at = Instant::now();
}
pub fn age(&self) -> Duration {
self.updated_at.elapsed()
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum StoreError {
#[error("task store unavailable: {0}")]
Unavailable(String),
#[error("task store capacity exceeded: {0}")]
CapacityExceeded(String),
}
#[async_trait]
pub trait TaskStore: Send + Sync {
async fn upsert(&self, stored: StoredTask) -> Result<(), StoreError>;
async fn get(&self, task_id: &str) -> Result<Option<StoredTask>, StoreError>;
async fn list(&self, filter: &TaskFilter) -> Result<Vec<StoredTask>, StoreError>;
async fn delete(&self, task_id: &str) -> Result<bool, StoreError>;
async fn compare_and_update(
&self,
task_id: &str,
update: StoredTask,
) -> Result<bool, StoreError> {
match self.get(task_id).await? {
Some(current) if current.task.status.can_transition_to(&update.task.status) => {
self.upsert(update).await?;
Ok(true)
}
_ => Ok(false),
}
}
}
pub const DEFAULT_MAX_TASKS: usize = 10_000;
#[derive(Debug, Clone)]
pub struct InMemoryTaskStore {
inner: Arc<RwLock<HashMap<String, StoredTask>>>,
max_tasks: usize,
}
impl InMemoryTaskStore {
pub fn new() -> Self {
Self::with_max_tasks(DEFAULT_MAX_TASKS)
}
pub fn with_max_tasks(max_tasks: usize) -> Self {
Self {
inner: Arc::new(RwLock::new(HashMap::new())),
max_tasks,
}
}
}
impl Default for InMemoryTaskStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl TaskStore for InMemoryTaskStore {
async fn upsert(&self, stored: StoredTask) -> Result<(), StoreError> {
let mut guard = self.inner.write().await;
let inserting_new = !guard.contains_key(&stored.task.id);
if inserting_new && self.max_tasks > 0 && guard.len() >= self.max_tasks {
let oldest_key = guard
.iter()
.min_by_key(|(_, t)| t.updated_at)
.map(|(k, _)| k.clone());
if let Some(key) = oldest_key {
guard.remove(&key);
}
}
guard.insert(stored.task.id.clone(), stored);
Ok(())
}
async fn get(&self, task_id: &str) -> Result<Option<StoredTask>, StoreError> {
Ok(self.inner.read().await.get(task_id).cloned())
}
async fn list(&self, filter: &TaskFilter) -> Result<Vec<StoredTask>, StoreError> {
let guard = self.inner.read().await;
let mut out: Vec<StoredTask> = guard
.values()
.filter(|t| filter.matches(&t.task))
.cloned()
.collect();
out.sort_by_key(|t| t.created_at);
Ok(out)
}
async fn delete(&self, task_id: &str) -> Result<bool, StoreError> {
Ok(self.inner.write().await.remove(task_id).is_some())
}
async fn compare_and_update(
&self,
task_id: &str,
update: StoredTask,
) -> Result<bool, StoreError> {
let mut guard = self.inner.write().await;
match guard.get(task_id) {
Some(current) if current.task.status.can_transition_to(&update.task.status) => {
guard.insert(task_id.to_string(), update);
Ok(true)
}
_ => Ok(false),
}
}
}
pub fn in_memory_store() -> Arc<dyn TaskStore> {
Arc::new(InMemoryTaskStore::new())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::{A2AMessage, TaskStatus};
fn sample_task(id: &str, status: TaskStatus) -> A2ATask {
A2ATask::new(id, A2AMessage::user("hi")).with_status(status)
}
#[tokio::test]
async fn upsert_get_roundtrip() {
let store = InMemoryTaskStore::new();
let mut stored = StoredTask::new(sample_task("t1", TaskStatus::Working));
stored.result = Some(A2ATaskResult::new("done"));
store.upsert(stored).await.unwrap();
let got = store.get("t1").await.unwrap().expect("task present");
assert_eq!(got.task.id, "t1");
assert_eq!(got.result.as_ref().unwrap().output, "done");
assert_eq!(got.task.status, TaskStatus::Working);
assert_eq!(got.created_at, got.updated_at);
}
#[tokio::test]
async fn upsert_updates_existing_in_place() {
let store = InMemoryTaskStore::new();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Submitted)))
.await
.unwrap();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Completed)))
.await
.unwrap();
let got = store.get("t1").await.unwrap().unwrap();
assert_eq!(got.task.status, TaskStatus::Completed);
}
#[tokio::test]
async fn get_missing_returns_none() {
let store = InMemoryTaskStore::new();
assert!(store.get("nope").await.unwrap().is_none());
}
#[tokio::test]
async fn list_filters_by_owner_and_status() {
let store = InMemoryTaskStore::new();
store
.upsert(StoredTask::new(
sample_task("t1", TaskStatus::Working).with_owner("a"),
))
.await
.unwrap();
store
.upsert(StoredTask::new(
sample_task("t2", TaskStatus::Completed).with_owner("a"),
))
.await
.unwrap();
store
.upsert(StoredTask::new(
sample_task("t3", TaskStatus::Working).with_owner("b"),
))
.await
.unwrap();
let all = store.list(&TaskFilter::new()).await.unwrap();
assert_eq!(all.len(), 3);
let only_a = store
.list(&TaskFilter::new().with_owner("a"))
.await
.unwrap();
assert_eq!(only_a.len(), 2);
let a_working = store
.list(
&TaskFilter::new()
.with_owner("a")
.with_statuses(vec![TaskStatus::Working]),
)
.await
.unwrap();
assert_eq!(a_working.len(), 1);
assert_eq!(a_working[0].task.id, "t1");
}
#[tokio::test]
async fn delete_removes_and_reports() {
let store = InMemoryTaskStore::new();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Submitted)))
.await
.unwrap();
assert!(store.delete("t1").await.unwrap());
assert!(!store.delete("t1").await.unwrap());
assert!(store.get("t1").await.unwrap().is_none());
}
#[tokio::test]
async fn evicts_oldest_when_full() {
let store = InMemoryTaskStore::with_max_tasks(2);
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Submitted)))
.await
.unwrap();
store
.upsert(StoredTask::new(sample_task("t2", TaskStatus::Submitted)))
.await
.unwrap();
store
.upsert(StoredTask::new(sample_task("t3", TaskStatus::Submitted)))
.await
.unwrap();
assert!(store.get("t1").await.unwrap().is_none());
assert!(store.get("t2").await.unwrap().is_some());
assert!(store.get("t3").await.unwrap().is_some());
}
#[tokio::test]
async fn touch_bumps_updated_at() {
let store = InMemoryTaskStore::new();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Submitted)))
.await
.unwrap();
let mut stored = store.get("t1").await.unwrap().unwrap();
stored.touch();
assert!(stored.updated_at >= stored.created_at);
}
#[tokio::test]
async fn store_is_clone_shareable() {
let store = InMemoryTaskStore::new();
let clone = store.clone();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Submitted)))
.await
.unwrap();
assert!(clone.get("t1").await.unwrap().is_some());
}
#[tokio::test]
async fn compare_and_update_rejects_stale_status() {
let store = InMemoryTaskStore::new();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Working)))
.await
.unwrap();
let mut stale = store.get("t1").await.unwrap().unwrap();
let mut cancelled = store.get("t1").await.unwrap().unwrap();
cancelled.task.status = TaskStatus::Cancelled;
cancelled.touch();
store.upsert(cancelled).await.unwrap();
stale.task.status = TaskStatus::Completed;
stale.touch();
assert!(!store.compare_and_update("t1", stale).await.unwrap());
assert_eq!(
store.get("t1").await.unwrap().unwrap().task.status,
TaskStatus::Cancelled
);
}
#[tokio::test]
async fn compare_and_update_applies_valid_transition() {
let store = InMemoryTaskStore::new();
store
.upsert(StoredTask::new(sample_task("t1", TaskStatus::Working)))
.await
.unwrap();
let mut done = store.get("t1").await.unwrap().unwrap();
done.task.status = TaskStatus::Completed;
done.touch();
assert!(store.compare_and_update("t1", done).await.unwrap());
assert_eq!(
store.get("t1").await.unwrap().unwrap().task.status,
TaskStatus::Completed
);
}
#[tokio::test]
async fn compare_and_update_missing_task_is_noop() {
let store = InMemoryTaskStore::new();
let update = StoredTask::new(sample_task("ghost", TaskStatus::Completed));
assert!(!store.compare_and_update("ghost", update).await.unwrap());
assert!(store.get("ghost").await.unwrap().is_none());
}
}