1use crate::error::A2aError;
13use crate::server::{TaskEvent, TaskEventSink};
14use crate::types::Task;
15use bytes::Bytes;
16use dashmap::DashMap;
17use futures::future::try_join_all;
18use klieo_core::KvStore;
19use std::sync::atomic::AtomicU64;
20use std::sync::Arc;
21use tracing::instrument;
22
23pub const DEFAULT_BUCKET: &str = "a2a.tasks";
25
26pub struct A2aTaskStore {
28 kv: Arc<dyn KvStore>,
29 bucket: String,
30 event_sink: Option<TaskEventSink>,
31 runtime: DashMap<String, Arc<AtomicU64>>,
34}
35
36impl A2aTaskStore {
37 pub fn new(kv: Arc<dyn KvStore>, bucket: String) -> Self {
40 Self {
41 kv,
42 bucket,
43 event_sink: None,
44 runtime: DashMap::new(),
45 }
46 }
47
48 pub fn with_event_sink(mut self, sink: TaskEventSink) -> Self {
51 self.event_sink = Some(sink);
52 self
53 }
54
55 pub(crate) fn next_event_id(&self, task_id: &str) -> u64 {
59 let counter = self
60 .runtime
61 .entry(task_id.to_string())
62 .or_insert_with(|| Arc::new(AtomicU64::new(0)))
63 .clone();
64 counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1
65 }
66
67 #[instrument(
79 skip_all,
80 fields(klieo.stream_id = %event.task_id),
81 level = "debug",
82 )]
83 async fn emit(&self, event: TaskEvent) {
84 if let Some(sink) = &self.event_sink {
85 let task_id = event.task_id.clone();
86 if let Err(e) = sink.send(event).await {
87 tracing::warn!(
88 target: "a2a.store",
89 task_id = %task_id,
90 error = %e,
91 source = ?std::error::Error::source(&e),
92 "task event publish failed; cross-replica fanout degraded",
93 );
94 }
95 }
96 }
97
98 fn task_key(&self, id: &str) -> String {
99 format!("task.{id}")
100 }
101
102 fn index_key(&self, context_id: &str) -> String {
103 format!("index.context.{context_id}")
104 }
105
106 #[instrument(
108 skip_all,
109 fields(
110 db.system = "klieo-kv",
111 db.namespace = %self.bucket,
112 db.operation = "put",
113 klieo.stream_id = %task.id,
114 ),
115 err,
116 )]
117 pub async fn put(&self, task: &Task) -> Result<(), A2aError> {
118 let bytes = Bytes::from(serde_json::to_vec(task)?);
119 self.kv
120 .put(&self.bucket, &self.task_key(&task.id), bytes)
121 .await?;
122 let key = self.index_key(&task.contextId);
124 let current = self.kv.get(&self.bucket, &key).await?;
125 let mut ids: Vec<String> = match current {
126 Some(entry) => serde_json::from_slice(&entry.value)?,
127 None => vec![],
128 };
129 if !ids.iter().any(|existing_id| existing_id == &task.id) {
130 ids.push(task.id.clone());
131 }
132 let updated = Bytes::from(serde_json::to_vec(&ids)?);
133 self.kv.put(&self.bucket, &key, updated).await?;
134 let event_id = self.next_event_id(&task.id);
135 self.emit(
136 TaskEvent::new(
137 task.id.clone(),
138 task.status,
139 task.history.last().cloned(),
140 task.status.is_terminal(),
141 )
142 .with_event_id(event_id),
143 )
144 .await;
145 Ok(())
146 }
147
148 #[instrument(
150 skip_all,
151 fields(
152 db.system = "klieo-kv",
153 db.namespace = %self.bucket,
154 db.operation = "get",
155 klieo.stream_id = %id,
156 ),
157 err,
158 )]
159 pub async fn get(&self, id: &str) -> Result<Option<Task>, A2aError> {
160 match self.kv.get(&self.bucket, &self.task_key(id)).await? {
161 Some(entry) => Ok(Some(serde_json::from_slice(&entry.value)?)),
162 None => Ok(None),
163 }
164 }
165
166 pub async fn list(&self, context_id: Option<&str>) -> Result<Vec<Task>, A2aError> {
170 let Some(ctx) = context_id else {
171 return Ok(vec![]);
172 };
173 let entry = self.kv.get(&self.bucket, &self.index_key(ctx)).await?;
174 let ids: Vec<String> = match entry {
175 Some(e) => serde_json::from_slice(&e.value)?,
176 None => return Ok(vec![]),
177 };
178 let tasks: Vec<Option<Task>> =
180 try_join_all(ids.iter().map(|id| self.get(id.as_str()))).await?;
181 Ok(tasks.into_iter().flatten().collect())
182 }
183
184 pub async fn delete(&self, id: &str) -> Result<(), A2aError> {
186 let task = self.get(id).await?;
187 self.kv.delete(&self.bucket, &self.task_key(id)).await?;
188 self.runtime.remove(id);
193 if let Some(t) = task {
194 let key = self.index_key(&t.contextId);
195 if let Some(entry) = self.kv.get(&self.bucket, &key).await? {
196 let mut ids: Vec<String> = serde_json::from_slice(&entry.value)?;
197 ids.retain(|existing_id| existing_id != id);
198 let updated = Bytes::from(serde_json::to_vec(&ids)?);
199 self.kv.put(&self.bucket, &key, updated).await?;
200 }
201 }
202 Ok(())
203 }
204}
205
206#[cfg(test)]
207mod tests {
208 use super::*;
209 use crate::server::{TaskEvent, TaskEventSink};
210 use crate::types::TaskStatus;
211 use klieo_bus_memory::MemoryBus;
212 use klieo_core::{DurableName, Pubsub};
213 use std::sync::Arc;
214 use tokio_stream::StreamExt as _;
215
216 fn make_task(id: &str, status: TaskStatus) -> Task {
217 Task {
218 id: id.into(),
219 contextId: "ctx-1".into(),
220 status,
221 artifacts: vec![],
222 history: vec![],
223 metadata: None,
224 }
225 }
226
227 #[tokio::test]
228 async fn task_store_emits_event_on_put() {
229 let bus = Arc::new(MemoryBus::new());
230 let pubsub: Arc<dyn Pubsub> = bus.pubsub.clone();
231 let sink = TaskEventSink::new(pubsub.clone());
232
233 let subject = "klieo.a2a.task.t-1";
235 let durable = DurableName::new("test-task-store-t1");
236 let mut stream = pubsub.subscribe(subject, durable).await.unwrap();
237
238 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into()).with_event_sink(sink);
239 store
240 .put(&make_task("t-1", TaskStatus::Submitted))
241 .await
242 .unwrap();
243
244 let msg = tokio::time::timeout(std::time::Duration::from_millis(500), stream.next())
245 .await
246 .expect("timeout")
247 .expect("stream ended")
248 .expect("bus error");
249 let event: TaskEvent = serde_json::from_slice(&msg.payload).unwrap();
250 msg.ack.ack().await.unwrap();
251
252 assert_eq!(event.task_id, "t-1");
253 assert!(matches!(event.status, TaskStatus::Submitted));
254 assert!(!event.final_event);
255 }
256
257 #[tokio::test]
258 async fn task_store_emits_final_event_on_terminal_status() {
259 assert_final_event_for_status(TaskStatus::Completed).await;
260 }
261
262 #[tokio::test]
263 async fn task_store_emits_final_event_for_failed() {
264 assert_final_event_for_status(TaskStatus::Failed).await;
265 }
266
267 #[tokio::test]
268 async fn task_store_emits_final_event_for_canceled() {
269 assert_final_event_for_status(TaskStatus::Canceled).await;
270 }
271
272 #[tokio::test]
273 async fn task_store_emits_final_event_for_rejected() {
274 assert_final_event_for_status(TaskStatus::Rejected).await;
275 }
276
277 async fn assert_final_event_for_status(status: TaskStatus) {
278 let bus = Arc::new(MemoryBus::new());
279 let pubsub: Arc<dyn Pubsub> = bus.pubsub.clone();
280 let sink = TaskEventSink::new(pubsub.clone());
281 let task_id = format!("t-terminal-{status:?}");
282
283 let subject = format!("klieo.a2a.task.{task_id}");
285 let durable = DurableName::new(format!("test-final-{status:?}"));
286 let mut stream = pubsub.subscribe(&subject, durable).await.unwrap();
287
288 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into()).with_event_sink(sink);
289 store.put(&make_task(&task_id, status)).await.unwrap();
290
291 let msg = tokio::time::timeout(std::time::Duration::from_millis(500), stream.next())
292 .await
293 .expect("timeout")
294 .expect("stream ended")
295 .expect("bus error");
296 let event: TaskEvent = serde_json::from_slice(&msg.payload).unwrap();
297 msg.ack.ack().await.unwrap();
298
299 assert!(
300 event.final_event,
301 "{:?} must set final_event=true",
302 event.status
303 );
304 }
305
306 #[tokio::test]
307 async fn delete_clears_runtime_counter_so_new_task_starts_at_one() {
308 let bus = Arc::new(MemoryBus::new());
309 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
310
311 store
312 .put(&make_task("t-del", TaskStatus::Submitted))
313 .await
314 .unwrap();
315 assert_eq!(store.next_event_id("t-del"), 2);
316
317 store.delete("t-del").await.unwrap();
318
319 assert_eq!(store.next_event_id("t-del"), 1);
320 }
321
322 #[tokio::test]
323 async fn next_event_id_is_per_task_and_monotonic() {
324 let bus = Arc::new(MemoryBus::new());
325 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
326 assert_eq!(store.next_event_id("t1"), 1);
327 assert_eq!(store.next_event_id("t1"), 2);
328 assert_eq!(store.next_event_id("t2"), 1);
329 assert_eq!(store.next_event_id("t1"), 3);
330 }
331
332 #[tokio::test]
333 async fn list_returns_tasks_in_index_order() {
334 let bus = Arc::new(MemoryBus::new());
335 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
336
337 store
338 .put(&make_task("t-a", TaskStatus::Submitted))
339 .await
340 .unwrap();
341 store
342 .put(&make_task("t-b", TaskStatus::Working))
343 .await
344 .unwrap();
345 store
346 .put(&make_task("t-c", TaskStatus::Submitted))
347 .await
348 .unwrap();
349
350 let tasks = store.list(Some("ctx-1")).await.unwrap();
351 assert_eq!(tasks.len(), 3);
352 let ids: Vec<&str> = tasks.iter().map(|t| t.id.as_str()).collect();
353 assert_eq!(ids, vec!["t-a", "t-b", "t-c"]);
354 }
355
356 #[tokio::test]
357 async fn list_with_none_context_returns_empty() {
358 let bus = Arc::new(MemoryBus::new());
359 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
360 store
361 .put(&make_task("t-x", TaskStatus::Submitted))
362 .await
363 .unwrap();
364 let tasks = store.list(None).await.unwrap();
365 assert!(tasks.is_empty());
366 }
367
368 #[tokio::test]
369 async fn list_with_unknown_context_returns_empty() {
370 let bus = Arc::new(MemoryBus::new());
371 let store = A2aTaskStore::new(bus.kv.clone(), DEFAULT_BUCKET.into());
372 let tasks = store.list(Some("no-such-ctx")).await.unwrap();
373 assert!(tasks.is_empty());
374 }
375}