Skip to main content

klieo_a2a/
task_store.rs

1//! `A2aTaskStore` — durable Task persistence over [`klieo_core::KvStore`].
2//!
3//! Layout in bucket `a2a.tasks`:
4//! - `task.<task_id>` → JSON-encoded [`Task`].
5//! - `index.context.<context_id>` → JSON array of task ids in that context.
6//!
7//! Index entries are maintained on `put` / `delete`. `list(Some(ctx))`
8//! reads the index, then fetches each task. `list(None)` is not
9//! supported in v0.0.1 (would require a global index entry which is a
10//! hot key on writes); callers must filter by context.
11
12use 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
23/// Default KV bucket name for [`A2aTaskStore`] entries.
24pub const DEFAULT_BUCKET: &str = "a2a.tasks";
25
26/// Task store over a CAS-style KV.
27pub struct A2aTaskStore {
28    kv: Arc<dyn KvStore>,
29    bucket: String,
30    event_sink: Option<TaskEventSink>,
31    // Per-task monotonic event counters. Shared with the HTTP transport
32    // via `next_event_id` to stamp pubsub and SSE events consistently.
33    runtime: DashMap<String, Arc<AtomicU64>>,
34}
35
36impl A2aTaskStore {
37    /// Build a new store backed by `kv` writing under `bucket` (typical:
38    /// [`DEFAULT_BUCKET`]).
39    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    /// Wire a [`TaskEventSink`] so state transitions surface on the
49    /// configured pubsub. Returns `Self` for builder-style chaining.
50    pub fn with_event_sink(mut self, sink: TaskEventSink) -> Self {
51        self.event_sink = Some(sink);
52        self
53    }
54
55    /// Atomically increment and return the next event id for the
56    /// given task. Returns 1 on first call per task; subsequent
57    /// calls are monotonic per task_id.
58    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    /// Emit a task event to the cross-replica fanout sink.
68    ///
69    /// Pubsub publish failures (bus disconnect, encode failure, invalid
70    /// subject segment) are logged at `warn` server-side but do not
71    /// propagate — `emit` is best-effort by design: the resume buffer
72    /// (NATS-KV) is the durable layer; the bus is the live-tail layer.
73    /// Per ADR-018.
74    ///
75    /// Silently dropping the publish error here would hide cross-replica
76    /// fanout regressions, so the typed cause is preserved via
77    /// `Error::source()` in the trace.
78    #[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    /// Persist a task and update its context index.
107    #[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        // Update index.
123        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    /// Fetch a task by id.
149    #[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    /// List tasks. `context_id = Some(ctx)` reads the index; `None` is
167    /// not supported in v0.0.1 and returns an empty vec — callers must
168    /// supply a context.
169    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        // try_join_all preserves input order — tasks returned in index order.
179        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    /// Delete a task and remove its id from the context index.
185    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        // Remove the counter after the KV delete. A concurrent next_event_id
189        // call between kv.delete and this remove can re-insert the counter, but
190        // the next delete() invocation will remove it again. The counter is a
191        // soft in-process cache; the KV entry is the authoritative state.
192        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        // Subscribe before put so the message is not missed.
234        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        // Subscribe before put.
284        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}