use crate::error::A2aError;
use crate::server::{TaskEvent, TaskEventSink};
use crate::types::Task;
use bytes::Bytes;
use dashmap::DashMap;
use futures::future::try_join_all;
use klieo_core::KvStore;
use std::sync::atomic::AtomicU64;
use std::sync::Arc;
use tracing::instrument;
pub const DEFAULT_BUCKET: &str = "a2a.tasks";
pub struct A2aTaskStore {
kv: Arc<dyn KvStore>,
bucket: String,
event_sink: Option<TaskEventSink>,
runtime: DashMap<String, Arc<AtomicU64>>,
}
impl A2aTaskStore {
pub fn new(kv: Arc<dyn KvStore>, bucket: String) -> Self {
Self {
kv,
bucket,
event_sink: None,
runtime: DashMap::new(),
}
}
pub fn with_event_sink(mut self, sink: TaskEventSink) -> Self {
self.event_sink = Some(sink);
self
}
pub(crate) fn next_event_id(&self, task_id: &str) -> u64 {
let counter = self
.runtime
.entry(task_id.to_string())
.or_insert_with(|| Arc::new(AtomicU64::new(0)))
.clone();
counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1
}
#[instrument(
skip_all,
fields(klieo.stream_id = %event.task_id),
level = "debug",
)]
async fn emit(&self, event: TaskEvent) {
if let Some(sink) = &self.event_sink {
let task_id = event.task_id.clone();
if let Err(e) = sink.send(event).await {
tracing::warn!(
target: "a2a.store",
task_id = %task_id,
error = %e,
source = ?std::error::Error::source(&e),
"task event publish failed; cross-replica fanout degraded",
);
}
}
}
fn task_key(&self, id: &str) -> String {
format!("task.{id}")
}
fn index_key(&self, context_id: &str) -> String {
format!("index.context.{context_id}")
}
#[instrument(
skip_all,
fields(
db.system = "klieo-kv",
db.namespace = %self.bucket,
db.operation = "put",
klieo.stream_id = %task.id,
),
err,
)]
pub async fn put(&self, task: &Task) -> Result<(), A2aError> {
let bytes = Bytes::from(serde_json::to_vec(task)?);
self.kv
.put(&self.bucket, &self.task_key(&task.id), bytes)
.await?;
let key = self.index_key(&task.contextId);
let current = self.kv.get(&self.bucket, &key).await?;
let mut ids: Vec<String> = match current {
Some(entry) => serde_json::from_slice(&entry.value)?,
None => vec![],
};
if !ids.iter().any(|existing_id| existing_id == &task.id) {
ids.push(task.id.clone());
}
let updated = Bytes::from(serde_json::to_vec(&ids)?);
self.kv.put(&self.bucket, &key, updated).await?;
let event_id = self.next_event_id(&task.id);
self.emit(
TaskEvent::new(
task.id.clone(),
task.status,
task.history.last().cloned(),
task.status.is_terminal(),
)
.with_event_id(event_id),
)
.await;
Ok(())
}
#[instrument(
skip_all,
fields(
db.system = "klieo-kv",
db.namespace = %self.bucket,
db.operation = "get",
klieo.stream_id = %id,
),
err,
)]
pub async fn get(&self, id: &str) -> Result<Option<Task>, A2aError> {
match self.kv.get(&self.bucket, &self.task_key(id)).await? {
Some(entry) => Ok(Some(serde_json::from_slice(&entry.value)?)),
None => Ok(None),
}
}
pub async fn list(&self, context_id: Option<&str>) -> Result<Vec<Task>, A2aError> {
let Some(ctx) = context_id else {
return Ok(vec![]);
};
let entry = self.kv.get(&self.bucket, &self.index_key(ctx)).await?;
let ids: Vec<String> = match entry {
Some(e) => serde_json::from_slice(&e.value)?,
None => return Ok(vec![]),
};
let tasks: Vec<Option<Task>> =
try_join_all(ids.iter().map(|id| self.get(id.as_str()))).await?;
Ok(tasks.into_iter().flatten().collect())
}
pub async fn delete(&self, id: &str) -> Result<(), A2aError> {
let task = self.get(id).await?;
self.kv.delete(&self.bucket, &self.task_key(id)).await?;
self.runtime.remove(id);
if let Some(t) = task {
let key = self.index_key(&t.contextId);
if let Some(entry) = self.kv.get(&self.bucket, &key).await? {
let mut ids: Vec<String> = serde_json::from_slice(&entry.value)?;
ids.retain(|existing_id| existing_id != id);
let updated = Bytes::from(serde_json::to_vec(&ids)?);
self.kv.put(&self.bucket, &key, updated).await?;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::server::{TaskEvent, TaskEventSink};
use crate::types::TaskStatus;
use klieo_bus_memory::MemoryBus;
use klieo_core::{DurableName, Pubsub};
use std::sync::Arc;
use tokio_stream::StreamExt as _;
fn make_task(id: &str, status: TaskStatus) -> Task {
Task {
id: id.into(),
contextId: "ctx-1".into(),
status,
artifacts: vec![],
history: vec![],
metadata: None,
}
}
#[tokio::test]
async fn task_store_emits_event_on_put() {
let bus = Arc::new(MemoryBus::new());
let pubsub: Arc<dyn Pubsub> = bus.pubsub.clone();
let sink = TaskEventSink::new(pubsub.clone());
let subject = "klieo.a2a.task.t-1";
let durable = DurableName::new("test-task-store-t1");
let mut stream = pubsub.subscribe(subject, durable).await.unwrap();
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into()).with_event_sink(sink);
store
.put(&make_task("t-1", TaskStatus::Submitted))
.await
.unwrap();
let msg = tokio::time::timeout(std::time::Duration::from_millis(500), stream.next())
.await
.expect("timeout")
.expect("stream ended")
.expect("bus error");
let event: TaskEvent = serde_json::from_slice(&msg.payload).unwrap();
msg.ack.ack().await.unwrap();
assert_eq!(event.task_id, "t-1");
assert!(matches!(event.status, TaskStatus::Submitted));
assert!(!event.final_event);
}
#[tokio::test]
async fn task_store_emits_final_event_on_terminal_status() {
assert_final_event_for_status(TaskStatus::Completed).await;
}
#[tokio::test]
async fn task_store_emits_final_event_for_failed() {
assert_final_event_for_status(TaskStatus::Failed).await;
}
#[tokio::test]
async fn task_store_emits_final_event_for_canceled() {
assert_final_event_for_status(TaskStatus::Canceled).await;
}
#[tokio::test]
async fn task_store_emits_final_event_for_rejected() {
assert_final_event_for_status(TaskStatus::Rejected).await;
}
async fn assert_final_event_for_status(status: TaskStatus) {
let bus = Arc::new(MemoryBus::new());
let pubsub: Arc<dyn Pubsub> = bus.pubsub.clone();
let sink = TaskEventSink::new(pubsub.clone());
let task_id = format!("t-terminal-{status:?}");
let subject = format!("klieo.a2a.task.{task_id}");
let durable = DurableName::new(format!("test-final-{status:?}"));
let mut stream = pubsub.subscribe(&subject, durable).await.unwrap();
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into()).with_event_sink(sink);
store.put(&make_task(&task_id, status)).await.unwrap();
let msg = tokio::time::timeout(std::time::Duration::from_millis(500), stream.next())
.await
.expect("timeout")
.expect("stream ended")
.expect("bus error");
let event: TaskEvent = serde_json::from_slice(&msg.payload).unwrap();
msg.ack.ack().await.unwrap();
assert!(
event.final_event,
"{:?} must set final_event=true",
event.status
);
}
#[tokio::test]
async fn delete_clears_runtime_counter_so_new_task_starts_at_one() {
let bus = Arc::new(MemoryBus::new());
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
store
.put(&make_task("t-del", TaskStatus::Submitted))
.await
.unwrap();
assert_eq!(store.next_event_id("t-del"), 2);
store.delete("t-del").await.unwrap();
assert_eq!(store.next_event_id("t-del"), 1);
}
#[tokio::test]
async fn next_event_id_is_per_task_and_monotonic() {
let bus = Arc::new(MemoryBus::new());
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
assert_eq!(store.next_event_id("t1"), 1);
assert_eq!(store.next_event_id("t1"), 2);
assert_eq!(store.next_event_id("t2"), 1);
assert_eq!(store.next_event_id("t1"), 3);
}
#[tokio::test]
async fn list_returns_tasks_in_index_order() {
let bus = Arc::new(MemoryBus::new());
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
store
.put(&make_task("t-a", TaskStatus::Submitted))
.await
.unwrap();
store
.put(&make_task("t-b", TaskStatus::Working))
.await
.unwrap();
store
.put(&make_task("t-c", TaskStatus::Submitted))
.await
.unwrap();
let tasks = store.list(Some("ctx-1")).await.unwrap();
assert_eq!(tasks.len(), 3);
let ids: Vec<&str> = tasks.iter().map(|t| t.id.as_str()).collect();
assert_eq!(ids, vec!["t-a", "t-b", "t-c"]);
}
#[tokio::test]
async fn list_with_none_context_returns_empty() {
let bus = Arc::new(MemoryBus::new());
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
store
.put(&make_task("t-x", TaskStatus::Submitted))
.await
.unwrap();
let tasks = store.list(None).await.unwrap();
assert!(tasks.is_empty());
}
#[tokio::test]
async fn list_with_unknown_context_returns_empty() {
let bus = Arc::new(MemoryBus::new());
let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
let tasks = store.list(Some("no-such-ctx")).await.unwrap();
assert!(tasks.is_empty());
}
}