klieo-a2a 3.3.0

Durable A2A v1.0 protocol layer atop klieo-bus traits.
Documentation
//! `A2aTaskStore` — durable Task persistence over [`klieo_core::KvStore`].
//!
//! Layout in bucket `a2a.tasks`:
//! - `task.<task_id>` → JSON-encoded [`Task`].
//! - `index.context.<context_id>` → JSON array of task ids in that context.
//!
//! Index entries are maintained on `put` / `delete`. `list(Some(ctx))`
//! reads the index, then fetches each task. `list(None)` is not
//! supported in v0.0.1 (would require a global index entry which is a
//! hot key on writes); callers must filter by context.

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;

/// Default KV bucket name for [`A2aTaskStore`] entries.
pub const DEFAULT_BUCKET: &str = "a2a.tasks";

/// Task store over a CAS-style KV.
pub struct A2aTaskStore {
    kv: Arc<dyn KvStore>,
    bucket: String,
    event_sink: Option<TaskEventSink>,
    // Per-task monotonic event counters. Shared with the HTTP transport
    // via `next_event_id` to stamp pubsub and SSE events consistently.
    runtime: DashMap<String, Arc<AtomicU64>>,
}

impl A2aTaskStore {
    /// Build a new store backed by `kv` writing under `bucket` (typical:
    /// [`DEFAULT_BUCKET`]).
    pub fn new(kv: Arc<dyn KvStore>, bucket: String) -> Self {
        Self {
            kv,
            bucket,
            event_sink: None,
            runtime: DashMap::new(),
        }
    }

    /// Wire a [`TaskEventSink`] so state transitions surface on the
    /// configured pubsub. Returns `Self` for builder-style chaining.
    pub fn with_event_sink(mut self, sink: TaskEventSink) -> Self {
        self.event_sink = Some(sink);
        self
    }

    /// Atomically increment and return the next event id for the
    /// given task. Returns 1 on first call per task; subsequent
    /// calls are monotonic per task_id.
    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
    }

    /// Emit a task event to the cross-replica fanout sink.
    ///
    /// Pubsub publish failures (bus disconnect, encode failure, invalid
    /// subject segment) are logged at `warn` server-side but do not
    /// propagate — `emit` is best-effort by design: the resume buffer
    /// (NATS-KV) is the durable layer; the bus is the live-tail layer.
    /// Per ADR-018.
    ///
    /// Silently dropping the publish error here would hide cross-replica
    /// fanout regressions, so the typed cause is preserved via
    /// `Error::source()` in the trace.
    #[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}")
    }

    /// Persist a task and update its context index.
    #[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?;
        // Update index.
        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(())
    }

    /// Fetch a task by id.
    #[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),
        }
    }

    /// List tasks. `context_id = Some(ctx)` reads the index; `None` is
    /// not supported in v0.0.1 and returns an empty vec — callers must
    /// supply a context.
    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![]),
        };
        // try_join_all preserves input order — tasks returned in index order.
        let tasks: Vec<Option<Task>> =
            try_join_all(ids.iter().map(|id| self.get(id.as_str()))).await?;
        Ok(tasks.into_iter().flatten().collect())
    }

    /// Delete a task and remove its id from the context index.
    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?;
        // Remove the counter after the KV delete. A concurrent next_event_id
        // call between kv.delete and this remove can re-insert the counter, but
        // the next delete() invocation will remove it again. The counter is a
        // soft in-process cache; the KV entry is the authoritative state.
        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());

        // Subscribe before put so the message is not missed.
        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:?}");

        // Subscribe before put.
        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());
    }
}