diff --git a/.gitmodules b/.gitmodules index 3c1225634..c3c0ed1dc 100644 --- a/.gitmodules +++ b/.gitmodules @@ -10,3 +10,6 @@ [submodule "vendor/tinyliveagents"] path = vendor/tinyliveagents url = https://github.com/tinyhumansai/tinyliveagents.git +[submodule "vendor/tinystoragedrivers"] + path = vendor/tinystoragedrivers + url = https://github.com/tinyhumansai/tinystoragedrivers.git diff --git a/AGENTS.md b/AGENTS.md index f22a6fb6b..76418a921 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -21,8 +21,10 @@ dedicated `types.rs` file and keep module-local unit tests in a sibling smallest useful API. Cargo features are package-local. `tinyagents-harness` exposes `sqlite`, -`tools`, `multimodal`, and `tracing`; `tinyagents-graph` exposes `sqlite` and -`tracing`; registry and session expose `tracing`. Tracing instrumentation is +`tools`, `multimodal`, `storage-drivers`, and `tracing`; `tinyagents-graph` +exposes `sqlite`, `storage-drivers`, and `tracing`; session exposes +`storage-drivers` and `tracing`; registry exposes `tracing`. The +`storage-drivers` features put each store on `tinystoragedrivers` ports. Tracing instrumentation is compiled out by default. Integration tests are in `crates/tinyagents-integration-tests/tests/`, covering serialization, graph routing, diff --git a/Cargo.lock b/Cargo.lock index 25bb8ca2e..cdd254dd5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1717,9 +1717,11 @@ dependencies = [ "rusqlite", "serde", "serde_json", + "sha2", "tempfile", "tinyagents-harness", "tinyinference-llm", + "tinystoragedrivers-core", "tinytools", "tokio", "tracing", @@ -1755,6 +1757,7 @@ dependencies = [ "tinyinference-image", "tinyinference-llm", "tinyinference-video", + "tinystoragedrivers-core", "tinytools", "tinytools-agent", "tokio", @@ -1884,6 +1887,8 @@ dependencies = [ "tempfile", "thiserror", "tinyagents-harness", + "tinystoragedrivers-core", + "tinystoragedrivers-sqlite", "tinytools-agent", "tokio", "tracing", @@ -1983,6 +1988,27 @@ dependencies = [ "url", ] +[[package]] +name = "tinystoragedrivers-core" +version = "0.3.0" +dependencies = [ + "async-trait", + "serde", + "serde_json", + "thiserror", + "tokio", +] + +[[package]] +name = "tinystoragedrivers-sqlite" +version = "0.3.0" +dependencies = [ + "rusqlite", + "serde_json", + "tinystoragedrivers-core", + "tokio", +] + [[package]] name = "tinystr" version = "0.8.3" diff --git a/Cargo.toml b/Cargo.toml index 951665032..0cef4eea5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -31,6 +31,9 @@ serde_json = "1" tempfile = "3" tokio = { version = "1", default-features = false, features = ["macros", "rt-multi-thread"] } tracing = "0.1" +# Storage ports from the vendored tinystoragedrivers submodule. Hosts that +# vendor tinyagents use this same copy, so there is one set of port types. +tinystoragedrivers-core = { path = "vendor/tinystoragedrivers/crates/tinystoragedrivers-core", version = "0.3.0" } uuid = { version = "1", features = ["v4"] } [workspace.lints.rust] @@ -52,3 +55,4 @@ strip = "debuginfo" [patch."https://github.com/tinyhumansai/tinytools"] tinytools-agent = { path = "vendor/tinytools/crates/tinytools-agent" } tinytools = { path = "vendor/tinytools/crates/tinytools" } + diff --git a/crates/tinyagents-graph/Cargo.toml b/crates/tinyagents-graph/Cargo.toml index ec4e01a53..894a68ffe 100644 --- a/crates/tinyagents-graph/Cargo.toml +++ b/crates/tinyagents-graph/Cargo.toml @@ -22,10 +22,17 @@ tinyinference-llm = { path = "../../vendor/tinyinference/crates/tinyinference-ll tinytools = { path = "../../vendor/tinytools/crates/tinytools", version = "0.5.0" } tokio = { workspace = true, features = ["sync", "time", "macros", "rt", "rt-multi-thread", "fs"] } tracing = { workspace = true } +# Storage ports (no database client) for `DriverCheckpointer`. +tinystoragedrivers-core = { workspace = true, optional = true } +# Collision-resistant ids for `DriverCheckpointer` keys too long to store. +sha2 = { version = "0.11", optional = true } [features] default = [] sqlite = ["dep:rusqlite"] +# `DriverCheckpointer`: checkpoints, pending writes and leases on any +# tinystoragedrivers backend the host opened. +storage-drivers = ["dep:tinystoragedrivers-core", "dep:sha2"] # Tracing instrumentation is now always compiled in (via the `tracing` crate # dependency above). This feature is retained as a no-op so downstream # feature forwards keep compiling. diff --git a/crates/tinyagents-graph/src/checkpoint/README.md b/crates/tinyagents-graph/src/checkpoint/README.md index c7f39a44d..8f28acd69 100644 --- a/crates/tinyagents-graph/src/checkpoint/README.md +++ b/crates/tinyagents-graph/src/checkpoint/README.md @@ -93,3 +93,25 @@ crashes mid-execution. - `DurabilityMode` trades persistence frequency for write volume — a mode that only checkpoints on interrupt/failure means a mid-run crash loses all progress since the last such boundary, not just the current node. + +## `DriverCheckpointer` (`storage-drivers` feature) + +`DriverCheckpointer` stores checkpoints, pending writes and execution leases in +a [tinystoragedrivers](https://github.com/tinyhumansai/tinystoragedrivers) +`DocumentStore`. Hosts use it to put graph durability on whichever backend they +opened, such as SQLite or MongoDB. The handle is bound to a tenant scope, so +tenants that share a thread id still never see each other's checkpoints. + +It uses four collections, all prefixed `graph` by default: + +- `_checkpoints`: one document per stored checkpoint, keyed by a + per-thread insertion sequence. Duplicate ids resolve to the latest write, + as in the append-only backends. +- `_threads`: the per-thread sequence counter, advanced with + compare-and-swap. +- `_writes`: merged pending writes per `(thread, namespace, checkpoint)`. +- `_leases`: per-thread execution leases, claimed, renewed and + released with compare-and-swap. + +It passes the same `testkit::conformance` contracts as the bundled backends: +checkpointer, writes, lineage and concurrency. diff --git a/crates/tinyagents-graph/src/checkpoint/drivers.rs b/crates/tinyagents-graph/src/checkpoint/drivers.rs new file mode 100644 index 000000000..cc4f24310 --- /dev/null +++ b/crates/tinyagents-graph/src/checkpoint/drivers.rs @@ -0,0 +1,495 @@ +//! A [`Checkpointer`] over `tinystoragedrivers` document ports. +//! +//! [`DriverCheckpointer`] puts graph checkpoints, their pending writes and the +//! per-thread execution lease on whatever backend the host opened (SQLite on a +//! desktop, MongoDB in the cloud, memory in tests). The document handle is +//! already bound to a tenant scope, so two tenants' threads never meet even +//! when they share a thread id. +//! +//! # Layout +//! +//! Three collections, named from a prefix (default `graph`): +//! +//! - `_checkpoints`: one document per stored checkpoint, +//! `{thread, namespace, seq, checkpoint_id, record}` (`namespace` is the +//! encoded subgraph namespace, so scoped reads filter on it). `seq` is a per-thread insertion +//! counter, so listing is a query sorted on it and duplicate checkpoint ids +//! resolve to the latest write, exactly like the append-only backends. +//! - `_threads`: one counter document per thread, advanced with a +//! compare-and-swap so concurrent writers never share a `seq`. +//! - `_writes`: one document per `(thread, namespace, checkpoint)` +//! holding the merged pending writes (see [`merge_writes`]). +//! - `_leases`: one document per leased thread, claimed and renewed +//! with compare-and-swap. + +use std::marker::PhantomData; +use std::sync::Arc; +use std::time::Duration; + +use async_trait::async_trait; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::{Value, json}; +use tinyagents_harness::error::{Result, TinyAgentsError}; +use tinystoragedrivers_core::{ + CollectionSpec, DocumentStore, DocumentStoreExt, ErrorKind, Filter, IndexSpec, Precondition, + Query, Sort, StorageError, Versioned, +}; +use tokio::sync::OnceCell; + +use super::{ + Checkpoint, CheckpointConfig, CheckpointId, CheckpointMetadata, Checkpointer, PendingWrite, + decode_json_err, merge_writes, require_checkpoint_id, +}; + +/// How many times a compare-and-swap loop retries before giving up. +const CAS_ATTEMPTS: usize = 64; + +/// Map a storage driver failure onto the graph's checkpoint error. +fn map_error(error: StorageError) -> TinyAgentsError { + TinyAgentsError::Checkpoint(format!("storage driver: {error}")) +} + +/// Longest id stored as is; longer ones are hashed. +const MAX_KEY_LEN: usize = 400; + +/// A document id for `parts`: length-prefixed so no two tuples collide, and +/// replaced by its SHA-256 when it would exceed the driver's id limit. +fn key(parts: &[&str]) -> String { + let joined: String = parts + .iter() + .map(|part| format!("{}:{part}", part.len())) + .collect::>() + .join("/"); + if joined.len() <= MAX_KEY_LEN { + joined + } else { + use sha2::{Digest, Sha256}; + let digest = Sha256::digest(joined.as_bytes()); + let hex: String = digest.iter().map(|byte| format!("{byte:02x}")).collect(); + format!("h:{hex}") + } +} + +/// `namespace` as one string, injectively: each component is +/// length-prefixed, so `["a", "b"]` and `["a/b"]` never meet. +fn namespace_key(namespace: &[String]) -> String { + namespace + .iter() + .map(|part| format!("{}:{part};", part.len())) + .collect() +} + +/// A [`Checkpointer`] that stores everything in a driver [`DocumentStore`]. +pub struct DriverCheckpointer { + docs: Arc, + checkpoints: String, + threads: String, + writes: String, + leases: String, + declared: Arc>, + _state: PhantomData State>, +} + +impl Clone for DriverCheckpointer { + fn clone(&self) -> Self { + Self { + docs: Arc::clone(&self.docs), + checkpoints: self.checkpoints.clone(), + threads: self.threads.clone(), + writes: self.writes.clone(), + leases: self.leases.clone(), + declared: Arc::clone(&self.declared), + _state: PhantomData, + } + } +} + +impl std::fmt::Debug for DriverCheckpointer { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("DriverCheckpointer") + .field("checkpoints", &self.checkpoints) + .finish_non_exhaustive() + } +} + +impl DriverCheckpointer { + /// Store checkpoints in the `graph_*` collections of `docs`. + pub fn new(docs: Arc) -> Self { + Self::with_prefix(docs, "graph") + } + + /// Store checkpoints in the `_*` collections of `docs`, so + /// several independent checkpointers can share one backend. + pub fn with_prefix(docs: Arc, prefix: &str) -> Self { + Self { + docs, + checkpoints: format!("{prefix}_checkpoints"), + threads: format!("{prefix}_threads"), + writes: format!("{prefix}_writes"), + leases: format!("{prefix}_leases"), + declared: Arc::new(OnceCell::new()), + _state: PhantomData, + } + } + + async fn declared(&self) -> Result<()> { + self.declared + .get_or_try_init(|| async { + let specs = [ + CollectionSpec::new(&self.checkpoints) + .index(IndexSpec::new("by_thread_seq", ["thread", "seq"])) + .index(IndexSpec::new("by_thread_id", ["thread", "checkpoint_id"])) + .index(IndexSpec::new( + "by_thread_namespace", + ["thread", "namespace", "seq"], + )), + CollectionSpec::new(&self.threads), + CollectionSpec::new(&self.writes) + .index(IndexSpec::new("by_thread", ["thread"])), + CollectionSpec::new(&self.leases), + ]; + for spec in &specs { + self.docs.ensure_collection(spec).await.map_err(map_error)?; + } + Ok::<(), TinyAgentsError>(()) + }) + .await + .map(|_| ()) + } + + /// Reserve the next insertion sequence number for `thread`. + async fn next_seq(&self, thread: &str) -> Result { + let id = key(&[thread]); + for _ in 0..CAS_ATTEMPTS { + let current = self.docs.get(&self.threads, &id).await.map_err(map_error)?; + let (seq, precondition) = match ¤t { + Some(found) => ( + found.doc.get("next").and_then(Value::as_u64).unwrap_or(0), + found.unchanged(), + ), + None => (0, Precondition::Absent), + }; + let doc = json!({ "thread": thread, "next": seq + 1 }); + match self.docs.put(&self.threads, &id, doc, precondition).await { + Ok(_) => return Ok(seq), + Err(error) if error.kind() == ErrorKind::Conflict => continue, + Err(error) => return Err(map_error(error)), + } + } + Err(TinyAgentsError::Checkpoint(format!( + "could not reserve a checkpoint sequence for thread `{thread}`" + ))) + } + + /// Every stored checkpoint document of `thread`, in insertion order. + async fn thread_docs(&self, thread: &str) -> Result>> { + self.declared().await?; + let query = Query::filter(Filter::eq("thread", thread)).sort(Sort::asc("seq")); + self.docs + .query_all(&self.checkpoints, &query) + .await + .map_err(map_error) + } + + fn writes_id(thread: &str, namespace: &[String], checkpoint_id: &str) -> String { + key(&[thread, &namespace_key(namespace), checkpoint_id]) + } + + async fn drop_writes(&self, filter: Filter) -> Result<()> { + self.docs + .delete_where(&self.writes, &filter) + .await + .map(|_| ()) + .map_err(map_error) + } +} + +impl DriverCheckpointer +where + State: DeserializeOwned, +{ + fn decode(stored: Versioned) -> Result> { + let record = stored.doc.get("record").cloned().unwrap_or(Value::Null); + // Through the shared classifier, so a `State` that no longer decodes + // is tagged `[schema]` exactly as the file and SQLite backends tag it + // (durable delegations prune such a checkpoint and start fresh). + let mut checkpoint: Checkpoint = serde_json::from_value(record) + .map_err(|error| decode_json_err("storage driver checkpointer", "record", error))?; + checkpoint.normalize(); + Ok(checkpoint) + } +} + +#[async_trait] +impl Checkpointer for DriverCheckpointer +where + State: Serialize + DeserializeOwned + Send + Sync + 'static, +{ + async fn put(&self, checkpoint: Checkpoint) -> Result { + self.declared().await?; + let id = CheckpointId::new(checkpoint.checkpoint_id.clone()); + let seq = self.next_seq(&checkpoint.thread_id).await?; + let doc = json!({ + "thread": checkpoint.thread_id, + "namespace": namespace_key(&checkpoint.namespace), + "seq": seq, + "checkpoint_id": checkpoint.checkpoint_id, + "record": serde_json::to_value(&checkpoint)?, + }); + let doc_id = key(&[&checkpoint.thread_id, &format!("{seq:020}")]); + self.docs + .put(&self.checkpoints, &doc_id, doc, Precondition::Absent) + .await + .map_err(map_error)?; + Ok(id) + } + + async fn get( + &self, + thread_id: &str, + checkpoint_id: Option<&str>, + ) -> Result>> { + self.declared().await?; + let mut filter = Filter::eq("thread", thread_id); + if let Some(id) = checkpoint_id { + filter = filter.and(Filter::eq("checkpoint_id", id)); + } + let query = Query::filter(filter).sort(Sort::desc("seq")).limit(1); + let page = self + .docs + .query(&self.checkpoints, &query) + .await + .map_err(map_error)?; + page.items.into_iter().next().map(Self::decode).transpose() + } + + /// One indexed query on the stored namespace, so a parent run and a + /// subgraph sharing a thread never load each other's checkpoints, even + /// when they reuse a checkpoint id. + async fn get_scoped( + &self, + thread_id: &str, + checkpoint_id: Option<&str>, + namespace: &[String], + ) -> Result>> { + self.declared().await?; + let mut filter = + Filter::eq("thread", thread_id).and(Filter::eq("namespace", namespace_key(namespace))); + if let Some(id) = checkpoint_id { + filter = filter.and(Filter::eq("checkpoint_id", id)); + } + let query = Query::filter(filter).sort(Sort::desc("seq")).limit(1); + let page = self + .docs + .query(&self.checkpoints, &query) + .await + .map_err(map_error)?; + page.items.into_iter().next().map(Self::decode).transpose() + } + + async fn list(&self, thread_id: &str) -> Result> { + Ok(self + .get_thread(thread_id) + .await? + .iter() + .map(Checkpoint::to_metadata) + .collect()) + } + + async fn get_thread(&self, thread_id: &str) -> Result>> { + self.thread_docs(thread_id) + .await? + .into_iter() + .map(Self::decode) + .collect() + } + + async fn list_threads(&self) -> Result> { + self.declared().await?; + let counters = self + .docs + .query_all(&self.threads, &Query::all()) + .await + .map_err(map_error)?; + let mut threads = Vec::new(); + for counter in counters { + let Some(thread) = counter.doc.get("thread").and_then(Value::as_str) else { + continue; + }; + let live = self + .docs + .count(&self.checkpoints, &Filter::eq("thread", thread)) + .await + .map_err(map_error)?; + if live > 0 { + threads.push(thread.to_owned()); + } + } + Ok(threads) + } + + async fn delete_thread(&self, thread_id: &str) -> Result<()> { + self.declared().await?; + self.docs + .delete_where(&self.checkpoints, &Filter::eq("thread", thread_id)) + .await + .map_err(map_error)?; + // Keep the sequence counter: a later thread of the same name continues + // after it, so its records never interleave with stale cursors. + self.drop_writes(Filter::eq("thread", thread_id)).await + } + + async fn delete_checkpoints(&self, thread_id: &str, ids: &[String]) -> Result { + if ids.is_empty() { + return Ok(0); + } + self.declared().await?; + let filter = Filter::eq("thread", thread_id).and(Filter::one_of( + "checkpoint_id", + ids.iter().map(String::as_str), + )); + let removed = self + .docs + .delete_where(&self.checkpoints, &filter) + .await + .map_err(map_error)?; + self.drop_writes(Filter::eq("thread", thread_id).and(Filter::one_of( + "checkpoint_id", + ids.iter().map(String::as_str), + ))) + .await?; + Ok(usize::try_from(removed).unwrap_or(usize::MAX)) + } + + async fn put_writes(&self, config: &CheckpointConfig, writes: &[PendingWrite]) -> Result<()> { + let checkpoint_id = require_checkpoint_id(config)?; + if writes.is_empty() { + return Ok(()); + } + self.declared().await?; + let id = Self::writes_id(&config.thread_id, &config.namespace, &checkpoint_id); + for _ in 0..CAS_ATTEMPTS { + let current = self.docs.get(&self.writes, &id).await.map_err(map_error)?; + let (mut stored, precondition): (Vec, Precondition) = match ¤t { + Some(found) => ( + serde_json::from_value( + found.doc.get("writes").cloned().unwrap_or(Value::Null), + )?, + found.unchanged(), + ), + None => (Vec::new(), Precondition::Absent), + }; + merge_writes(&mut stored, writes); + let doc = json!({ + "thread": config.thread_id, + "namespace": config.namespace, + "checkpoint_id": checkpoint_id, + "writes": serde_json::to_value(&stored)?, + }); + match self.docs.put(&self.writes, &id, doc, precondition).await { + Ok(_) => return Ok(()), + Err(error) if error.kind() == ErrorKind::Conflict => continue, + Err(error) => return Err(map_error(error)), + } + } + Err(TinyAgentsError::Checkpoint(format!( + "could not record pending writes for thread `{}`", + config.thread_id + ))) + } + + async fn get_writes(&self, config: &CheckpointConfig) -> Result> { + let Some(checkpoint_id) = self.resolve_write_target(config).await? else { + return Ok(Vec::new()); + }; + self.declared().await?; + let id = Self::writes_id(&config.thread_id, &config.namespace, &checkpoint_id); + let Some(found) = self.docs.get(&self.writes, &id).await.map_err(map_error)? else { + return Ok(Vec::new()); + }; + Ok(serde_json::from_value( + found.doc.get("writes").cloned().unwrap_or(Value::Null), + )?) + } + + async fn try_claim(&self, thread: &str, owner: &str, ttl: Duration) -> Result { + self.declared().await?; + let id = key(&[thread]); + for _ in 0..CAS_ATTEMPTS { + let now = tinyagents_harness::ids::now_ms(); + let current = self.docs.get(&self.leases, &id).await.map_err(map_error)?; + let precondition = match ¤t { + None => Precondition::Absent, + Some(found) => { + let held_by = found.doc.get("owner").and_then(Value::as_str); + let expires = found.doc.get("expires_at_ms").and_then(Value::as_u64); + let live = expires.is_some_and(|at| at > now); + if live && held_by != Some(owner) { + return Ok(false); + } + found.unchanged() + } + }; + let expires_at_ms = + now.saturating_add(u64::try_from(ttl.as_millis()).unwrap_or(u64::MAX)); + let doc = json!({ "owner": owner, "expires_at_ms": expires_at_ms }); + match self.docs.put(&self.leases, &id, doc, precondition).await { + Ok(_) => return Ok(true), + Err(error) if error.kind() == ErrorKind::Conflict => continue, + Err(error) => return Err(map_error(error)), + } + } + Ok(false) + } + + async fn renew(&self, thread: &str, owner: &str, ttl: Duration) -> Result { + self.declared().await?; + let id = key(&[thread]); + let now = tinyagents_harness::ids::now_ms(); + let Some(found) = self.docs.get(&self.leases, &id).await.map_err(map_error)? else { + return Ok(false); + }; + let held_by = found.doc.get("owner").and_then(Value::as_str); + let live = found + .doc + .get("expires_at_ms") + .and_then(Value::as_u64) + .is_some_and(|at| at > now); + if held_by != Some(owner) || !live { + return Ok(false); + } + let expires_at_ms = now.saturating_add(u64::try_from(ttl.as_millis()).unwrap_or(u64::MAX)); + let doc = json!({ "owner": owner, "expires_at_ms": expires_at_ms }); + match self + .docs + .put(&self.leases, &id, doc, found.unchanged()) + .await + { + Ok(_) => Ok(true), + Err(error) if error.kind() == ErrorKind::Conflict => Ok(false), + Err(error) => Err(map_error(error)), + } + } + + async fn release(&self, thread: &str, owner: &str) -> Result<()> { + self.declared().await?; + let id = key(&[thread]); + let Some(found) = self.docs.get(&self.leases, &id).await.map_err(map_error)? else { + return Ok(()); + }; + if found.doc.get("owner").and_then(Value::as_str) != Some(owner) { + return Ok(()); + } + match self.docs.delete(&self.leases, &id, found.unchanged()).await { + Ok(_) => Ok(()), + // Someone else reclaimed it in between; it is no longer ours. + Err(error) if error.kind() == ErrorKind::Conflict => Ok(()), + Err(error) => Err(map_error(error)), + } + } +} + +#[cfg(test)] +#[path = "drivers_tests.rs"] +mod tests; diff --git a/crates/tinyagents-graph/src/checkpoint/drivers_tests.rs b/crates/tinyagents-graph/src/checkpoint/drivers_tests.rs new file mode 100644 index 000000000..2a8ee4d35 --- /dev/null +++ b/crates/tinyagents-graph/src/checkpoint/drivers_tests.rs @@ -0,0 +1,224 @@ +use std::time::Duration; + +use super::*; +use crate::testkit::conformance::{ + checkpointer_concurrent_contract, checkpointer_contract, checkpointer_lineage_contract, + checkpointer_writes_contract, +}; +use tinystoragedrivers_core::{MemoryStorage, Scope, StorageBackend}; + +fn docs(storage: &MemoryStorage, scope: &str) -> Arc { + Arc::clone( + storage + .for_scope(&Scope::new(scope).unwrap()) + .unwrap() + .documents(), + ) +} + +fn sample(thread: &str, id: &str, _parent: Option<&str>, step: i32) -> Checkpoint { + Checkpoint::new(step, Vec::new()) + .with_thread_id(thread.to_string()) + .with_checkpoint_id(id.to_string()) +} + +fn checkpointer() -> DriverCheckpointer { + DriverCheckpointer::new(docs(&MemoryStorage::new(), "local")) +} + +#[tokio::test] +async fn passes_the_checkpointer_contract() { + checkpointer_contract(checkpointer()).await; +} + +#[tokio::test] +async fn passes_the_writes_contract() { + checkpointer_writes_contract(checkpointer()).await; +} + +#[tokio::test] +async fn passes_the_lineage_contract() { + checkpointer_lineage_contract(checkpointer()).await; +} + +#[tokio::test] +async fn passes_the_concurrent_contract() { + checkpointer_concurrent_contract(Arc::new(checkpointer())).await; +} + +#[tokio::test] +async fn scopes_and_prefixes_keep_threads_apart() { + let storage = MemoryStorage::new(); + let alice = DriverCheckpointer::::new(docs(&storage, "alice")); + let bob = DriverCheckpointer::::new(docs(&storage, "bob")); + let other = DriverCheckpointer::::with_prefix(docs(&storage, "alice"), "other"); + let checkpoint = sample("t", "c1", None, 1); + alice.put(checkpoint).await.unwrap(); + assert!(bob.get("t", None).await.unwrap().is_none()); + assert!(other.get("t", None).await.unwrap().is_none()); + assert_eq!(alice.list_threads().await.unwrap(), vec!["t".to_owned()]); + assert!(bob.list_threads().await.unwrap().is_empty()); + assert!(format!("{alice:?}").contains("graph_checkpoints")); + let _ = alice.clone(); +} + +#[tokio::test] +async fn leases_follow_the_claim_protocol() { + let cp = checkpointer(); + let minute = Duration::from_secs(60); + assert!(cp.try_claim("t", "a", minute).await.unwrap()); + assert!( + !cp.try_claim("t", "b", minute).await.unwrap(), + "live lease is refused" + ); + assert!( + cp.try_claim("t", "a", minute).await.unwrap(), + "same owner re-claims" + ); + assert!(cp.renew("t", "a", minute).await.unwrap()); + assert!( + !cp.renew("t", "b", minute).await.unwrap(), + "only the owner renews" + ); + assert!(!cp.renew("missing", "a", minute).await.unwrap()); + + cp.release("t", "b").await.unwrap(); + assert!( + !cp.try_claim("t", "b", minute).await.unwrap(), + "foreign release is a no-op" + ); + cp.release("t", "a").await.unwrap(); + cp.release("t", "a").await.unwrap(); + assert!( + cp.try_claim("t", "b", minute).await.unwrap(), + "released lease is free" + ); + + assert!(cp.try_claim("z", "dead", Duration::ZERO).await.unwrap()); + tokio::time::sleep(Duration::from_millis(5)).await; + assert!( + !cp.renew("z", "dead", minute).await.unwrap(), + "expired lease cannot renew" + ); + assert!( + cp.try_claim("z", "new", minute).await.unwrap(), + "expired lease is reclaimable" + ); +} + +#[test] +fn long_keys_are_hashed_and_tuples_never_collide() { + assert_ne!(key(&["a/b", "c"]), key(&["a", "b/c"])); + let long = "t".repeat(500); + let hashed = key(&[&long]); + assert!(hashed.starts_with("h:") && hashed.len() == 66, "{hashed}"); + assert_ne!(hashed, key(&[&"u".repeat(500)])); + assert_ne!( + key(&[&"t".repeat(500)]), + key(&[&"t".repeat(499), "t"]), + "the hash covers the length-prefixed tuple" + ); +} + +#[test] +fn namespaces_encode_injectively() { + let ns = |parts: &[&str]| parts.iter().map(|p| (*p).to_string()).collect::>(); + assert_ne!( + namespace_key(&ns(&["a", "b"])), + namespace_key(&ns(&["a\u{1f}b"])) + ); + assert_ne!( + namespace_key(&ns(&["a", "b"])), + namespace_key(&ns(&["a/b"])) + ); + assert_ne!(namespace_key(&ns(&["ab"])), namespace_key(&ns(&["a", "b"]))); + assert_eq!(namespace_key(&[]), ""); +} + +#[tokio::test] +async fn scoped_reads_stay_in_their_namespace_when_ids_repeat() { + let saver = checkpointer(); + let child = vec!["sub".to_string()]; + saver.put(sample("t", "same", None, 1)).await.unwrap(); + saver + .put(sample("t", "same", None, 2).with_namespace(child.clone())) + .await + .unwrap(); + let root = saver + .get_scoped("t", Some("same"), &[]) + .await + .unwrap() + .unwrap(); + assert_eq!( + root.state, 1, + "the root never loads the subgraph's checkpoint" + ); + let nested = saver + .get_scoped("t", Some("same"), &child) + .await + .unwrap() + .unwrap(); + assert_eq!(nested.state, 2); + assert_eq!( + saver + .get_scoped("t", None, &[]) + .await + .unwrap() + .unwrap() + .state, + 1 + ); + assert!( + saver + .get_scoped("t", None, &["other".to_string()]) + .await + .unwrap() + .is_none() + ); +} + +#[tokio::test] +async fn pending_writes_never_merge_across_lookalike_namespaces() { + let saver = checkpointer(); + let config = |namespace: Vec| CheckpointConfig { + thread_id: "t".to_string(), + checkpoint_id: Some("c".to_string()), + namespace, + }; + let split = config(vec!["a".to_string(), "b".to_string()]); + let joined = config(vec!["a\u{1f}b".to_string()]); + saver + .put_writes( + &split, + &[PendingWrite::data("n", "task", 0, "out", json!("split"))], + ) + .await + .unwrap(); + assert!(saver.get_writes(&joined).await.unwrap().is_empty()); + assert_eq!(saver.get_writes(&split).await.unwrap().len(), 1); +} + +#[test] +fn driver_failures_are_checkpoint_errors() { + let error = map_error(StorageError::unavailable("busy")); + assert!(matches!(error, TinyAgentsError::Checkpoint(ref m) if m.contains("busy"))); +} + +#[tokio::test] +async fn a_corrupt_record_is_a_checkpoint_error() { + let storage = MemoryStorage::new(); + let docs = docs(&storage, "local"); + let cp = DriverCheckpointer::::new(Arc::clone(&docs)); + cp.put(sample("t", "c1", None, 1)).await.unwrap(); + let stored = cp.thread_docs("t").await.unwrap().remove(0); + docs.put( + "graph_checkpoints", + &stored.id, + json!({"thread": "t", "seq": 0, "checkpoint_id": "c1", "record": "nope"}), + Precondition::None, + ) + .await + .unwrap(); + let error = cp.get("t", None).await.unwrap_err(); + assert!(matches!(error, TinyAgentsError::Checkpoint(_)), "{error:?}"); +} diff --git a/crates/tinyagents-graph/src/checkpoint/mod.rs b/crates/tinyagents-graph/src/checkpoint/mod.rs index ddd201b9d..87463a386 100644 --- a/crates/tinyagents-graph/src/checkpoint/mod.rs +++ b/crates/tinyagents-graph/src/checkpoint/mod.rs @@ -13,11 +13,15 @@ //! at superstep boundaries only — never mid-node — so resuming always reruns a //! node from its start. +#[cfg(feature = "storage-drivers")] +mod drivers; mod file; #[cfg(feature = "sqlite")] mod sqlite; mod types; +#[cfg(feature = "storage-drivers")] +pub use drivers::DriverCheckpointer; pub use file::FileCheckpointer; #[cfg(feature = "sqlite")] pub use sqlite::SqliteCheckpointer; diff --git a/crates/tinyagents-graph/src/lib.rs b/crates/tinyagents-graph/src/lib.rs index 3f836b69d..bf04be749 100644 --- a/crates/tinyagents-graph/src/lib.rs +++ b/crates/tinyagents-graph/src/lib.rs @@ -58,6 +58,8 @@ pub use channel::{ Barrier, BinaryAggregate, Channel, ChannelSet, ChannelState, ChannelUpdate, ChannelWrite, Delta, Ephemeral, LastValue, Messages, NamedBarrier, ReducerRegistry, Topic, Untracked, }; +#[cfg(feature = "storage-drivers")] +pub use checkpoint::DriverCheckpointer; #[cfg(feature = "sqlite")] pub use checkpoint::SqliteCheckpointer; pub use checkpoint::{ diff --git a/crates/tinyagents-harness/Cargo.toml b/crates/tinyagents-harness/Cargo.toml index 29dbc1e26..6accef0e3 100644 --- a/crates/tinyagents-harness/Cargo.toml +++ b/crates/tinyagents-harness/Cargo.toml @@ -40,6 +40,9 @@ tinyinference-image = { path = "../../vendor/tinyinference/crates/tinyinference- tinyinference-video = { path = "../../vendor/tinyinference/crates/tinyinference-video", version = "0.3.0", optional = true } tinytools = { path = "../../vendor/tinytools/crates/tinytools", version = "0.5.0" } tokio = { workspace = true, features = ["sync", "time", "macros", "rt", "rt-multi-thread", "fs", "io-util", "process"] } +# Storage ports (no database client): the `storage-drivers` adapters put +# the harness `Store` / `AppendStore` on any tinystoragedrivers backend. +tinystoragedrivers-core = { workspace = true, optional = true } tempfile = { workspace = true } wait-timeout = { version = "0.2", optional = true } uuid = { workspace = true } @@ -47,6 +50,10 @@ uuid = { workspace = true } [features] default = ["claude-code", "langfuse"] sqlite = ["dep:rusqlite"] +# `store::DriverStore` / `store::DriverAppendStore` over tinystoragedrivers +# ports, so a host's chosen backend (SQLite, MongoDB, memory) holds harness +# data. +storage-drivers = ["dep:tinystoragedrivers-core"] # `tools` is a deprecated alias kept so downstream feature forwards that # still spell out the old name keep compiling. tools = ["builtin-tools"] diff --git a/crates/tinyagents-harness/src/store/README.md b/crates/tinyagents-harness/src/store/README.md index 93ff45322..831a6add6 100644 --- a/crates/tinyagents-harness/src/store/README.md +++ b/crates/tinyagents-harness/src/store/README.md @@ -75,3 +75,26 @@ The richer, TTL-aware, batch-oriented sibling trait — see its own README. writers. - None of the in-memory backends are durable: data is lost when the value is dropped. + +## Storage-driver backends (`storage-drivers` feature) + +`DriverStore` and `DriverAppendStore` adapt the harness `Store` and +`AppendStore` traits onto [tinystoragedrivers](https://github.com/tinyhumansai/tinystoragedrivers) +ports. A host opens one storage backend (SQLite on desktop, MongoDB in the cloud, +memory in tests), takes a scoped `DocumentStore` / `StreamStore` for the tenant, +and wraps it: + +```rust,ignore +let scoped = backend.for_scope(&Scope::new(agent_id)?)?; +let kv = DriverStore::new(scoped.documents().clone()); +let journal = DriverAppendStore::with_prefix(scoped.streams().clone(), "journal/"); +``` + +- `DriverStore` keeps every namespace in one collection (`harness_store` by + default) as `{ns, key, value}` documents, so namespaces and keys may hold any + characters, and `list` is an indexed query. +- `DriverAppendStore` maps each stream to a driver stream (optionally + prefixed). Driver offsets are already dense and zero-based, matching the + `AppendStore` contract. +- Driver `InvalidInput` errors surface as `TinyAgentsError::Validation`. Every + other driver error surfaces as `TinyAgentsError::Storage`. diff --git a/crates/tinyagents-harness/src/store/drivers.rs b/crates/tinyagents-harness/src/store/drivers.rs new file mode 100644 index 000000000..252e5c3dc --- /dev/null +++ b/crates/tinyagents-harness/src/store/drivers.rs @@ -0,0 +1,241 @@ +//! [`Store`] and [`AppendStore`] over `tinystoragedrivers` ports. +//! +//! A host that has opened a storage backend (SQLite on a desktop, MongoDB in +//! the cloud, memory in tests) hands the harness a scoped +//! [`DocumentStore`] / [`StreamStore`] and wraps it here; the harness keeps +//! talking to its own narrow traits and never learns which backend it runs +//! on. The handles are already bound to a tenant scope, so nothing here takes +//! one. +//! +//! # Layout +//! +//! - [`DriverStore`] keeps every harness namespace in one document collection +//! (default [`DriverStore::DEFAULT_COLLECTION`]). Each entry is a document +//! `{ "ns": , "key": , "value": }`, so a namespace +//! or key may hold any characters and `list` is an indexed query on `ns`. +//! - [`DriverAppendStore`] maps each harness stream to the driver stream +//! `:`, the prefix empty when there is none. The +//! length is always written, so no two `(prefix, stream)` pairs, prefixed +//! or not, ever address one stream. Offsets are the +//! driver's dense offsets, which already match the [`AppendStore`] +//! contract. + +use std::sync::Arc; + +use async_trait::async_trait; +use serde_json::{Value, json}; +use tinystoragedrivers_core::{ + CollectionSpec, DocumentStore, DocumentStoreExt, ErrorKind, Filter, IndexSpec, Precondition, + Query, StorageError, StreamStore, +}; +use tokio::sync::OnceCell; + +use super::{AppendStore, Store}; +use crate::error::{Result, TinyAgentsError}; + +/// Map a storage driver failure onto the harness error type. Invalid input +/// (a name the backend cannot store) is a validation error; everything else +/// is a storage error carrying the driver's message. +pub(crate) fn map_error(error: StorageError) -> TinyAgentsError { + match error.kind() { + ErrorKind::InvalidInput => TinyAgentsError::Validation(error.to_string()), + _ => TinyAgentsError::Storage(error.to_string()), + } +} + +/// A harness [`Store`] over a driver [`DocumentStore`]. +#[derive(Debug, Clone)] +pub struct DriverStore { + docs: Arc, + collection: String, + declared: Arc>, +} + +impl DriverStore { + /// The collection used by [`DriverStore::new`]. + pub const DEFAULT_COLLECTION: &'static str = "harness_store"; + + /// Store harness namespaces in [`Self::DEFAULT_COLLECTION`]. + pub fn new(docs: Arc) -> Self { + Self::with_collection(docs, Self::DEFAULT_COLLECTION) + } + + /// Store harness namespaces in `collection`, so several independent + /// harness stores can share one backend. + pub fn with_collection(docs: Arc, collection: impl Into) -> Self { + Self { + docs, + collection: collection.into(), + declared: Arc::new(OnceCell::new()), + } + } + + /// The document id of `key` in `namespace`. The namespace is + /// length-prefixed so `("a/b", "c")` and `("a", "b/c")` never collide. + fn id(namespace: &str, key: &str) -> String { + format!("{}:{namespace}/{key}", namespace.len()) + } + + /// Declare the collection (with its `ns` index) once per handle. + async fn declared(&self) -> Result<()> { + self.declared + .get_or_try_init(|| async { + let spec = + CollectionSpec::new(&self.collection).index(IndexSpec::new("by_ns", ["ns"])); + self.docs.ensure_collection(&spec).await.map_err(map_error) + }) + .await + .map(|_| ()) + } +} + +#[async_trait] +impl Store for DriverStore { + async fn get(&self, namespace: &str, key: &str) -> Result> { + self.declared().await?; + let found = self + .docs + .get(&self.collection, &Self::id(namespace, key)) + .await + .map_err(map_error)?; + Ok(found.and_then(|mut stored| stored.doc.get_mut("value").map(Value::take))) + } + + async fn put(&self, namespace: &str, key: &str, value: Value) -> Result<()> { + self.declared().await?; + let doc = json!({ "ns": namespace, "key": key, "value": value }); + self.docs + .put( + &self.collection, + &Self::id(namespace, key), + doc, + Precondition::None, + ) + .await + .map(|_| ()) + .map_err(map_error) + } + + async fn delete(&self, namespace: &str, key: &str) -> Result<()> { + self.declared().await?; + self.docs + .delete( + &self.collection, + &Self::id(namespace, key), + Precondition::None, + ) + .await + .map(|_| ()) + .map_err(map_error) + } + + async fn list(&self, namespace: &str) -> Result> { + self.declared().await?; + let query = Query::filter(Filter::eq("ns", namespace)); + let found = self + .docs + .query_all(&self.collection, &query) + .await + .map_err(map_error)?; + Ok(found + .into_iter() + .filter_map(|stored| { + stored + .doc + .get("key") + .and_then(Value::as_str) + .map(str::to_owned) + }) + .collect()) + } +} + +/// A harness [`AppendStore`] over a driver [`StreamStore`]. +#[derive(Debug, Clone)] +pub struct DriverAppendStore { + streams: Arc, + prefix: String, +} + +impl DriverAppendStore { + /// Map each harness stream `s` to the driver stream `0:s`. + pub fn new(streams: Arc) -> Self { + Self::with_prefix(streams, "") + } + + /// Map each harness stream to `:`, so several + /// independent append stores can share one backend. The prefix length is + /// part of the name, so `("a", "bc")` and `("ab", "c")` stay apart. + pub fn with_prefix(streams: Arc, prefix: impl Into) -> Self { + Self { + streams, + prefix: prefix.into(), + } + } + + fn name(&self, stream: &str) -> String { + format!("{}:{}{stream}", self.prefix.len(), self.prefix) + } +} + +/// Entries read per driver call when following a stream to its end. +const READ_PAGE: usize = 1_000; + +#[async_trait] +impl AppendStore for DriverAppendStore { + async fn append(&self, stream: &str, value: Value) -> Result { + self.streams + .append(&self.name(stream), value) + .await + .map_err(map_error) + } + + async fn read_from(&self, stream: &str, offset: u64) -> Result> { + let name = self.name(stream); + let mut out = Vec::new(); + let mut next = offset; + loop { + let page = self + .streams + .read_window(&name, next, READ_PAGE) + .await + .map_err(map_error)?; + let full = page.len() == READ_PAGE; + for entry in page { + next = entry.offset + 1; + out.push((entry.offset, entry.value)); + } + if !full { + return Ok(out); + } + } + } + + async fn read_window( + &self, + stream: &str, + offset: u64, + limit: usize, + ) -> Result> { + let page = self + .streams + .read_window(&self.name(stream), offset, limit) + .await + .map_err(map_error)?; + Ok(page + .into_iter() + .map(|entry| (entry.offset, entry.value)) + .collect()) + } + + async fn len(&self, stream: &str) -> Result { + self.streams + .len(&self.name(stream)) + .await + .map_err(map_error) + } +} + +#[cfg(test)] +#[path = "drivers_tests.rs"] +mod tests; diff --git a/crates/tinyagents-harness/src/store/drivers_tests.rs b/crates/tinyagents-harness/src/store/drivers_tests.rs new file mode 100644 index 000000000..d044d4ce8 --- /dev/null +++ b/crates/tinyagents-harness/src/store/drivers_tests.rs @@ -0,0 +1,127 @@ +use super::*; +use crate::store::conformance::run_store_conformance; +use tinystoragedrivers_core::{MemoryStorage, Scope, StorageBackend}; + +fn scoped(storage: &MemoryStorage, scope: &str) -> tinystoragedrivers_core::ScopedStorage { + storage.for_scope(&Scope::new(scope).unwrap()).unwrap() +} + +#[tokio::test] +async fn driver_store_passes_the_store_conformance_suite() { + let storage = MemoryStorage::new(); + let store = DriverStore::new(Arc::clone(scoped(&storage, "local").documents())); + run_store_conformance(&store).await; +} + +#[tokio::test] +async fn namespaces_and_keys_may_hold_any_characters() { + let storage = MemoryStorage::new(); + let store = DriverStore::new(Arc::clone(scoped(&storage, "local").documents())); + store.put("a/b", "c", json!(1)).await.unwrap(); + store.put("a", "b/c", json!(2)).await.unwrap(); + store.put("odd ns:é", "k y", json!(3)).await.unwrap(); + assert_eq!(store.get("a/b", "c").await.unwrap(), Some(json!(1))); + assert_eq!(store.get("a", "b/c").await.unwrap(), Some(json!(2))); + assert_eq!(store.get("odd ns:é", "k y").await.unwrap(), Some(json!(3))); + assert_eq!(store.list("a").await.unwrap(), vec!["b/c".to_owned()]); +} + +#[tokio::test] +async fn stores_in_different_scopes_and_collections_are_isolated() { + let storage = MemoryStorage::new(); + let alice = DriverStore::new(Arc::clone(scoped(&storage, "alice").documents())); + let bob = DriverStore::new(Arc::clone(scoped(&storage, "bob").documents())); + let other = DriverStore::with_collection( + Arc::clone(scoped(&storage, "alice").documents()), + "other_store", + ); + alice.put("ns", "k", json!("a")).await.unwrap(); + assert_eq!(bob.get("ns", "k").await.unwrap(), None); + assert_eq!(other.get("ns", "k").await.unwrap(), None); + assert!(bob.list("ns").await.unwrap().is_empty()); +} + +#[tokio::test] +async fn a_list_follows_every_page() { + let storage = MemoryStorage::new(); + let store = DriverStore::new(Arc::clone(scoped(&storage, "local").documents())); + for i in 0..1_005 { + store + .put("big", &format!("k{i:04}"), json!(i)) + .await + .unwrap(); + } + assert_eq!(store.list("big").await.unwrap().len(), 1_005); +} + +#[tokio::test] +async fn an_unstorable_key_is_a_validation_error() { + let storage = MemoryStorage::new(); + let store = DriverStore::new(Arc::clone(scoped(&storage, "local").documents())); + let long = "k".repeat(600); + let error = store.put("ns", &long, json!(1)).await.unwrap_err(); + assert!(matches!(error, TinyAgentsError::Validation(_)), "{error:?}"); + let error = DriverStore::with_collection( + Arc::clone(scoped(&storage, "local").documents()), + "bad name", + ) + .get("ns", "k") + .await + .unwrap_err(); + assert!(matches!(error, TinyAgentsError::Validation(_)), "{error:?}"); +} + +#[test] +fn backend_failures_map_to_storage_errors() { + let error = map_error(StorageError::unavailable("busy")); + assert!(matches!(error, TinyAgentsError::Storage(ref m) if m.contains("busy"))); +} + +#[tokio::test] +async fn driver_append_store_keeps_dense_offsets() { + let storage = MemoryStorage::new(); + let journal = DriverAppendStore::new(Arc::clone(scoped(&storage, "local").streams())); + assert_eq!(journal.len("run-1").await.unwrap(), 0); + assert!(journal.read_from("run-1", 0).await.unwrap().is_empty()); + for i in 0..3 { + assert_eq!(journal.append("run-1", json!({ "i": i })).await.unwrap(), i); + } + assert_eq!(journal.len("run-1").await.unwrap(), 3); + assert_eq!( + journal.read_from("run-1", 1).await.unwrap(), + vec![(1, json!({ "i": 1 })), (2, json!({ "i": 2 }))] + ); + assert_eq!( + journal.read_window("run-1", 0, 2).await.unwrap(), + vec![(0, json!({ "i": 0 })), (1, json!({ "i": 1 }))] + ); + assert!(journal.read_from("run-1", 9).await.unwrap().is_empty()); +} + +#[tokio::test] +async fn read_from_follows_long_streams_and_prefixes_isolate() { + let storage = MemoryStorage::new(); + let streams = Arc::clone(scoped(&storage, "local").streams()); + let journal = DriverAppendStore::with_prefix(Arc::clone(&streams), "journal/"); + let other = DriverAppendStore::with_prefix(Arc::clone(&streams), "other/"); + for i in 0..(READ_PAGE as u64 + 5) { + journal.append("s", json!(i)).await.unwrap(); + } + let all = journal.read_from("s", 0).await.unwrap(); + assert_eq!(all.len(), READ_PAGE + 5); + assert_eq!(all.last().unwrap().0, READ_PAGE as u64 + 4); + assert_eq!(other.len("s").await.unwrap(), 0); + let error = journal.append("", json!(1)).await; + assert!(error.is_ok(), "a prefix makes the empty stream name valid"); + let ab = DriverAppendStore::with_prefix(Arc::clone(&streams), "ab"); + let a = DriverAppendStore::with_prefix(Arc::clone(&streams), "a"); + ab.append("c", json!("ab/c")).await.unwrap(); + assert_eq!(a.len("bc").await.unwrap(), 0, "prefixes never overlap"); + let bare = DriverAppendStore::new(Arc::clone(scoped(&storage, "local").streams())); + bare.append("", json!("bare")).await.unwrap(); + assert_eq!(bare.len("").await.unwrap(), 1, "a bare name is encoded too"); + // A bare stream spelled like a prefixed one stays its own stream. + bare.append("1:ab", json!("bare")).await.unwrap(); + assert_eq!(a.len("b").await.unwrap(), 0); + assert_eq!(a.len("").await.unwrap(), 0); +} diff --git a/crates/tinyagents-harness/src/store/mod.rs b/crates/tinyagents-harness/src/store/mod.rs index 3c21fb3de..f0fb005d6 100644 --- a/crates/tinyagents-harness/src/store/mod.rs +++ b/crates/tinyagents-harness/src/store/mod.rs @@ -31,6 +31,8 @@ //! consistent names make multi-store applications easier to audit. pub mod conformance; +#[cfg(feature = "storage-drivers")] +mod drivers; pub mod namespaced; mod types; @@ -43,6 +45,9 @@ use serde_json::Value; pub use types::*; +#[cfg(feature = "storage-drivers")] +pub use drivers::{DriverAppendStore, DriverStore}; + use crate::error::{Result, TinyAgentsError}; use crate::ids::now_ms; diff --git a/crates/tinyagents-session/Cargo.toml b/crates/tinyagents-session/Cargo.toml index de3587e74..eff5b34d9 100644 --- a/crates/tinyagents-session/Cargo.toml +++ b/crates/tinyagents-session/Cargo.toml @@ -28,6 +28,8 @@ tracing = { workspace = true } tokio = { workspace = true, default-features = true, features = ["sync"] } # `threads` names its atomic-rewrite temp files `.conversations-.tmp`. uuid = { workspace = true } +# Session stores over `tinystoragedrivers` ports (`storage-drivers` feature). +tinystoragedrivers-core = { workspace = true, optional = true, features = ["blocking"] } [features] default = [] @@ -35,8 +37,17 @@ default = [] # dependency above). This feature is retained as a no-op so downstream # feature forwards keep compiling. tracing = ["tinyagents-harness/tracing"] +# A `SessionStoreProvider` over `tinystoragedrivers` ports, so a host can keep +# every agent's transcripts, turn states, records and journal in whichever +# backend it opened (SQLite, MongoDB, memory, ...). +storage-drivers = [ + "dep:tinystoragedrivers-core", + "tinyagents-harness/storage-drivers", +] [dev-dependencies] +# The `storage-drivers` provider is checked against a real SQLite backend too. +tinystoragedrivers-sqlite = { path = "../../vendor/tinystoragedrivers/crates/tinystoragedrivers-sqlite" } tokio = { workspace = true, default-features = true, features = ["macros", "rt-multi-thread"] } [lints] diff --git a/crates/tinyagents-session/src/README.md b/crates/tinyagents-session/src/README.md index f76dad831..846c414e1 100644 --- a/crates/tinyagents-session/src/README.md +++ b/crates/tinyagents-session/src/README.md @@ -68,6 +68,10 @@ reachable under `session::`, `session::run_ledger::`, and - **Chat threads** — `ConversationStore` and its wire types (re-exported at the root), plus the free functions and channel subscriber under `session::threads` — see its own [README](./threads/README.md) +- **Store port** — `SessionStoreProvider`, `AgentStores`, `TurnStates`, + `InMemorySessionStores`, and with feature `storage-drivers` + `DriverSessionStores` over a `tinystoragedrivers` backend — see + [`docs/modules/session/store-port.md`](../../../docs/modules/session/store-port.md) - **Connections** — `with_connection` (autocommit) and `with_transaction` (`BEGIN IMMEDIATE`) - **Testkit** — `testkit::conformance::run_ledger_conformance` and diff --git a/crates/tinyagents-session/src/lib.rs b/crates/tinyagents-session/src/lib.rs index 07c094eee..09114b5e8 100644 --- a/crates/tinyagents-session/src/lib.rs +++ b/crates/tinyagents-session/src/lib.rs @@ -112,6 +112,10 @@ pub use port::{ AgentStores, InMemorySessionStores, InMemoryTranscriptLocator, InMemoryTurnStates, SessionStoreProvider, TurnStates, }; +#[cfg(feature = "storage-drivers")] +pub use port::{ + DriverSessionStores, DriverTranscriptHistory, DriverTranscriptLocator, DriverTurnStates, +}; pub use retention::{ RetentionReport, apply_retention, prune_run_events_before, prune_run_telemetry_before, prune_sessions_before, prune_tool_calls_before, reindex_fts, trim_session_messages, diff --git a/crates/tinyagents-session/src/port/drivers/mod.rs b/crates/tinyagents-session/src/port/drivers/mod.rs new file mode 100644 index 000000000..abcc75821 --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/mod.rs @@ -0,0 +1,277 @@ +//! A [`SessionStoreProvider`] over `tinystoragedrivers` ports. +//! +//! A host opens one storage backend (SQLite on a desktop, MongoDB in the +//! cloud, memory in tests) and hands it to [`DriverSessionStores`]. Each +//! agent's stores then live in that backend under the agent's own +//! [`Scope`], so one shared database serves every agent and the driver, not +//! this crate, keeps them apart: a handle bound to one scope cannot name +//! another scope's records. +//! +//! # What goes where +//! +//! | Store | Backend shape | +//! | --- | --- | +//! | transcripts | `session_transcripts` (one index document per stem) and `session_transcript_entries` (an append-only log per stem) | +//! | turn states | `session_turn_states`, one document per turn | +//! | key-value | `session_kv`, through the harness [`DriverStore`] | +//! | journal | streams `16:session_journal/`, through the harness [`DriverAppendStore`] | +//! +//! # Sync seams +//! +//! The transcript and turn-state traits are synchronous, so their +//! implementations here run each driver call on a [`Blocking`] bridge: a +//! dedicated runtime thread, which works from any caller (inside a tokio +//! runtime or not) and never needs `block_in_place`. +//! +//! # Scopes +//! +//! An agent id that is a valid [`Scope`] is used as is. Any other id (empty, +//! over-long, holding whitespace) and any id that itself starts with +//! `sha256:` maps to `sha256:` of the id, so every agent gets a scope of +//! its own, no raw id can name a hashed scope, and the mapping never changes. + +mod refused; +mod transcripts; +mod turn_states; + +use std::collections::{HashMap, HashSet}; +use std::sync::{Arc, Mutex, PoisonError}; + +use sha2::{Digest, Sha256}; +use tinyagents_harness::store::{DriverAppendStore, DriverStore}; +use tinystoragedrivers_core::{Blocking, Scope, ScopedStorage, StorageBackend, StorageError}; + +use super::{AgentStores, SessionStoreProvider}; + +pub use transcripts::{DriverTranscriptHistory, DriverTranscriptLocator}; +pub use turn_states::DriverTurnStates; + +/// Prefix of the scopes agent ids are hashed into. Reserved: an agent id +/// starting with it is hashed too, so it cannot name another agent's scope. +const HASHED_SCOPE: &str = "sha256:"; + +/// Collection holding each agent's key-value records. +const KV_COLLECTION: &str = "session_kv"; +/// Prefix of each agent's journal streams. +const JOURNAL_PREFIX: &str = "session_journal/"; + +/// A [`SessionStoreProvider`] keeping every agent's stores in one +/// `tinystoragedrivers` backend, each agent in its own [`Scope`]. +pub struct DriverSessionStores { + backend: Arc, + bridge: Blocking, + agents: Mutex>, + recover_on_open: bool, + recovered: Mutex>, + /// Names this provider in destination keys and handle paths. + id: uuid::Uuid, +} + +impl std::fmt::Debug for DriverSessionStores { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let agents = self.agents.lock().unwrap_or_else(PoisonError::into_inner); + f.debug_struct("DriverSessionStores") + .field("driver", &self.backend.driver()) + .field("agents", &agents.len()) + .field("recover_on_open", &self.recover_on_open) + .finish_non_exhaustive() + } +} + +impl DriverSessionStores { + /// A provider over `backend`, with its own [`Blocking`] bridge. + /// + /// # Errors + /// + /// When the bridge's runtime thread cannot start. + pub fn new(backend: Arc) -> Result { + Ok(Self::with_bridge(backend, Blocking::new()?)) + } + + /// A provider over `backend` that runs its synchronous seams on `bridge`, + /// so several providers (or other sync adapters) can share one thread. + pub fn with_bridge(backend: Arc, bridge: Blocking) -> Self { + Self { + backend, + bridge, + agents: Mutex::new(HashMap::new()), + recover_on_open: false, + recovered: Mutex::new(HashSet::new()), + id: uuid::Uuid::new_v4(), + } + } + + /// Marks an agent's in-flight turns interrupted the first time this + /// provider opens its stores. + /// + /// For a single-process host (the desktop app): there, a turn still in + /// flight when the process starts can only be one an earlier process + /// left behind. A host that runs several processes against one database + /// must leave this off, since another process may own those turns. + #[must_use] + pub fn recover_on_open(mut self, enabled: bool) -> Self { + self.recover_on_open = enabled; + self + } + + /// The scope agent `agent_id` is stored under. + pub fn scope_for(agent_id: &str) -> Scope { + if !agent_id.starts_with(HASHED_SCOPE) + && let Ok(scope) = Scope::new(agent_id) + { + return scope; + } + let digest = Sha256::digest(agent_id.as_bytes()); + Scope::new(format!("{HASHED_SCOPE}{}", hex::encode(digest))) + .expect("a sha256 hex scope is always valid") + } + + /// The stores of `agent_id`, or the error that kept the backend from + /// binding its scope. + /// + /// # Errors + /// + /// Whatever [`StorageBackend::for_scope`] reports. + pub fn try_for_agent(&self, agent_id: &str) -> Result { + if let Some(stores) = self + .agents + .lock() + .unwrap_or_else(PoisonError::into_inner) + .get(agent_id) + { + return Ok(stores.clone()); + } + let scoped = self.backend.for_scope(&Self::scope_for(agent_id))?; + let stores = self.build(&scoped); + if self.recover_on_open { + // Fail closed: stores whose recovery did not run must not start + // turns a retried sweep would then mistake for crash residue. + // Nothing is cached, so the next open tries again. + self.recover_agent(agent_id, &stores)?; + } + Ok(self + .agents + .lock() + .unwrap_or_else(PoisonError::into_inner) + .entry(agent_id.to_string()) + .or_insert(stores) + .clone()) + } + + /// Interrupts `agent_id`'s in-flight turns unless this provider already + /// did. + /// + /// The `recovered` lock is held across the sweep, so a concurrent first + /// open of the same agent waits for it instead of handing out stores a + /// still-running sweep could interrupt a new turn on. The agent is + /// recorded only once the sweep succeeds. + fn recover_agent(&self, agent_id: &str, stores: &AgentStores) -> Result<(), StorageError> { + let mut recovered = self + .recovered + .lock() + .unwrap_or_else(PoisonError::into_inner); + if recovered.contains(agent_id) { + return Ok(()); + } + let now = chrono::Utc::now().to_rfc3339(); + match stores.turn_states.mark_all_interrupted(&now) { + Ok(count) => { + if count > 0 { + tracing::info!( + target: "tinyagents_session::port::drivers", + agent_id, + count, + "[session-store] marked turns left in flight interrupted" + ); + } + recovered.insert(agent_id.to_string()); + Ok(()) + } + Err(error) => { + tracing::warn!( + target: "tinyagents_session::port::drivers", + agent_id, + %error, + "[session-store] could not recover in-flight turns; retried on the next open" + ); + Err(StorageError::unavailable(format!( + "recovering in-flight turns failed: {error}" + ))) + } + } + } + + fn build(&self, scoped: &ScopedStorage) -> AgentStores { + let docs = Arc::clone(scoped.documents()); + // The provider's own id keeps two providers with the same driver and + // scope from claiming one destination; unlike an address, it is + // never reused. + let label = format!("{}://{}/{}", scoped.driver(), self.id, scoped.scope()); + AgentStores { + transcripts: Arc::new(DriverTranscriptLocator::new( + Arc::clone(&docs), + self.bridge.clone(), + label, + )), + turn_states: Arc::new(DriverTurnStates::new( + Arc::clone(&docs), + self.bridge.clone(), + )), + kv: Arc::new(DriverStore::with_collection(docs, KV_COLLECTION)), + journal: Arc::new(DriverAppendStore::with_prefix( + Arc::clone(scoped.streams()), + JOURNAL_PREFIX, + )), + } + } +} + +impl SessionStoreProvider for DriverSessionStores { + /// The stores of `agent_id`. + /// + /// Fails closed: when the backend cannot bind the agent's scope, or + /// [`Self::recover_on_open`]'s sweep fails, the stores returned refuse + /// every call with that error rather than fall + /// back to somewhere the agent's data does not belong. They are not + /// cached, so the next call tries the backend again. + fn for_agent(&self, agent_id: &str) -> AgentStores { + self.try_for_agent(agent_id).unwrap_or_else(|error| { + tracing::error!( + target: "tinyagents_session::port::drivers", + agent_id, + %error, + "[session-store] backend refused the agent's scope" + ); + refused::stores(&error, &self.bridge) + }) + } + + /// Marks in-flight turns interrupted for every agent this provider has + /// opened. A backend cannot enumerate its scopes, so agents not yet + /// opened are recovered by [`Self::recover_on_open`] instead. + fn recover(&self) -> anyhow::Result<()> { + let now = chrono::Utc::now().to_rfc3339(); + let agents: Vec = self + .agents + .lock() + .unwrap_or_else(PoisonError::into_inner) + .values() + .cloned() + .collect(); + for stores in agents { + stores + .turn_states + .mark_all_interrupted(&now) + .map_err(anyhow::Error::msg)?; + } + Ok(()) + } + + fn destination_key(&self) -> Option { + Some(format!("{}://{}", self.backend.driver(), self.id)) + } +} + +#[cfg(test)] +#[path = "mod_tests.rs"] +mod tests; diff --git a/crates/tinyagents-session/src/port/drivers/mod_tests.rs b/crates/tinyagents-session/src/port/drivers/mod_tests.rs new file mode 100644 index 000000000..66ebeffd0 --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/mod_tests.rs @@ -0,0 +1,432 @@ +use super::*; +use crate::testkit::conformance::{session_store_conformance, session_store_isolation_conformance}; +use crate::turn_state::{TurnLifecycle, TurnState}; +use tinystoragedrivers_core::MemoryStorage; +use tinystoragedrivers_sqlite::SqliteStorage; + +fn memory() -> DriverSessionStores { + DriverSessionStores::new(Arc::new(MemoryStorage::new())).unwrap() +} + +fn sqlite(dir: &tempfile::TempDir) -> DriverSessionStores { + let storage = SqliteStorage::open(dir.path().join("sessions.db")).unwrap(); + DriverSessionStores::new(Arc::new(storage)).unwrap() +} + +#[tokio::test] +async fn memory_backed_stores_meet_the_session_store_contract() { + session_store_conformance(&memory()).await; +} + +#[tokio::test] +async fn memory_backed_stores_keep_agents_apart() { + session_store_isolation_conformance(&memory()).await; +} + +#[tokio::test] +async fn sqlite_backed_stores_meet_the_session_store_contract() { + let dir = tempfile::tempdir().unwrap(); + session_store_conformance(&sqlite(&dir)).await; +} + +#[tokio::test] +async fn sqlite_backed_stores_keep_agents_apart() { + let dir = tempfile::tempdir().unwrap(); + session_store_isolation_conformance(&sqlite(&dir)).await; +} + +#[tokio::test(flavor = "multi_thread")] +async fn stores_work_from_a_multi_threaded_runtime() { + session_store_conformance(&memory()).await; +} + +#[test] +fn stores_work_outside_any_runtime() { + let stores = memory().for_agent("plain"); + let turn = TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z"); + stores.turn_states.put(&turn).unwrap(); + assert_eq!(stores.turn_states.get("t").unwrap(), Some(turn)); +} + +#[test] +fn a_valid_agent_id_is_its_own_scope() { + assert_eq!( + DriverSessionStores::scope_for("agent-7").as_str(), + "agent-7" + ); +} + +#[test] +fn any_other_agent_id_maps_to_a_stable_hashed_scope() { + let spaced = DriverSessionStores::scope_for("has a space"); + assert!(spaced.as_str().starts_with("sha256:")); + assert_eq!(spaced, DriverSessionStores::scope_for("has a space")); + assert_ne!(spaced, DriverSessionStores::scope_for("has a space")); + assert!( + DriverSessionStores::scope_for("") + .as_str() + .starts_with("sha256:") + ); + let long = "a".repeat(300); + assert!( + DriverSessionStores::scope_for(&long) + .as_str() + .starts_with("sha256:") + ); +} + +#[test] +fn a_raw_id_can_never_name_a_hashed_scope() { + let hashed = DriverSessionStores::scope_for("has a space"); + let impostor = DriverSessionStores::scope_for(hashed.as_str()); + assert_ne!( + impostor, hashed, + "ids in the reserved prefix are hashed too" + ); + assert!(impostor.as_str().starts_with("sha256:")); +} + +#[test] +fn two_backends_never_share_a_destination() { + let one = memory().for_agent("a"); + let two = memory().for_agent("a"); + assert_ne!( + one.transcripts.destination_key(), + two.transcripts.destination_key() + ); +} + +#[test] +fn an_agent_gets_the_same_stores_every_time() { + let provider = memory(); + let first = provider.for_agent("a"); + let again = provider.for_agent("a"); + assert!(Arc::ptr_eq(&first.turn_states, &again.turn_states)); + assert_eq!( + first.transcripts.destination_key(), + again.transcripts.destination_key() + ); + assert_ne!( + first.transcripts.destination_key(), + provider.for_agent("b").transcripts.destination_key() + ); + let debug = format!("{provider:?}"); + assert!(debug.contains("agents: 2"), "{debug}"); + assert!(debug.contains("memory"), "{debug}"); + assert!( + provider + .destination_key() + .is_some_and(|key| key.starts_with("memory://")) + ); +} + +#[test] +fn data_survives_a_new_provider_on_the_same_backend() { + let backend: Arc = Arc::new(MemoryStorage::new()); + let turn = TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z"); + DriverSessionStores::new(Arc::clone(&backend)) + .unwrap() + .for_agent("a") + .turn_states + .put(&turn) + .unwrap(); + let reopened = DriverSessionStores::new(backend).unwrap().for_agent("a"); + assert_eq!(reopened.turn_states.get("t").unwrap(), Some(turn)); +} + +#[test] +fn recover_marks_open_agents_turns_interrupted() { + let provider = memory(); + let stores = provider.for_agent("a"); + stores + .turn_states + .put(&TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z")) + .unwrap(); + provider.recover().unwrap(); + assert_eq!( + stores.turn_states.get("t").unwrap().unwrap().lifecycle, + TurnLifecycle::Interrupted + ); +} + +#[test] +fn recover_on_open_interrupts_turns_an_earlier_process_left() { + let backend: Arc = Arc::new(MemoryStorage::new()); + let earlier = DriverSessionStores::new(Arc::clone(&backend)).unwrap(); + earlier + .for_agent("a") + .turn_states + .put(&TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z")) + .unwrap(); + + let plain = DriverSessionStores::new(Arc::clone(&backend)).unwrap(); + assert_eq!( + plain + .for_agent("a") + .turn_states + .get("t") + .unwrap() + .unwrap() + .lifecycle, + TurnLifecycle::Started, + "without the option nothing is touched" + ); + + let recovering = DriverSessionStores::new(backend) + .unwrap() + .recover_on_open(true); + assert!(format!("{recovering:?}").contains("recover_on_open: true")); + let stores = recovering.for_agent("a"); + assert_eq!( + stores.turn_states.get("t").unwrap().unwrap().lifecycle, + TurnLifecycle::Interrupted + ); + + // Only the first open recovers: a turn started afterwards stays live. + stores + .turn_states + .put(&TurnState::started("t2", "r", 8, "2026-01-01T00:01:00Z")) + .unwrap(); + recovering.agents.lock().unwrap().clear(); + assert_eq!( + recovering + .for_agent("a") + .turn_states + .get("t2") + .unwrap() + .unwrap() + .lifecycle, + TurnLifecycle::Started + ); +} + +/// A backend that cannot bind any scope. +#[derive(Debug)] +struct Unbindable; + +impl StorageBackend for Unbindable { + fn driver(&self) -> &'static str { + "unbindable" + } + + fn capabilities(&self) -> tinystoragedrivers_core::Capabilities { + tinystoragedrivers_core::Capabilities::default() + } + + fn for_scope(&self, _scope: &Scope) -> Result { + Err(StorageError::unavailable("database is down")) + } + + fn database(&self, _name: &str) -> Result, StorageError> { + Err(StorageError::unavailable("database is down")) + } +} + +#[tokio::test] +async fn an_unbindable_scope_fails_closed() { + let provider = DriverSessionStores::new(Arc::new(Unbindable)).unwrap(); + assert!(provider.try_for_agent("a").is_err()); + let stores = provider.for_agent("a"); + + let error = stores + .turn_states + .put(&TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z")) + .unwrap_err(); + assert!(error.contains("database is down"), "{error}"); + assert!(stores.turn_states.list().is_err()); + assert!(stores.kv.get("ns", "k").await.is_err()); + assert!( + stores + .journal + .append("s", serde_json::json!(1)) + .await + .is_err() + ); + assert!(stores.journal.len("s").await.is_err()); + + let session = crate::transcript::SessionRef::scoped("t", "a"); + assert!(!stores.transcripts.session_exists(&session)); + assert!(stores.transcripts.root_for_thread("t").is_none()); + assert!(stores.transcripts.latest_for_agent("a").is_none()); + let handle = stores + .transcripts + .open_stem("stem", transcripts::tests::meta("t")) + .unwrap(); + assert!(handle.read_session().is_err()); + assert!(handle.messages().is_err()); + assert!( + stores + .transcripts + .begin_generation(&session, transcripts::tests::meta("t")) + .is_err() + ); + + assert!( + provider.agents.lock().unwrap().is_empty(), + "refusals are not cached" + ); + provider.recover().unwrap(); +} + +/// A memory backend whose documents fail every query while `down` is set. +#[derive(Debug)] +struct Flaky { + storage: MemoryStorage, + down: Arc, +} + +#[derive(Debug)] +struct FlakyDocs { + inner: Arc, + down: Arc, +} + +impl FlakyDocs { + fn check(&self) -> Result<(), StorageError> { + if self.down.load(std::sync::atomic::Ordering::SeqCst) { + Err(StorageError::unavailable("flaky backend is down")) + } else { + Ok(()) + } + } +} + +#[tinystoragedrivers_core::async_trait] +impl tinystoragedrivers_core::DocumentStore for FlakyDocs { + fn capabilities(&self) -> tinystoragedrivers_core::Capabilities { + self.inner.capabilities() + } + async fn ensure_collection( + &self, + spec: &tinystoragedrivers_core::CollectionSpec, + ) -> Result<(), StorageError> { + self.inner.ensure_collection(spec).await + } + async fn get( + &self, + collection: &str, + id: &str, + ) -> Result>, StorageError> { + self.inner.get(collection, id).await + } + async fn put( + &self, + collection: &str, + id: &str, + doc: serde_json::Value, + precondition: tinystoragedrivers_core::Precondition, + ) -> Result { + self.inner.put(collection, id, doc, precondition).await + } + async fn delete( + &self, + collection: &str, + id: &str, + precondition: tinystoragedrivers_core::Precondition, + ) -> Result { + self.inner.delete(collection, id, precondition).await + } + async fn query( + &self, + collection: &str, + query: &tinystoragedrivers_core::Query, + ) -> Result< + tinystoragedrivers_core::Page>, + StorageError, + > { + self.check()?; + self.inner.query(collection, query).await + } + async fn count( + &self, + collection: &str, + filter: &tinystoragedrivers_core::Filter, + ) -> Result { + self.inner.count(collection, filter).await + } + async fn delete_where( + &self, + collection: &str, + filter: &tinystoragedrivers_core::Filter, + ) -> Result { + self.inner.delete_where(collection, filter).await + } + async fn claim( + &self, + collection: &str, + filter: &tinystoragedrivers_core::Filter, + sort: &[tinystoragedrivers_core::Sort], + patch: &serde_json::Value, + ) -> Result>, StorageError> { + self.inner.claim(collection, filter, sort, patch).await + } + async fn drop_collection(&self, collection: &str) -> Result<(), StorageError> { + self.inner.drop_collection(collection).await + } +} + +impl StorageBackend for Flaky { + fn driver(&self) -> &'static str { + "flaky" + } + fn capabilities(&self) -> tinystoragedrivers_core::Capabilities { + self.storage.capabilities() + } + fn for_scope(&self, scope: &Scope) -> Result { + let inner = self.storage.for_scope(scope)?; + Ok(ScopedStorage::new( + scope.clone(), + "flaky", + Arc::new(FlakyDocs { + inner: Arc::clone(inner.documents()), + down: Arc::clone(&self.down), + }), + Arc::clone(inner.streams()), + Arc::clone(inner.blobs()), + )) + } + fn database(&self, name: &str) -> Result, StorageError> { + self.storage.database(name) + } +} + +#[test] +fn a_failed_recovery_fails_closed_and_is_retried() { + let down = Arc::new(std::sync::atomic::AtomicBool::new(true)); + let provider = DriverSessionStores::new(Arc::new(Flaky { + storage: MemoryStorage::new(), + down: Arc::clone(&down), + })) + .unwrap() + .recover_on_open(true); + + let error = provider.try_for_agent("a").unwrap_err(); + assert!(error.message().contains("recovering"), "{error}"); + let refused = provider.for_agent("a"); + assert!( + refused + .turn_states + .put(&TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z")) + .is_err(), + "no turn can start before recovery ran" + ); + assert!(provider.agents.lock().unwrap().is_empty()); + + down.store(false, std::sync::atomic::Ordering::SeqCst); + let stores = provider.for_agent("a"); + stores + .turn_states + .put(&TurnState::started("t", "r", 8, "2026-01-01T00:00:00Z")) + .unwrap(); + assert_eq!( + provider + .for_agent("a") + .turn_states + .get("t") + .unwrap() + .unwrap() + .lifecycle, + TurnLifecycle::Started, + "a recovered agent is never swept again" + ); +} diff --git a/crates/tinyagents-session/src/port/drivers/refused.rs b/crates/tinyagents-session/src/port/drivers/refused.rs new file mode 100644 index 000000000..7e3968497 --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/refused.rs @@ -0,0 +1,142 @@ +//! Stores that refuse every call: what an agent gets when the backend could +//! not bind its scope, so the failure surfaces on use instead of the agent's +//! data landing somewhere it does not belong. + +use std::sync::Arc; + +use serde_json::Value; +use tinystoragedrivers_core::{ + Blocking, Capabilities, CollectionSpec, DocumentStore, ErrorKind, Filter, Page, Precondition, + Query, Result, Sort, StorageError, StreamEntry, StreamStore, Version, Versioned, async_trait, +}; + +use super::super::AgentStores; +use super::{DriverAppendStore, DriverStore, DriverTranscriptLocator, DriverTurnStates}; + +/// A document and stream store whose every call fails with one error. +#[derive(Debug)] +pub(super) struct Refused { + kind: ErrorKind, + message: String, +} + +impl Refused { + fn error(&self) -> StorageError { + StorageError::new(self.kind, self.message.clone()) + } +} + +/// Stores over [`Refused`], reporting `error` on every call. +pub(super) fn stores(error: &StorageError, bridge: &Blocking) -> AgentStores { + let refused = Arc::new(Refused { + kind: error.kind(), + message: format!("session stores unavailable: {}", error.message()), + }); + let docs: Arc = refused.clone(); + let streams: Arc = refused; + AgentStores { + transcripts: Arc::new(DriverTranscriptLocator::new( + Arc::clone(&docs), + bridge.clone(), + "refused", + )), + turn_states: Arc::new(DriverTurnStates::new(Arc::clone(&docs), bridge.clone())), + kv: Arc::new(DriverStore::new(docs)), + journal: Arc::new(DriverAppendStore::new(streams)), + } +} + +#[async_trait] +impl DocumentStore for Refused { + fn capabilities(&self) -> Capabilities { + Capabilities::default() + } + + async fn ensure_collection(&self, _spec: &CollectionSpec) -> Result<()> { + Err(self.error()) + } + + async fn get(&self, _collection: &str, _id: &str) -> Result>> { + Err(self.error()) + } + + async fn put( + &self, + _collection: &str, + _id: &str, + _doc: Value, + _precondition: Precondition, + ) -> Result { + Err(self.error()) + } + + async fn delete(&self, _collection: &str, _id: &str, _pre: Precondition) -> Result { + Err(self.error()) + } + + async fn query(&self, _collection: &str, _query: &Query) -> Result>> { + Err(self.error()) + } + + async fn count(&self, _collection: &str, _filter: &Filter) -> Result { + Err(self.error()) + } + + async fn delete_where(&self, _collection: &str, _filter: &Filter) -> Result { + Err(self.error()) + } + + async fn claim( + &self, + _collection: &str, + _filter: &Filter, + _sort: &[Sort], + _patch: &Value, + ) -> Result>> { + Err(self.error()) + } + + async fn drop_collection(&self, _collection: &str) -> Result<()> { + Err(self.error()) + } +} + +#[async_trait] +impl StreamStore for Refused { + async fn append(&self, _stream: &str, _value: Value) -> Result { + Err(self.error()) + } + + async fn append_batch(&self, _stream: &str, _values: Vec) -> Result { + Err(self.error()) + } + + async fn read_window( + &self, + _stream: &str, + _from: u64, + _limit: usize, + ) -> Result> { + Err(self.error()) + } + + async fn len(&self, _stream: &str) -> Result { + Err(self.error()) + } + + async fn truncate_before(&self, _stream: &str, _offset: u64) -> Result { + Err(self.error()) + } + + async fn delete_stream(&self, _stream: &str) -> Result { + Err(self.error()) + } + + async fn streams(&self, _prefix: &str) -> Result> { + Err(self.error()) + } +} + +#[cfg(test)] +#[path = "refused_tests.rs"] +mod tests; diff --git a/crates/tinyagents-session/src/port/drivers/refused_tests.rs b/crates/tinyagents-session/src/port/drivers/refused_tests.rs new file mode 100644 index 000000000..acef2fd7f --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/refused_tests.rs @@ -0,0 +1,54 @@ +use super::*; + +fn refused() -> Refused { + Refused { + kind: ErrorKind::Unavailable, + message: "down".into(), + } +} + +#[tokio::test] +async fn every_document_call_reports_the_refusal() { + let docs = refused(); + let kind = |error: StorageError| { + assert_eq!(error.message(), "down"); + error.kind() + }; + assert_eq!(docs.capabilities(), Capabilities::default()); + assert_eq!( + kind( + docs.ensure_collection(&CollectionSpec::new("c")) + .await + .unwrap_err() + ), + ErrorKind::Unavailable + ); + assert!(docs.get("c", "i").await.is_err()); + assert!( + docs.put("c", "i", Value::Null, Precondition::None) + .await + .is_err() + ); + assert!(docs.delete("c", "i", Precondition::None).await.is_err()); + assert!(docs.query("c", &Query::all()).await.is_err()); + assert!(docs.count("c", &Filter::All).await.is_err()); + assert!(docs.delete_where("c", &Filter::All).await.is_err()); + assert!( + docs.claim("c", &Filter::All, &[], &Value::Null) + .await + .is_err() + ); + assert!(docs.drop_collection("c").await.is_err()); +} + +#[tokio::test] +async fn every_stream_call_reports_the_refusal() { + let streams = refused(); + assert!(streams.append("s", Value::Null).await.is_err()); + assert!(streams.append_batch("s", vec![]).await.is_err()); + assert!(streams.read_window("s", 0, 1).await.is_err()); + assert!(StreamStore::len(&streams, "s").await.is_err()); + assert!(streams.truncate_before("s", 0).await.is_err()); + assert!(streams.delete_stream("s").await.is_err()); + assert!(streams.streams("").await.is_err()); +} diff --git a/crates/tinyagents-session/src/port/drivers/transcripts.rs b/crates/tinyagents-session/src/port/drivers/transcripts.rs new file mode 100644 index 000000000..fd0cc1855 --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/transcripts.rs @@ -0,0 +1,1137 @@ +//! [`TranscriptLocator`] and [`TranscriptHistory`] over a driver +//! [`DocumentStore`]. +//! +//! # Layout +//! +//! Each transcript stem is an append-only log of entries in +//! [`ENTRIES`], one document per write, numbered `0, 1, 2, …`: +//! +//! ```text +//! { "stem", "seq", "written", "meta"?, "set"?, "extend"?, "tools"?, +//! "partial"?, "request_id"?, "clear"?, "seal"? } +//! ``` +//! +//! Replaying the entries in order rebuilds the transcript: `set` replaces the +//! logical messages (a first write or a compaction), `extend` appends to them +//! (an ordinary turn only stores its new rows), `partial` records a +//! display-only partial, `clear` empties both, and `seal` closes the +//! generation. A writer claims the next number with an insert-only write, so +//! two writers racing for the same transcript (two processes sharing a +//! database) are serialized: the loser re-reads, sees the winner's entry, and +//! decides again — which is also how a write that races a seal is refused. +//! +//! [`INDEX`] holds one small document per stem with the lookup fields +//! (`thread_id`, `agent_id`, `agent_name`, `subagent`, `created_at`), so +//! "newest root transcript of this thread" is one indexed query instead of a +//! scan of every log. It is refreshed after each write that changes them. A +//! successor generation is reserved there before its predecessor is sealed. + +use std::future::Future; +use std::path::{Path, PathBuf}; +use std::sync::Arc; + +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; +use tinystoragedrivers_core::{ + Blocking, CollectionSpec, DocumentStore, DocumentStoreExt, ErrorKind, Filter, IndexSpec, + Precondition, Query, Sort, StorageError, Version, +}; +use tokio::sync::{Mutex, OnceCell}; + +use super::super::memory::MAX_GENERATIONS; +use crate::transcript::{ + SessionRef, SessionTranscript, TranscriptHistory, TranscriptLocator, TranscriptMessage, + TranscriptMeta, TranscriptPartial, TranscriptRead, TranscriptTurn, TurnUsage, + same_transcript_messages, session_stem, stamped_rows, +}; + +/// One document per transcript stem: the lookup fields. +pub(super) const INDEX: &str = "session_transcripts"; +/// One document per write to a transcript: its append-only log. +pub(super) const ENTRIES: &str = "session_transcript_entries"; + +/// Index refresh attempts after a durable write before leaving the repair +/// to the next write or read. +const INDEX_ATTEMPTS: usize = 3; + +/// Insert attempts before a contended write gives up. +const CAS_ATTEMPTS: usize = 64; + +/// How old an unwritten generation reservation must be before another +/// compaction may take it over. A reservation that old was left by a process +/// that stopped between reserving and writing; without a takeover its +/// sealed predecessor would refuse every later write. +const STALE_RESERVATION_MS: i64 = 30_000; + +/// Longest id before it is hashed, leaving room under the driver's limit for +/// an entry's `#` suffix. +const MAX_KEY_LEN: usize = 400; + +/// A document id for `parts`: length-prefixed so no two tuples collide, and +/// replaced by its SHA-256 when it would exceed [`MAX_KEY_LEN`]. +pub(super) fn doc_key(parts: &[&str]) -> String { + let joined: String = parts + .iter() + .map(|part| format!("{}:{part}", part.len())) + .collect::>() + .join("/"); + if joined.len() <= MAX_KEY_LEN { + joined + } else { + use sha2::{Digest, Sha256}; + format!("h:{}", hex::encode(Sha256::digest(joined.as_bytes()))) + } +} + +fn entry_id(stem: &str, seq: u64) -> String { + format!("{}#{seq:010}", doc_key(&[stem])) +} + +/// Whether `stem` names a sub-agent transcript: `__` separates a parent stem +/// from its child's. +fn is_subagent(stem: &str) -> bool { + stem.split_once("__") + .is_some_and(|(parent, child)| !parent.is_empty() && !child.is_empty()) +} + +fn now_rfc3339() -> String { + chrono::Utc::now().to_rfc3339_opts(chrono::SecondsFormat::Nanos, true) +} + +fn serialization(error: &serde_json::Error) -> StorageError { + StorageError::serialization(error.to_string()) +} + +/// A display-only partial as stored. +#[derive(Debug, Clone, Serialize, Deserialize)] +struct StoredPartial { + content: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + reasoning_content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + iteration: Option, +} + +impl From<&TranscriptPartial> for StoredPartial { + fn from(partial: &TranscriptPartial) -> Self { + Self { + content: partial.content.clone(), + reasoning_content: partial.reasoning_content.clone(), + iteration: partial.iteration, + } + } +} + +impl From for TranscriptPartial { + fn from(stored: StoredPartial) -> Self { + Self { + content: stored.content, + reasoning_content: stored.reasoning_content, + iteration: stored.iteration, + } + } +} + +/// One entry of a transcript's log. +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +struct Entry { + /// Whether this entry makes the transcript exist. Only a bare seal does + /// not: sealing a generation nobody wrote leaves it unwritten. + #[serde(default)] + written: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + meta: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + set: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + extend: Option>, + #[serde(default, skip_serializing_if = "Option::is_none")] + tools: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + partial: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + request_id: Option, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + clear: bool, + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + seal: bool, +} + +/// A transcript rebuilt from its log, up to `next_seq`. +#[derive(Debug, Default)] +struct Replay { + next_seq: u64, + meta: Option, + messages: Vec, + tools: Option, + partials: Vec<(TranscriptPartial, Option)>, + written: bool, + sealed: bool, + /// The lookup fields last written to [`INDEX`] by this handle. + indexed: Option, + /// The version of the generation reservation this handle was opened + /// with, until its first write claims it (see + /// [`HistoryInner::claim_reservation`]). + reservation: Option, +} + +impl Replay { + fn apply(&mut self, seq: u64, entry: Entry) { + self.next_seq = seq + 1; + if let Some(meta) = entry.meta { + self.meta = Some(meta); + } + if let Some(set) = entry.set { + self.messages = normalized(set); + } + if let Some(extend) = entry.extend { + self.messages.extend(normalized(extend)); + } + if let Some(tools) = entry.tools { + self.tools = Some(tools); + } + if let Some(partial) = entry.partial { + self.partials.push((partial.into(), entry.request_id)); + } + if entry.clear { + self.messages.clear(); + self.partials.clear(); + } + self.sealed |= entry.seal; + self.written |= entry.written; + } +} + +fn normalized(rows: Vec) -> Vec { + rows.into_iter() + .map(TranscriptMessage::normalized) + .collect() +} + +/// Declares both collections, once per locator. +#[derive(Debug, Default)] +struct Declared(OnceCell<()>); + +impl Declared { + async fn ensure(&self, docs: &Arc) -> Result<(), StorageError> { + self.0 + .get_or_try_init(|| async { + docs.ensure_collection( + &CollectionSpec::new(ENTRIES) + .index(IndexSpec::new("by_stem", ["stem", "seq"])) + .index(IndexSpec::new("by_stem_written", ["stem", "written"])), + ) + .await?; + docs.ensure_collection( + &CollectionSpec::new(INDEX) + .index(IndexSpec::new("by_thread", ["thread_id", "created_at"])) + .index(IndexSpec::new( + "by_agent_name", + ["agent_name", "created_at"], + )) + .index(IndexSpec::new("by_agent_id", ["agent_id", "created_at"])), + ) + .await + }) + .await + .map(|_| ()) + } +} + +/// Runs `future` on `bridge`, flattening the bridge's own failure into the +/// driver error. +fn run_on(bridge: &Blocking, future: Fut) -> Result +where + Fut: Future> + Send + 'static, + T: Send + 'static, +{ + bridge.run(future)? +} + +// ── History ────────────────────────────────────────────────────────── + +/// A [`TranscriptHistory`] bound to one stem of a driver-backed locator. +/// +/// Handles are cheap and independent: two handles on one stem (or two +/// processes) stay consistent because every write claims the next log number +/// with an insert-only write and re-reads on a clash. +pub struct DriverTranscriptHistory { + inner: Arc, + bridge: Blocking, + path: PathBuf, +} + +struct HistoryInner { + docs: Arc, + declared: Arc, + stem: String, + seed: TranscriptMeta, + replay: Mutex, +} + +impl std::fmt::Debug for DriverTranscriptHistory { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("DriverTranscriptHistory") + .field("path", &self.path) + .finish_non_exhaustive() + } +} + +impl DriverTranscriptHistory { + /// The display-only partials recorded so far, oldest first, with their + /// request ids. + /// + /// # Errors + /// + /// When the log cannot be read. + pub fn partials(&self) -> anyhow::Result)>> { + let inner = Arc::clone(&self.inner); + Ok(run_on(&self.bridge, async move { + let mut replay = inner.replay.lock().await; + inner.refresh(&mut replay).await?; + Ok(replay.partials.clone()) + })?) + } + + /// Appends `entry` built by `build` from the current replay; returns + /// whether anything was written (`build` answers `None` to write + /// nothing). + fn commit(&self, build: B) -> anyhow::Result + where + B: Fn(&Replay, &TranscriptMeta) -> anyhow::Result> + Send + 'static, + { + let inner = Arc::clone(&self.inner); + run_on(&self.bridge, async move { Ok(inner.commit(build).await) })? + } + + /// Seals this generation; later writes to it are refused. + /// + /// With a `baseline`, the seal is conditional on the transcript still + /// holding exactly those messages, checked against the same replay the + /// seal is claimed on: a turn another writer committed after the caller + /// read the baseline makes the seal fail instead of being dropped from + /// the successor. + fn seal(&self, baseline: Option>) -> anyhow::Result<()> { + self.commit(move |replay, _| { + if let Some(baseline) = &baseline { + let current: &[TranscriptMessage] = if replay.written { + &replay.messages + } else { + &[] + }; + anyhow::ensure!( + same_transcript_messages(current, baseline), + "transcript baseline is stale; reload the session before creating a generation" + ); + } + Ok((!replay.sealed).then(|| Entry { + seal: true, + ..Entry::default() + })) + }) + .map(|_| ()) + } + + /// Records the display-only `partial` of an interrupted turn, unless the + /// generation is sealed or the partial is empty. + fn record_partial( + &self, + partial: &TranscriptPartial, + request_id: Option<&str>, + ) -> anyhow::Result { + if partial.content.is_empty() { + return Ok(false); + } + let partial = StoredPartial::from(partial); + let request_id = request_id.map(str::to_string); + self.commit(move |replay, seed| { + Ok((!replay.sealed).then(|| Entry { + written: true, + meta: (!replay.written).then(|| seed.clone()), + partial: Some(partial.clone()), + request_id: request_id.clone(), + ..Entry::default() + })) + }) + } +} + +impl HistoryInner { + /// Applies every entry written since `replay` was last brought up to date. + async fn refresh(&self, replay: &mut Replay) -> Result<(), StorageError> { + self.declared.ensure(&self.docs).await?; + let query = Query::filter( + Filter::eq("stem", self.stem.as_str()).and(Filter::gte("seq", replay.next_seq)), + ) + .sort(Sort::asc("seq")); + for stored in self.docs.query_all(ENTRIES, &query).await? { + let seq = stored + .doc + .get("seq") + .and_then(Value::as_u64) + .ok_or_else(|| { + StorageError::serialization(format!( + "transcript entry {} has no seq", + stored.id + )) + })?; + let entry: Entry = serde_json::from_value(stored.doc).map_err(|error| { + StorageError::serialization(format!( + "transcript entry {} is unreadable: {error}", + stored.id + )) + })?; + replay.apply(seq, entry); + } + Ok(()) + } + + async fn commit(&self, build: B) -> anyhow::Result + where + B: Fn(&Replay, &TranscriptMeta) -> anyhow::Result>, + { + let mut replay = self.replay.lock().await; + for _ in 0..CAS_ATTEMPTS { + self.refresh(&mut replay).await?; + let Some(entry) = build(&replay, &self.seed)? else { + return Ok(false); + }; + self.claim_reservation(&mut replay).await?; + let seq = replay.next_seq; + let mut doc = serde_json::to_value(&entry).map_err(|error| serialization(&error))?; + if let Value::Object(fields) = &mut doc { + fields.insert("stem".into(), json!(self.stem)); + fields.insert("seq".into(), json!(seq)); + } + match self + .docs + .put( + ENTRIES, + &entry_id(&self.stem, seq), + doc, + Precondition::Absent, + ) + .await + { + Ok(_) => { + replay.apply(seq, entry); + // The entry is durable: report the write as done. A + // failed index refresh is retried a few times here, then + // by this handle's next write and by any read of the + // transcript; it never makes the durable write fail. + let mut last = None; + for _ in 0..INDEX_ATTEMPTS { + match self.index(&mut replay).await { + Ok(()) => { + last = None; + break; + } + Err(error) => last = Some(error), + } + } + if let Some(error) = last { + tracing::warn!( + target: "tinyagents_session::port::drivers", + stem = %self.stem, + seq, + %error, + "[session-store] transcript index refresh failed; retried on the next write or read" + ); + } + return Ok(true); + } + Err(error) if error.kind() == ErrorKind::Conflict => {} + Err(error) => return Err(error.into()), + } + } + anyhow::bail!( + "transcript {} kept changing under {CAS_ATTEMPTS} write attempts", + self.stem + ) + } + + /// Claims the generation reservation this handle was opened with, once, + /// before its first write: a compare-and-swap from the reserved version. + /// If another process took the reservation over in the meantime (this + /// one stalled past [`STALE_RESERVATION_MS`]), the generation is theirs + /// and this handle's writes are refused instead of overwriting it. + async fn claim_reservation(&self, replay: &mut Replay) -> anyhow::Result<()> { + let Some(reserved) = replay.reservation else { + return Ok(()); + }; + let claim = json!({ + "stem": self.stem, + "subagent": is_subagent(&self.stem), + "written": false, + "created_at": now_rfc3339(), + "reserved_at": chrono::Utc::now().timestamp_millis(), + }); + match self + .docs + .put( + INDEX, + &doc_key(&[&self.stem]), + claim, + Precondition::Version(reserved), + ) + .await + { + Ok(_) => { + replay.reservation = None; + Ok(()) + } + Err(error) if error.kind() == ErrorKind::Conflict => anyhow::bail!( + "session generation {} was taken over by another writer; reload the session", + self.stem + ), + Err(error) => Err(error.into()), + } + } + + /// Releases the reservation of a generation this handle never wrote. + async fn release_reservation(&self, replay: &mut Replay) -> Result<(), StorageError> { + let Some(reserved) = replay.reservation.take() else { + return Ok(()); + }; + match self + .docs + .delete( + INDEX, + &doc_key(&[&self.stem]), + Precondition::Version(reserved), + ) + .await + { + Err(error) if error.kind() != ErrorKind::Conflict => Err(error), + _ => Ok(()), + } + } + + /// Brings the index up to this replay on a read, best effort. A write + /// whose index refresh failed is otherwise only repaired by the next + /// write, and a transcript nobody writes again would stay invisible to + /// thread and agent lookups. Usually one read of a fresh index document. + async fn repair_index(&self, replay: &mut Replay) { + if let Err(error) = self.index(replay).await { + tracing::debug!( + target: "tinyagents_session::port::drivers", + stem = %self.stem, + %error, + "[session-store] transcript index repair on read failed" + ); + } + } + + /// Refreshes this stem's [`INDEX`] document when its lookup fields + /// changed. + /// + /// Fenced by log position: the document records the entry it was built + /// from (`indexed_seq`), and a handle whose replay is older than that + /// leaves it alone, so a delayed writer never puts back stale lookup + /// fields. The write itself is a compare-and-swap on the version read. + async fn index(&self, replay: &mut Replay) -> Result<(), StorageError> { + if !replay.written { + return Ok(()); + } + let meta = replay.meta.as_ref().unwrap_or(&self.seed); + let fields = json!({ + "stem": self.stem, + "subagent": is_subagent(&self.stem), + "written": true, + "thread_id": meta.thread_id, + "agent_id": meta.agent_id, + "agent_name": meta.agent_name, + }); + if replay.indexed.as_ref() == Some(&fields) { + return Ok(()); + } + let seq = replay.next_seq.saturating_sub(1); + let id = doc_key(&[&self.stem]); + for _ in 0..CAS_ATTEMPTS { + let existing = self.docs.get(INDEX, &id).await?; + let newer = existing.as_ref().is_some_and(|found| { + found.doc.get("written") == Some(&json!(true)) + && found + .doc + .get("indexed_seq") + .and_then(Value::as_u64) + .is_some_and(|indexed| indexed >= seq) + }); + if newer { + // Someone indexed this entry or a later one: the document is + // at least as fresh as this replay. + replay.indexed = Some(fields); + return Ok(()); + } + let created_at = existing + .as_ref() + .and_then(|found| found.doc.get("created_at").cloned()) + .unwrap_or_else(|| json!(now_rfc3339())); + let precondition = existing + .as_ref() + .map_or(Precondition::Absent, |found| found.unchanged()); + let mut doc = fields.clone(); + doc["created_at"] = created_at; + doc["indexed_seq"] = json!(seq); + match self.docs.put(INDEX, &id, doc, precondition).await { + Ok(_) => { + replay.indexed = Some(fields); + return Ok(()); + } + Err(error) if error.kind() == ErrorKind::Conflict => {} + Err(error) => return Err(error), + } + } + Err(StorageError::conflict(format!( + "transcript index for {} kept changing under {CAS_ATTEMPTS} attempts", + self.stem + ))) + } +} + +impl TranscriptRead for DriverTranscriptHistory { + fn path(&self) -> &Path { + &self.path + } + + fn read_session(&self) -> anyhow::Result> { + let inner = Arc::clone(&self.inner); + Ok(run_on(&self.bridge, async move { + let mut replay = inner.replay.lock().await; + inner.refresh(&mut replay).await?; + inner.repair_index(&mut replay).await; + Ok(replay.written.then(|| SessionTranscript { + meta: replay.meta.clone().unwrap_or_else(|| inner.seed.clone()), + messages: replay.messages.clone(), + tools: replay.tools.clone(), + })) + })?) + } +} + +/// The entry recording a turn whose logical messages are now `next`. +/// +/// Mirrors the JSONL writer: the turn's `prev` must be what is stored +/// (otherwise the caller's view is stale and the write is refused rather than +/// allowed to replace newer rows), an extension stores only its new rows and +/// anything else stores the whole set, and the written rows carry the turn's +/// usage and request id exactly as a transcript file records them. +fn turn_entry(replay: &Replay, turn: &TurnRecord) -> anyhow::Result> { + anyhow::ensure!(!replay.sealed, "transcript generation is sealed"); + let stored: &[TranscriptMessage] = if replay.written { + &replay.messages + } else { + &[] + }; + anyhow::ensure!( + !replay.written || same_transcript_messages(stored, &turn.prev), + "transcript baseline is stale; reload the session before persisting" + ); + let common = stored + .iter() + .zip(&turn.next) + .take_while(|(left, right)| left.same_row_as(right)) + .count(); + let usage = turn.turn_usage.as_ref(); + let request_id = turn.request_id.as_deref(); + let (set, extend) = if replay.written && common == stored.len() { + ( + None, + Some(stamped_rows(&turn.next[common..], usage, request_id)), + ) + } else { + (Some(stamped_rows(&turn.next, usage, request_id)), None) + }; + Ok(Some(Entry { + written: true, + meta: Some(turn.meta.clone()), + set, + extend, + tools: turn.tools.clone(), + partial: turn.partial.clone(), + request_id: turn.request_id.clone(), + ..Entry::default() + })) +} + +/// An owned [`TranscriptTurn`], so it can cross onto the bridge. +struct TurnRecord { + prev: Vec, + next: Vec, + meta: TranscriptMeta, + turn_usage: Option, + tools: Option, + request_id: Option, + partial: Option, +} + +impl TurnRecord { + fn new(turn: &TranscriptTurn<'_>, partial: Option<&TranscriptPartial>) -> Self { + Self { + prev: normalized(turn.prev.to_vec()), + next: normalized(turn.next.to_vec()), + meta: turn.meta.clone(), + turn_usage: turn.turn_usage.cloned(), + tools: turn.tools.cloned(), + request_id: turn.request_id.map(str::to_string), + partial: partial + .filter(|partial| !partial.content.is_empty()) + .map(StoredPartial::from), + } + } +} + +impl TranscriptHistory for DriverTranscriptHistory { + fn append_turn(&self, turn: TranscriptTurn<'_>) -> anyhow::Result<()> { + let record = TurnRecord::new(&turn, None); + self.commit(move |replay, _| turn_entry(replay, &record)) + .map(|_| ()) + } + + fn append_turn_with_partial( + &self, + turn: TranscriptTurn<'_>, + partial: Option<&TranscriptPartial>, + ) -> anyhow::Result<()> { + let record = TurnRecord::new(&turn, partial); + self.commit(move |replay, _| turn_entry(replay, &record)) + .map(|_| ()) + } + + fn messages(&self) -> anyhow::Result> { + let inner = Arc::clone(&self.inner); + Ok(run_on(&self.bridge, async move { + let mut replay = inner.replay.lock().await; + inner.refresh(&mut replay).await?; + Ok(replay.messages.clone()) + })?) + } + + fn append(&self, message: TranscriptMessage) -> anyhow::Result<()> { + let message = message.normalized(); + self.commit(move |replay, seed| { + anyhow::ensure!(!replay.sealed, "transcript generation is sealed"); + Ok(Some(Entry { + written: true, + meta: (!replay.written).then(|| seed.clone()), + extend: Some(vec![message.clone()]), + ..Entry::default() + })) + }) + .map(|_| ()) + } + + fn replace(&self, messages: &[TranscriptMessage]) -> anyhow::Result<()> { + let messages = normalized(messages.to_vec()); + self.commit(move |replay, seed| { + anyhow::ensure!(!replay.sealed, "transcript generation is sealed"); + Ok(Some(Entry { + written: true, + meta: (!replay.written).then(|| seed.clone()), + set: Some(messages.clone()), + ..Entry::default() + })) + }) + .map(|_| ()) + } + + fn clear(&self) -> anyhow::Result<()> { + // Clearing a reserved generation nobody wrote gives the reservation + // back, so the compaction can be retried at once instead of after + // the stale timeout. + let inner = Arc::clone(&self.inner); + let released = run_on(&self.bridge, async move { + let mut replay = inner.replay.lock().await; + inner.refresh(&mut replay).await?; + if replay.written || replay.reservation.is_none() { + return Ok(false); + } + inner.release_reservation(&mut replay).await?; + Ok(true) + })?; + if released { + return Ok(()); + } + self.commit(|replay, _| { + anyhow::ensure!(!replay.sealed, "transcript generation is sealed"); + Ok(replay.written.then(|| Entry { + written: true, + clear: true, + ..Entry::default() + })) + }) + .map(|_| ()) + } +} + +// ── Locator ────────────────────────────────────────────────────────── + +/// A [`TranscriptLocator`] keeping transcripts in a driver +/// [`DocumentStore`]; see the module docs for the layout. +#[derive(Clone)] +pub struct DriverTranscriptLocator { + docs: Arc, + bridge: Blocking, + label: String, + declared: Arc, +} + +impl std::fmt::Debug for DriverTranscriptLocator { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("DriverTranscriptLocator") + .field("label", &self.label) + .finish_non_exhaustive() + } +} + +impl DriverTranscriptLocator { + /// Transcripts in `docs`, with synchronous calls run on `bridge`. + /// + /// `label` names where `docs` points (driver, database, scope). It is + /// the locator's [`destination_key`](TranscriptLocator::destination_key) + /// and the prefix of every handle's path, so two locators should share a + /// label only when they address the same data. + pub fn new(docs: Arc, bridge: Blocking, label: impl Into) -> Self { + Self { + docs, + bridge, + label: label.into(), + declared: Arc::new(Declared::default()), + } + } + + fn handle(&self, stem: &str, seed: TranscriptMeta) -> DriverTranscriptHistory { + self.reserved_handle(stem, seed, None) + } + + /// A handle that holds the generation reservation `reservation`; its + /// first write claims it (or is refused if it was taken over). + fn reserved_handle( + &self, + stem: &str, + seed: TranscriptMeta, + reservation: Option, + ) -> DriverTranscriptHistory { + DriverTranscriptHistory { + inner: Arc::new(HistoryInner { + docs: Arc::clone(&self.docs), + declared: Arc::clone(&self.declared), + stem: stem.to_string(), + seed, + replay: Mutex::new(Replay { + reservation, + ..Replay::default() + }), + }), + bridge: self.bridge.clone(), + path: PathBuf::from(format!("{}/{stem}", self.label)), + } + } + + fn run(&self, op: F) -> Result + where + F: FnOnce(Arc, Arc) -> Fut, + Fut: Future> + Send + 'static, + T: Send + 'static, + { + run_on( + &self.bridge, + op(Arc::clone(&self.docs), Arc::clone(&self.declared)), + ) + } + + /// The stem of the newest written root transcript matching `filter`. + fn newest_root_stem(&self, filter: Filter, what: &str) -> Option { + let found = self.run(|docs, declared| async move { + declared.ensure(&docs).await?; + let query = Query::filter( + filter + .and(Filter::eq("subagent", false)) + .and(Filter::eq("written", true)), + ) + .sort(Sort::desc("created_at")) + .limit(1); + Ok(docs + .query(INDEX, &query) + .await? + .items + .into_iter() + .next() + .and_then(|found| { + found + .doc + .get("stem") + .and_then(Value::as_str) + .map(str::to_string) + })) + }); + found.unwrap_or_else(|error| { + tracing::warn!( + target: "tinyagents_session::port::drivers", + locator = %self.label, + lookup = what, + %error, + "[session-store] transcript lookup failed" + ); + None + }) + } + + /// The newest written root transcript matching `filter`. + fn newest_root(&self, filter: Filter, what: &str) -> Option> { + self.newest_root_stem(filter, what).map(|stem| { + let seed = discovered_seed(&stem); + Arc::new(self.handle(&stem, seed)) as Arc + }) + } + + /// Reserves `stem` as a new generation: succeeds when nothing is written + /// there and no live reservation holds it. + async fn reserve(docs: &Arc, stem: &str) -> anyhow::Result { + let written = docs + .count( + ENTRIES, + &Filter::eq("stem", stem).and(Filter::eq("written", true)), + ) + .await?; + anyhow::ensure!(written == 0, "session generation already exists"); + let id = doc_key(&[stem]); + let now = chrono::Utc::now().timestamp_millis(); + let reservation = json!({ + "stem": stem, + "subagent": is_subagent(stem), + "written": false, + "created_at": now_rfc3339(), + "reserved_at": now, + }); + let precondition = match docs.get(INDEX, &id).await? { + None => Precondition::Absent, + Some(found) => { + let stale = found.doc.get("written") == Some(&json!(false)) + && found + .doc + .get("reserved_at") + .and_then(Value::as_i64) + .is_some_and(|at| now - at >= STALE_RESERVATION_MS); + anyhow::ensure!(stale, "session generation already exists or is reserved"); + found.unchanged() + } + }; + match docs.put(INDEX, &id, reservation, precondition).await { + Ok(version) => Ok(version), + Err(error) if error.kind() == ErrorKind::Conflict => { + anyhow::bail!("session generation already exists or is reserved") + } + Err(error) => Err(error.into()), + } + } +} + +/// A seed for a handle on a transcript found by lookup. It is used only if +/// the transcript turns out to be unwritten, which a lookup never returns. +fn discovered_seed(stem: &str) -> TranscriptMeta { + TranscriptMeta { + agent_name: stem.to_string(), + agent_id: None, + agent_type: None, + dispatcher: String::new(), + provider: None, + model: None, + created: String::new(), + updated: String::new(), + turn_count: 0, + prefix_message_count: None, + input_tokens: 0, + output_tokens: 0, + cached_input_tokens: 0, + charged_amount_usd: 0.0, + thread_id: None, + task_id: None, + session_id: None, + parent_session_id: None, + } +} + +impl TranscriptLocator for DriverTranscriptLocator { + fn destination_key(&self) -> Option { + Some(format!("{}/transcripts", self.label)) + } + + fn latest_for_agent(&self, agent_name: &str) -> Option> { + let filter = Filter::eq("agent_name", agent_name).or(Filter::eq("agent_id", agent_name)); + self.newest_root(filter, "latest_for_agent") + } + + fn root_for_thread(&self, thread_id: &str) -> Option> { + self.root_for_thread_scoped(thread_id, None) + } + + fn root_for_thread_scoped( + &self, + thread_id: &str, + agent_id: Option<&str>, + ) -> Option> { + let thread_id = thread_id.trim(); + if thread_id.is_empty() { + return None; + } + let mut filter = Filter::eq("thread_id", thread_id); + if let Some(agent_id) = agent_id { + filter = filter.and(Filter::eq("agent_id", agent_id)); + } + self.newest_root(filter, "root_for_thread") + } + + fn open_stem( + &self, + stem: &str, + seed: TranscriptMeta, + ) -> anyhow::Result> { + Ok(Arc::new(self.handle(stem, seed))) + } + + fn session_exists(&self, session: &SessionRef) -> bool { + let stem = session_stem(session); + let written = self.run(|docs, declared| async move { + declared.ensure(&docs).await?; + docs.count( + ENTRIES, + &Filter::eq("stem", stem).and(Filter::eq("written", true)), + ) + .await + }); + match written { + Ok(count) => count > 0, + Err(error) => { + tracing::warn!( + target: "tinyagents_session::port::drivers", + locator = %self.label, + %error, + "[session-store] transcript existence check failed" + ); + false + } + } + } + + fn read_session_transcript(&self, session: &SessionRef) -> Option> { + if !self.session_exists(session) { + return None; + } + let stem = session_stem(session); + let seed = discovered_seed(&stem); + Some(Arc::new(self.handle(&stem, seed))) + } + + fn append_interrupted_partial( + &self, + thread_id: &str, + agent_id: Option<&str>, + partial: &TranscriptPartial, + request_id: Option<&str>, + ) -> anyhow::Result { + if partial.content.is_empty() { + return Ok(false); + } + let thread_id = thread_id.trim(); + if thread_id.is_empty() { + return Ok(false); + } + let mut filter = Filter::eq("thread_id", thread_id); + if let Some(agent_id) = agent_id { + filter = filter.and(Filter::eq("agent_id", agent_id)); + } + let Some(stem) = self.newest_root_stem(filter, "append_interrupted_partial") else { + return Ok(false); + }; + self.handle(&stem, discovered_seed(&stem)) + .record_partial(partial, request_id) + } + + /// Reserves the successor generation, then seals `session`. The + /// reservation is what keeps two compactions from both opening the + /// successor; it is released again if the seal fails. + fn begin_generation( + &self, + session: &SessionRef, + seed: TranscriptMeta, + ) -> anyhow::Result<(SessionRef, Arc)> { + self.begin_generation_checked(session, seed, None) + } + + /// [`Self::begin_generation`], sealing only if the predecessor still + /// holds exactly `baseline` — checked as part of the seal itself, so a + /// turn committed in between cannot be lost from the successor. + fn begin_generation_from_baseline( + &self, + session: &SessionRef, + seed: TranscriptMeta, + baseline: &[TranscriptMessage], + ) -> anyhow::Result<(SessionRef, Arc)> { + self.begin_generation_checked(session, seed, Some(normalized(baseline.to_vec()))) + } +} + +impl DriverTranscriptLocator { + fn begin_generation_checked( + &self, + session: &SessionRef, + seed: TranscriptMeta, + baseline: Option>, + ) -> anyhow::Result<(SessionRef, Arc)> { + let successor = session.next_generation(); + anyhow::ensure!( + successor.generation <= MAX_GENERATIONS, + "session generation limit reached" + ); + let stem = session_stem(&successor); + let reservation = { + let stem = stem.clone(); + let docs = Arc::clone(&self.docs); + let declared = Arc::clone(&self.declared); + run_on(&self.bridge, async move { + declared.ensure(&docs).await?; + Ok(Self::reserve(&docs, &stem).await) + })?? + }; + let predecessor_stem = session_stem(session); + if let Err(error) = self.handle(&predecessor_stem, seed.clone()).seal(baseline) { + self.release(&stem, reservation); + return Err(error); + } + let mut meta = seed; + meta.session_id = Some(successor.session_id()); + meta.parent_session_id = successor.parent_session_id(); + let handle = self.reserved_handle(&stem, meta, Some(reservation)); + Ok((successor, Arc::new(handle))) + } + + /// Drops the reservation of `stem` this call made, at the version it + /// wrote: if another process has since taken it over or written the + /// successor, the document is theirs and stays. + fn release(&self, stem: &str, reservation: Version) { + let docs = Arc::clone(&self.docs); + let id = doc_key(&[stem]); + let released = run_on(&self.bridge, async move { + match docs + .delete(INDEX, &id, Precondition::Version(reservation)) + .await + { + Err(error) if error.kind() == ErrorKind::Conflict => Ok(false), + other => other, + } + }); + if let Err(error) = released { + tracing::warn!( + target: "tinyagents_session::port::drivers", + stem, + %error, + "[session-store] could not release a generation reservation" + ); + } + } +} + +#[cfg(test)] +#[path = "transcripts_tests.rs"] +pub(super) mod tests; diff --git a/crates/tinyagents-session/src/port/drivers/transcripts_tests.rs b/crates/tinyagents-session/src/port/drivers/transcripts_tests.rs new file mode 100644 index 000000000..69e762cfd --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/transcripts_tests.rs @@ -0,0 +1,669 @@ +use super::*; +use crate::testkit::conformance::transcript_history_conformance; +use tinystoragedrivers_core::{MemoryStorage, Scope, StorageBackend, Version}; + +pub(in crate::port::drivers) fn meta(thread_id: &str) -> TranscriptMeta { + let mut meta = discovered_seed("planner"); + meta.agent_id = Some("planner".to_string()); + meta.dispatcher = "native".to_string(); + meta.created = "2026-01-01T00:00:00Z".to_string(); + meta.updated = "2026-01-01T00:00:00Z".to_string(); + meta.thread_id = Some(thread_id.to_string()); + meta +} + +fn message(role: &str, content: &str) -> TranscriptMessage { + serde_json::from_value(json!({ "role": role, "content": content })).unwrap() +} + +fn docs() -> Arc { + let storage = MemoryStorage::new(); + Arc::clone(storage.for_scope(&Scope::local()).unwrap().documents()) +} + +fn locator(docs: &Arc) -> DriverTranscriptLocator { + DriverTranscriptLocator::new(Arc::clone(docs), Blocking::new().unwrap(), "memory://test") +} + +fn turn<'a>( + prev: &'a [TranscriptMessage], + next: &'a [TranscriptMessage], + meta: &'a TranscriptMeta, +) -> TranscriptTurn<'a> { + TranscriptTurn { + prev, + next, + meta, + turn_usage: None, + request_id: Some("req"), + tools: None, + } +} + +/// The raw log entries of `stem`, in order. +fn entries(docs: &Arc, stem: &str) -> Vec { + let docs = Arc::clone(docs); + let stem = stem.to_string(); + Blocking::new() + .unwrap() + .run(async move { + docs.query_all( + ENTRIES, + &Query::filter(Filter::eq("stem", stem)).sort(Sort::asc("seq")), + ) + .await + .unwrap() + .into_iter() + .map(|found| found.doc) + .collect() + }) + .unwrap() +} + +fn on_bridge( + future: impl Future> + Send + 'static, +) -> T { + Blocking::new().unwrap().run(future).unwrap().unwrap() +} + +#[test] +fn a_handle_meets_the_transcript_history_contract() { + let docs = docs(); + let history = locator(&docs).open_stem("contract", meta("t")).unwrap(); + transcript_history_conformance(history.as_ref()); +} + +#[test] +fn an_ordinary_turn_stores_only_its_new_rows() { + let docs = docs(); + let history = locator(&docs).open_stem("s", meta("t")).unwrap(); + let first = vec![message("user", "a"), message("assistant", "b")]; + history.append_turn(turn(&[], &first, &meta("t"))).unwrap(); + let mut second = first.clone(); + second.extend([message("user", "c"), message("assistant", "d")]); + history + .append_turn(turn(&first, &second, &meta("t"))) + .unwrap(); + let compacted = vec![message("user", "summary")]; + history + .append_turn(turn(&second, &compacted, &meta("t"))) + .unwrap(); + + let log = entries(&docs, "s"); + assert_eq!(log.len(), 3); + assert_eq!(log[0]["set"].as_array().unwrap().len(), 2); + assert_eq!(log[1]["extend"].as_array().unwrap().len(), 2); + assert!(log[1].get("set").is_none()); + assert_eq!(log[2]["set"].as_array().unwrap().len(), 1); + assert_eq!(log[2]["seq"], json!(2)); + let contents: Vec = history + .messages() + .unwrap() + .into_iter() + .map(|row| row.content) + .collect(); + assert_eq!(contents, ["summary"]); +} + +#[test] +fn two_handles_on_one_stem_never_lose_each_others_writes() { + let docs = docs(); + let locator = locator(&docs); + let one = locator.open_stem("s", meta("t")).unwrap(); + let two = locator.open_stem("s", meta("t")).unwrap(); + one.append(message("user", "from one")).unwrap(); + two.append(message("user", "from two")).unwrap(); + one.append(message("user", "one again")).unwrap(); + let contents: Vec = two + .messages() + .unwrap() + .into_iter() + .map(|row| row.content) + .collect(); + assert_eq!(contents, ["from one", "from two", "one again"]); +} + +#[test] +fn concurrent_writers_serialize_through_the_log() { + let docs = docs(); + let locator = locator(&docs); + let threads: Vec<_> = (0..4) + .map(|writer| { + let handle = locator.open_stem("busy", meta("t")).unwrap(); + std::thread::spawn(move || { + for i in 0..5 { + handle + .append(message("user", &format!("{writer}-{i}"))) + .unwrap(); + } + }) + }) + .collect(); + for thread in threads { + thread.join().unwrap(); + } + let reader = locator.open_stem("busy", meta("t")).unwrap(); + assert_eq!(reader.messages().unwrap().len(), 20); + let seqs: Vec = entries(&docs, "busy") + .iter() + .map(|entry| entry["seq"].as_u64().unwrap()) + .collect(); + assert_eq!(seqs, (0..20).collect::>(), "the log is dense"); +} + +#[test] +fn tools_and_meta_follow_the_latest_turn() { + let docs = docs(); + let history = locator(&docs).open_stem("s", meta("t")).unwrap(); + let rows = vec![message("user", "a")]; + let tools = json!([{ "name": "search" }]); + history + .append_turn(TranscriptTurn { + tools: Some(&tools), + ..turn(&[], &rows, &meta("t")) + }) + .unwrap(); + let mut later = meta("t"); + later.turn_count = 2; + history.append_turn(turn(&rows, &rows, &later)).unwrap(); + let read = history.read_session().unwrap().unwrap(); + assert_eq!( + read.tools, + Some(tools), + "a turn without tools keeps the last" + ); + assert_eq!(read.meta.turn_count, 2); +} + +#[test] +fn partials_stay_out_of_the_replay_and_a_clear_drops_them() { + let docs = docs(); + let locator = locator(&docs); + let handle = locator.handle("s", meta("t")); + assert!( + !handle + .record_partial(&TranscriptPartial::new(""), None) + .unwrap() + ); + let rows = vec![message("user", "a")]; + let mut partial = TranscriptPartial::new("half"); + partial.reasoning_content = Some("thinking".into()); + partial.iteration = Some(2); + handle + .append_turn_with_partial(turn(&[], &rows, &meta("t")), Some(&partial)) + .unwrap(); + assert!( + handle + .record_partial(&TranscriptPartial::new("more"), Some("r2")) + .unwrap() + ); + let partials = handle.partials().unwrap(); + assert_eq!(partials.len(), 2); + assert_eq!(partials[0], (partial, Some("req".to_string()))); + assert_eq!(partials[1].1.as_deref(), Some("r2")); + assert_eq!(handle.messages().unwrap().len(), 1); + handle.clear().unwrap(); + assert!(handle.partials().unwrap().is_empty()); + assert!(handle.messages().unwrap().is_empty()); + assert!( + handle.read_session().unwrap().is_some(), + "a cleared transcript still exists" + ); +} + +#[test] +fn a_partial_alone_makes_the_transcript_exist() { + let docs = docs(); + let handle = locator(&docs).handle("s", meta("t")); + assert!( + handle + .record_partial(&TranscriptPartial::new("half"), None) + .unwrap() + ); + let read = handle.read_session().unwrap().unwrap(); + assert!(read.messages.is_empty()); + assert_eq!( + read.meta.thread_id.as_deref(), + Some("t"), + "the seed is the meta" + ); +} + +#[test] +fn a_sealed_generation_refuses_writes() { + let docs = docs(); + let locator = locator(&docs); + let session = SessionRef::scoped("t", "planner"); + let history = locator.open_session(&session, meta("t")).unwrap(); + history.append(message("user", "a")).unwrap(); + let (successor, _next) = locator.begin_generation(&session, meta("t")).unwrap(); + assert_eq!(successor.generation, 1); + for error in [ + history.append(message("user", "late")).unwrap_err(), + history.replace(&[]).unwrap_err(), + history.clear().unwrap_err(), + history.append_turn(turn(&[], &[], &meta("t"))).unwrap_err(), + ] { + assert!(error.to_string().contains("sealed"), "{error}"); + } + let handle = locator.handle(&session_stem(&session), meta("t")); + assert!( + !handle + .record_partial(&TranscriptPartial::new("x"), None) + .unwrap() + ); + handle.seal(None).unwrap(); + assert_eq!( + entries(&docs, &session_stem(&session)) + .iter() + .filter(|entry| entry["seal"] == json!(true)) + .count(), + 1, + "sealing twice writes one seal" + ); +} + +#[test] +fn a_generation_is_opened_once() { + let docs = docs(); + let locator = locator(&docs); + let session = SessionRef::scoped("t", "planner"); + locator + .open_session(&session, meta("t")) + .unwrap() + .append(message("user", "a")) + .unwrap(); + let (successor, next) = locator.begin_generation(&session, meta("t")).unwrap(); + let error = locator.begin_generation(&session, meta("t")).err().unwrap(); + assert!(error.to_string().contains("reserved"), "{error}"); + next.append(message("user", "summary")).unwrap(); + let error = locator.begin_generation(&session, meta("t")).err().unwrap(); + assert!(error.to_string().contains("already exists"), "{error}"); + assert_eq!(locator.head_generation(&session), successor); + let read = next.read_session().unwrap().unwrap(); + assert_eq!(read.meta.session_id, Some(successor.session_id())); + assert_eq!(read.meta.parent_session_id, successor.parent_session_id()); +} + +#[test] +fn a_stale_reservation_is_taken_over() { + let docs = docs(); + let locator = locator(&docs); + let session = SessionRef::scoped("t", "planner"); + locator + .open_session(&session, meta("t")) + .unwrap() + .append(message("user", "a")) + .unwrap(); + locator.begin_generation(&session, meta("t")).unwrap(); + // The process that reserved generation 1 stopped before writing it. + let id = doc_key(&[&session_stem(&session.next_generation())]); + let aged = Arc::clone(&docs); + on_bridge(async move { + let found = aged.get(INDEX, &id).await?.unwrap(); + let mut doc = found.doc; + doc["reserved_at"] = json!(0); + aged.put(INDEX, &id, doc, Precondition::None).await + }); + let (successor, next) = locator.begin_generation(&session, meta("t")).unwrap(); + next.append(message("user", "summary")).unwrap(); + assert_eq!(locator.head_generation(&session), successor); +} + +#[test] +fn the_generation_limit_holds() { + let docs = docs(); + let mut session = SessionRef::scoped("t", "planner"); + session.generation = MAX_GENERATIONS; + let error = locator(&docs) + .begin_generation(&session, meta("t")) + .err() + .unwrap(); + assert!(error.to_string().contains("limit"), "{error}"); +} + +#[test] +fn lookups_find_the_newest_written_root() { + let docs = docs(); + let locator = locator(&docs); + assert!(locator.root_for_thread(" ").is_none()); + assert!( + !locator + .append_interrupted_partial(" ", None, &TranscriptPartial::new("x"), None) + .unwrap() + ); + assert!( + !locator + .append_interrupted_partial("t", None, &TranscriptPartial::new(""), None) + .unwrap() + ); + + // An opened but unwritten stem is not found. + let _unwritten = locator.open_stem("never", meta("t")).unwrap(); + assert!(locator.root_for_thread("t").is_none()); + + locator + .open_stem("old", meta("t")) + .unwrap() + .append(message("user", "old")) + .unwrap(); + locator + .open_stem("new", meta("t")) + .unwrap() + .append(message("user", "new")) + .unwrap(); + locator + .open_stem("new__child", meta("t")) + .unwrap() + .append(message("user", "child")) + .unwrap(); + let newest = locator.root_for_thread("t").unwrap(); + assert_eq!(newest.path(), Path::new("memory://test/new")); + assert_eq!( + locator.latest_for_agent("planner").unwrap().path(), + Path::new("memory://test/new"), + "sub-agent transcripts are never roots" + ); + assert!(locator.root_for_thread_scoped("t", Some("other")).is_none()); + assert!(locator.latest_for_agent("nobody").is_none()); + + assert!( + locator + .append_interrupted_partial("t", Some("planner"), &TranscriptPartial::new("p"), None) + .unwrap() + ); + assert_eq!( + locator.handle("new", meta("t")).partials().unwrap().len(), + 1 + ); +} + +#[test] +fn a_thread_change_moves_the_transcript_in_the_index() { + let docs = docs(); + let locator = locator(&docs); + let history = locator.open_stem("s", meta("first")).unwrap(); + let rows = vec![message("user", "a")]; + history + .append_turn(turn(&[], &rows, &meta("first"))) + .unwrap(); + history + .append_turn(turn(&rows, &rows, &meta("second"))) + .unwrap(); + assert!(locator.root_for_thread("first").is_none()); + assert!(locator.root_for_thread("second").is_some()); +} + +#[test] +fn ids_are_unambiguous_and_bounded() { + assert_ne!(doc_key(&["a/b", "c"]), doc_key(&["a", "b/c"])); + let long = "x".repeat(1_000); + let hashed = doc_key(&[&long]); + assert!(hashed.starts_with("h:") && hashed.len() == 66, "{hashed}"); + assert!(entry_id(&long, 9).len() < tinystoragedrivers_core::MAX_ID_LEN); + assert_eq!(entry_id("s", 7), "1:s#0000000007"); + assert!(is_subagent("parent__child")); + assert!(!is_subagent("__child")); + assert!(!is_subagent("plain")); +} + +#[test] +fn a_long_stem_still_round_trips() { + let docs = docs(); + let stem = "s".repeat(600); + let history = locator(&docs).open_stem(&stem, meta("t")).unwrap(); + history.append(message("user", "a")).unwrap(); + assert_eq!(history.messages().unwrap().len(), 1); +} + +#[test] +fn an_unreadable_entry_is_an_error() { + let docs = docs(); + let locator = locator(&docs); + let history = locator.open_stem("s", meta("t")).unwrap(); + history.append(message("user", "a")).unwrap(); + let raw = Arc::clone(&docs); + on_bridge(async move { + raw.put( + ENTRIES, + &entry_id("s", 1), + json!({ "stem": "s", "seq": 1, "set": "not rows" }), + Precondition::None, + ) + .await + }); + assert!(history.messages().is_err()); +} + +#[test] +fn the_index_keeps_its_creation_time() { + let docs = docs(); + let history = locator(&docs).open_stem("s", meta("t")).unwrap(); + let rows = vec![message("user", "a")]; + history.append_turn(turn(&[], &rows, &meta("t"))).unwrap(); + let read = Arc::clone(&docs); + let first = on_bridge(async move { read.get(INDEX, &doc_key(&["s"])).await }); + history + .append_turn(turn(&rows, &rows, &meta("t2"))) + .unwrap(); + let read = Arc::clone(&docs); + let second = on_bridge(async move { read.get(INDEX, &doc_key(&["s"])).await }); + let (first, second) = (first.unwrap(), second.unwrap()); + assert_eq!(first.doc["created_at"], second.doc["created_at"]); + assert_eq!(second.doc["thread_id"], json!("t2")); + assert!(second.version > first.version || first.version == Version(1)); +} + +#[test] +fn handles_describe_themselves() { + let docs = docs(); + let locator = locator(&docs); + assert!(format!("{locator:?}").contains("memory://test")); + assert!(format!("{:?}", locator.handle("s", meta("t"))).contains("memory://test/s")); + assert_eq!( + locator.destination_key().as_deref(), + Some("memory://test/transcripts") + ); +} + +#[test] +fn a_stale_turn_is_refused_instead_of_replacing_newer_rows() { + let docs = docs(); + let locator = locator(&docs); + let one = locator.open_stem("s", meta("t")).unwrap(); + let two = locator.open_stem("s", meta("t")).unwrap(); + let a = vec![message("user", "a")]; + one.append_turn(turn(&[], &a, &meta("t"))).unwrap(); + let mut ab = a.clone(); + ab.push(message("assistant", "b")); + one.append_turn(turn(&a, &ab, &meta("t"))).unwrap(); + // `two` still believes the transcript is `a`. + let mut ac = a.clone(); + ac.push(message("assistant", "c")); + let error = two.append_turn(turn(&a, &ac, &meta("t"))).unwrap_err(); + assert!(error.to_string().contains("stale"), "{error}"); + let contents: Vec = two + .messages() + .unwrap() + .into_iter() + .map(|row| row.content) + .collect(); + assert_eq!(contents, ["a", "b"], "the newer rows survive"); +} + +#[test] +fn a_turn_keeps_its_usage_and_request_id() { + let docs = docs(); + let history = locator(&docs).open_stem("s", meta("t")).unwrap(); + let rows = vec![message("user", "q"), message("assistant", "a")]; + let usage: crate::transcript::TurnUsage = serde_json::from_value(json!({ + "provider": "p", + "model": "m", + "usage": { "input": 7, "output": 3, "cached_input": 0, "cost_usd": 0.01 }, + })) + .unwrap(); + history + .append_turn(TranscriptTurn { + turn_usage: Some(&usage), + request_id: Some("req-1"), + ..turn(&[], &rows, &meta("t")) + }) + .unwrap(); + let read = history.messages().unwrap(); + assert!( + read.iter() + .all(|row| row.request_id.as_deref() == Some("req-1")) + ); + assert!( + read[1].turn_usage.is_some(), + "the assistant row carries the usage" + ); + assert!(read[0].turn_usage.is_none()); + + // The next turn, built from the rows as read back, extends them. + let mut next = read.clone(); + next.push(message("user", "again")); + history + .append_turn(TranscriptTurn { + request_id: Some("req-2"), + ..turn(&read, &next, &meta("t")) + }) + .unwrap(); + let log = entries(&docs, "s"); + assert_eq!(log[1]["extend"].as_array().unwrap().len(), 1); + let read = history.messages().unwrap(); + assert_eq!( + read[0].request_id.as_deref(), + Some("req-1"), + "old rows keep theirs" + ); + assert_eq!(read[2].request_id.as_deref(), Some("req-2")); +} + +#[test] +fn a_stale_baseline_fails_the_seal_and_frees_the_successor() { + let docs = docs(); + let locator = locator(&docs); + let session = SessionRef::scoped("t", "planner"); + let history = locator.open_session(&session, meta("t")).unwrap(); + let a = vec![message("user", "a")]; + history.append_turn(turn(&[], &a, &meta("t"))).unwrap(); + let mut ab = a.clone(); + ab.push(message("assistant", "b")); + history.append_turn(turn(&a, &ab, &meta("t"))).unwrap(); + + let error = locator + .begin_generation_from_baseline(&session, meta("t"), &a) + .err() + .unwrap(); + assert!(error.to_string().contains("stale"), "{error}"); + let mut abc = ab.clone(); + abc.push(message("user", "c")); + history + .append_turn(turn(&ab, &abc, &meta("t"))) + .expect("the predecessor was not sealed"); + + let (successor, _) = locator + .begin_generation_from_baseline(&session, meta("t"), &abc) + .expect("the released reservation is free again"); + assert_eq!(successor.generation, 1); +} + +#[test] +fn a_stale_handle_never_rewinds_the_index() { + let docs = docs(); + let locator = locator(&docs); + let rows = vec![message("user", "a")]; + let old = locator.handle("s", meta("old")); + old.append_turn(turn(&[], &rows, &meta("old"))).unwrap(); + let new = locator.open_stem("s", meta("new")).unwrap(); + new.append_turn(turn(&rows, &rows, &meta("new"))).unwrap(); + + // `old` re-indexes from a replay that predates `new`'s entry. + let inner = Arc::clone(&old.inner); + on_bridge(async move { + let mut replay = inner.replay.lock().await; + replay.indexed = None; + replay.next_seq = 1; + inner.index(&mut replay).await + }); + assert!(locator.root_for_thread("new").is_some()); + assert!(locator.root_for_thread("old").is_none()); +} + +#[test] +fn a_read_repairs_an_index_a_write_could_not_refresh() { + let docs = docs(); + let locator = locator(&docs); + let history = locator.handle("s", meta("t")); + history.append(message("user", "a")).unwrap(); + // The index refresh after that write was lost. + let raw = Arc::clone(&docs); + on_bridge(async move { + raw.delete(INDEX, &doc_key(&["s"]), Precondition::None) + .await + }); + assert!(locator.root_for_thread("t").is_none()); + locator + .handle("s", meta("t")) + .read_session() + .unwrap() + .unwrap(); + assert!( + locator.root_for_thread("t").is_some(), + "the read re-indexed it" + ); +} + +#[test] +fn clearing_an_unwritten_successor_frees_the_generation() { + let docs = docs(); + let locator = locator(&docs); + let session = SessionRef::scoped("t", "planner"); + locator + .open_session(&session, meta("t")) + .unwrap() + .append(message("user", "a")) + .unwrap(); + let (_, next) = locator.begin_generation(&session, meta("t")).unwrap(); + next.clear().unwrap(); + let (successor, again) = locator + .begin_generation(&session, meta("t")) + .expect("the cleared reservation is free at once"); + again.append(message("user", "summary")).unwrap(); + assert_eq!(locator.head_generation(&session), successor); +} + +#[test] +fn a_taken_over_generation_refuses_its_stalled_owner() { + let docs = docs(); + let locator = locator(&docs); + let session = SessionRef::scoped("t", "planner"); + locator + .open_session(&session, meta("t")) + .unwrap() + .append(message("user", "a")) + .unwrap(); + let (_, stalled) = locator.begin_generation(&session, meta("t")).unwrap(); + // The owner stalls past the stale timeout and another process takes over. + let id = doc_key(&[&session_stem(&session.next_generation())]); + let aged = Arc::clone(&docs); + on_bridge(async move { + let found = aged.get(INDEX, &id).await?.unwrap(); + let mut doc = found.doc; + doc["reserved_at"] = json!(0); + aged.put(INDEX, &id, doc, Precondition::None).await + }); + let (_, winner) = locator.begin_generation(&session, meta("t")).unwrap(); + winner.append(message("user", "winner")).unwrap(); + + let error = stalled.replace(&[message("user", "late")]).unwrap_err(); + assert!(error.to_string().contains("taken over"), "{error}"); + let contents: Vec = winner + .messages() + .unwrap() + .into_iter() + .map(|row| row.content) + .collect(); + assert_eq!(contents, ["winner"]); +} diff --git a/crates/tinyagents-session/src/port/drivers/turn_states.rs b/crates/tinyagents-session/src/port/drivers/turn_states.rs new file mode 100644 index 000000000..22d87f7d2 --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/turn_states.rs @@ -0,0 +1,410 @@ +//! [`TurnStates`] over a driver [`DocumentStore`]. +//! +//! One document per turn in [`COLLECTION`], keyed by `(thread, request)`: +//! `{ "thread_id", "request_id", "lifecycle", "turn": }`, with +//! the thread indexed so a thread's turns are one query. Every conditional +//! change (a write that must not clobber a completed turn, a settle, an +//! interruption sweep) is a compare-and-swap on the document's version, so +//! two processes sharing a database cannot lose each other's updates. +//! +//! Semantics match [`InMemoryTurnStates`](crate::port::InMemoryTurnStates): +//! "latest" is the newest `started_at`, completed turns are kept up to +//! [`COMPLETED_RETENTION`] per thread, and a completed turn is never +//! overwritten by a conditional write. + +use std::cmp::Ordering; +use std::collections::HashMap; +use std::future::Future; +use std::sync::Arc; + +use serde_json::{Value, json}; +use tinystoragedrivers_core::{ + Blocking, CollectionSpec, DocumentStore, DocumentStoreExt, ErrorKind, Filter, IndexSpec, + Precondition, Query, StorageError, Versioned, +}; +use tokio::sync::OnceCell; + +use super::super::TurnStates; +use super::super::memory::{ + COMPLETED_RETENTION, completed_newest_first, is_terminal, newest_first, +}; +use super::transcripts::doc_key; +use crate::turn_state::{TurnLifecycle, TurnState}; + +/// The collection turn snapshots live in. +pub(super) const COLLECTION: &str = "session_turn_states"; + +/// Compare-and-swap attempts before a contended update gives up. +const CAS_ATTEMPTS: usize = 64; + +/// [`TurnStates`] stored in a driver [`DocumentStore`]. +#[derive(Debug, Clone)] +pub struct DriverTurnStates { + inner: Arc, + bridge: Blocking, +} + +#[derive(Debug)] +struct Inner { + docs: Arc, + declared: OnceCell<()>, +} + +impl DriverTurnStates { + /// Turn states in `docs`, with synchronous calls run on `bridge`. + pub fn new(docs: Arc, bridge: Blocking) -> Self { + Self { + inner: Arc::new(Inner { + docs, + declared: OnceCell::new(), + }), + bridge, + } + } + + /// Runs `op` against the store on the bridge. + fn run(&self, op: F) -> Result + where + F: FnOnce(Arc) -> Fut, + Fut: Future> + Send + 'static, + T: Send + 'static, + { + let inner = Arc::clone(&self.inner); + let future = op(inner); + match self.bridge.run(future) { + Ok(Ok(value)) => Ok(value), + Ok(Err(error)) | Err(error) => Err(error.to_string()), + } + } +} + +fn id(thread_id: &str, request_id: &str) -> String { + doc_key(&[thread_id, request_id]) +} + +fn encode(state: &TurnState) -> Result { + let turn = serde_json::to_value(state) + .map_err(|error| StorageError::serialization(error.to_string()))?; + let lifecycle = serde_json::to_value(state.lifecycle) + .map_err(|error| StorageError::serialization(error.to_string()))?; + Ok(json!({ + "thread_id": state.thread_id, + "request_id": state.request_id, + "lifecycle": lifecycle, + "turn": turn, + })) +} + +fn decode(stored: Versioned) -> Result, StorageError> { + let Versioned { id, version, doc } = stored; + let turn = doc + .get("turn") + .cloned() + .ok_or_else(|| StorageError::serialization(format!("turn state {id} has no body")))?; + let turn = serde_json::from_value(turn).map_err(|error| { + StorageError::serialization(format!("turn state {id} is unreadable: {error}")) + })?; + Ok(Versioned { + id, + version, + doc: turn, + }) +} + +impl Inner { + async fn declared(&self) -> Result<(), StorageError> { + self.declared + .get_or_try_init(|| async { + let spec = CollectionSpec::new(COLLECTION) + .index(IndexSpec::new("by_thread", ["thread_id"])); + self.docs.ensure_collection(&spec).await + }) + .await + .map(|_| ()) + } + + async fn get_turn( + &self, + thread_id: &str, + request_id: &str, + ) -> Result>, StorageError> { + self.declared().await?; + self.docs + .get(COLLECTION, &id(thread_id, request_id)) + .await? + .map(decode) + .transpose() + } + + async fn query(&self, filter: Filter) -> Result>, StorageError> { + self.declared().await?; + self.docs + .query_all(COLLECTION, &Query::filter(filter)) + .await? + .into_iter() + .map(decode) + .collect() + } + + async fn thread(&self, thread_id: &str) -> Result, StorageError> { + let mut turns: Vec = self + .query(Filter::eq("thread_id", thread_id)) + .await? + .into_iter() + .map(|stored| stored.doc) + .collect(); + turns.sort_by(newest_first); + Ok(turns) + } + + async fn write(&self, state: &TurnState, pre: Precondition) -> Result<(), StorageError> { + self.declared().await?; + self.docs + .put( + COLLECTION, + &id(&state.thread_id, &state.request_id), + encode(state)?, + pre, + ) + .await + .map(|_| ()) + } + + /// Drops `thread_id`'s completed turns beyond the newest + /// [`COMPLETED_RETENTION`]. + /// + /// Each delete is conditional on the version the listing saw: a turn + /// rewritten in the meantime is someone else's newer state and stays. + async fn prune_completed(&self, thread_id: &str) -> Result<(), StorageError> { + let mut completed: Vec> = self + .query(Filter::eq("thread_id", thread_id)) + .await? + .into_iter() + .filter(|stored| stored.doc.lifecycle == TurnLifecycle::Completed) + .collect(); + completed.sort_by(|a, b| completed_newest_first(&a.doc, &b.doc)); + for stale in completed.iter().skip(COMPLETED_RETENTION) { + match self + .docs + .delete(COLLECTION, &stale.id, stale.unchanged()) + .await + { + Ok(_) => {} + Err(error) if is_race(&error) => {} + Err(error) => return Err(error), + } + } + Ok(()) + } + + /// Applies `change` to one turn under compare-and-swap; returns whether + /// it changed anything. `change` answers `None` to leave the turn alone. + async fn update( + &self, + thread_id: &str, + request_id: &str, + change: impl Fn(&TurnState) -> Option, + ) -> Result, StorageError> { + for _ in 0..CAS_ATTEMPTS { + let Some(stored) = self.get_turn(thread_id, request_id).await? else { + return Ok(None); + }; + let Some(next) = change(&stored.doc) else { + return Ok(None); + }; + match self.write(&next, stored.unchanged()).await { + Ok(()) => return Ok(Some(next)), + Err(error) if is_race(&error) => {} + Err(error) => return Err(error), + } + } + Err(contended(thread_id, request_id)) + } +} + +/// Whether `error` means another writer got there first. +fn is_race(error: &StorageError) -> bool { + error.kind() == ErrorKind::Conflict +} + +fn contended(thread_id: &str, request_id: &str) -> StorageError { + StorageError::conflict(format!( + "turn {thread_id}/{request_id} kept changing under {CAS_ATTEMPTS} attempts" + )) +} + +/// `turn` marked interrupted as of `now`. +fn interrupted(turn: &TurnState, now: &str) -> Option { + if is_terminal(turn) { + return None; + } + let mut next = turn.clone(); + next.lifecycle = TurnLifecycle::Interrupted; + next.updated_at = now.to_string(); + next.active_tool = None; + next.active_subagent = None; + Some(next) +} + +impl TurnStates for DriverTurnStates { + fn put(&self, state: &TurnState) -> Result<(), String> { + let state = state.clone(); + self.run(|inner| async move { + inner.write(&state, Precondition::None).await?; + if state.lifecycle == TurnLifecycle::Completed { + inner.prune_completed(&state.thread_id).await?; + } + Ok(()) + }) + } + + fn put_unless_completed(&self, state: &TurnState) -> Result { + let state = state.clone(); + self.run(|inner| async move { + for _ in 0..CAS_ATTEMPTS { + let pre = match inner.get_turn(&state.thread_id, &state.request_id).await? { + Some(stored) if stored.doc.lifecycle == TurnLifecycle::Completed => { + return Ok(false); + } + Some(stored) => stored.unchanged(), + None => Precondition::Absent, + }; + match inner.write(&state, pre).await { + Ok(()) => { + if state.lifecycle == TurnLifecycle::Completed { + inner.prune_completed(&state.thread_id).await?; + } + return Ok(true); + } + Err(error) if is_race(&error) => {} + Err(error) => return Err(error), + } + } + Err(contended(&state.thread_id, &state.request_id)) + }) + } + + fn get(&self, thread_id: &str) -> Result, String> { + Ok(self.list_thread(thread_id)?.into_iter().next()) + } + + fn get_turn(&self, thread_id: &str, request_id: &str) -> Result, String> { + let (thread_id, request_id) = (thread_id.to_string(), request_id.to_string()); + self.run(|inner| async move { + Ok(inner + .get_turn(&thread_id, &request_id) + .await? + .map(|stored| stored.doc)) + }) + } + + fn delete(&self, thread_id: &str) -> Result { + let thread_id = thread_id.to_string(); + self.run(|inner| async move { + inner.declared().await?; + let removed = inner + .docs + .delete_where(COLLECTION, &Filter::eq("thread_id", thread_id)) + .await?; + Ok(removed > 0) + }) + } + + fn delete_turn(&self, thread_id: &str, request_id: &str) -> Result { + let key = id(thread_id, request_id); + self.run(|inner| async move { + inner.declared().await?; + inner + .docs + .delete(COLLECTION, &key, Precondition::None) + .await + }) + } + + fn list(&self) -> Result, String> { + self.run(|inner| async move { + let mut latest: HashMap = HashMap::new(); + for turn in inner.query(Filter::All).await? { + let turn = turn.doc; + match latest.get(&turn.thread_id) { + Some(kept) if newest_first(&turn, kept) != Ordering::Less => {} + _ => { + latest.insert(turn.thread_id.clone(), turn); + } + } + } + Ok(latest.into_values().collect()) + }) + } + + fn list_thread(&self, thread_id: &str) -> Result, String> { + let thread_id = thread_id.to_string(); + self.run(|inner| async move { inner.thread(&thread_id).await }) + } + + fn clear_all(&self) -> Result { + self.run(|inner| async move { + inner.declared().await?; + let removed = inner.docs.delete_where(COLLECTION, &Filter::All).await?; + Ok(usize::try_from(removed).unwrap_or(usize::MAX)) + }) + } + + fn mark_all_interrupted(&self, now_rfc3339: &str) -> Result { + let now = now_rfc3339.to_string(); + self.run(|inner| async move { + let mut count = 0; + for stored in inner.query(Filter::All).await? { + let turn = stored.doc; + if is_terminal(&turn) { + continue; + } + if inner + .update(&turn.thread_id, &turn.request_id, |current| { + interrupted(current, &now) + }) + .await? + .is_some() + { + count += 1; + } + } + Ok(count) + }) + } + + fn settle_turn( + &self, + thread_id: &str, + request_id: &str, + lifecycle: TurnLifecycle, + now_rfc3339: &str, + ) -> Result { + let (thread_id, request_id) = (thread_id.to_string(), request_id.to_string()); + let now = now_rfc3339.to_string(); + self.run(|inner| async move { + let settled = inner + .update(&thread_id, &request_id, |current| { + if is_terminal(current) { + return None; + } + let mut next = current.clone(); + next.lifecycle = lifecycle; + next.phase = None; + next.active_tool = None; + next.active_subagent = None; + next.updated_at.clone_from(&now); + Some(next) + }) + .await?; + if settled.is_some() && lifecycle == TurnLifecycle::Completed { + inner.prune_completed(&thread_id).await?; + } + Ok(settled.is_some()) + }) + } +} + +#[cfg(test)] +#[path = "turn_states_tests.rs"] +mod tests; diff --git a/crates/tinyagents-session/src/port/drivers/turn_states_tests.rs b/crates/tinyagents-session/src/port/drivers/turn_states_tests.rs new file mode 100644 index 000000000..f271acadb --- /dev/null +++ b/crates/tinyagents-session/src/port/drivers/turn_states_tests.rs @@ -0,0 +1,242 @@ +use super::*; +use tinystoragedrivers_core::{MemoryStorage, Scope, StorageBackend}; + +fn store() -> (DriverTurnStates, Arc) { + let storage = MemoryStorage::new(); + let docs = Arc::clone(storage.for_scope(&Scope::local()).unwrap().documents()); + ( + DriverTurnStates::new(Arc::clone(&docs), Blocking::new().unwrap()), + docs, + ) +} + +fn started(thread: &str, request: &str, minute: u32) -> TurnState { + TurnState::started(thread, request, 8, format!("2026-01-01T00:{minute:02}:00Z")) +} + +fn completed(thread: &str, request: &str, minute: u32) -> TurnState { + let mut turn = started(thread, request, minute); + turn.lifecycle = TurnLifecycle::Completed; + turn.updated_at = turn.started_at.clone(); + turn +} + +#[test] +fn completed_turns_are_kept_up_to_the_retention_limit() { + let (turns, _) = store(); + for minute in 0..(COMPLETED_RETENTION as u32 + 3) { + turns + .put(&completed("t", &format!("r{minute:02}"), minute)) + .unwrap(); + } + turns.put(&started("t", "live", 59)).unwrap(); + let kept = turns.list_thread("t").unwrap(); + assert_eq!(kept.len(), COMPLETED_RETENTION + 1); + assert!( + turns.get_turn("t", "r00").unwrap().is_none(), + "the oldest completed turns go first" + ); + assert!( + turns.get_turn("t", "live").unwrap().is_some(), + "live turns stay" + ); +} + +#[test] +fn settling_prunes_too() { + let (turns, _) = store(); + for minute in 0..COMPLETED_RETENTION as u32 { + turns + .put(&completed("t", &format!("r{minute:02}"), minute)) + .unwrap(); + } + turns.put(&started("t", "last", 58)).unwrap(); + assert!( + turns + .settle_turn( + "t", + "last", + TurnLifecycle::Completed, + "2026-01-01T01:00:00Z" + ) + .unwrap() + ); + assert_eq!(turns.list_thread("t").unwrap().len(), COMPLETED_RETENTION); +} + +#[test] +fn put_unless_completed_inserts_updates_and_then_holds() { + let (turns, _) = store(); + let mut turn = started("t", "r", 0); + assert!( + turns.put_unless_completed(&turn).unwrap(), + "a new turn is written" + ); + turn.lifecycle = TurnLifecycle::Completed; + assert!( + turns.put_unless_completed(&turn).unwrap(), + "a live turn is updated" + ); + assert!( + !turns.put_unless_completed(&started("t", "r", 1)).unwrap(), + "a completed turn is never overwritten conditionally" + ); + assert_eq!( + turns.get_turn("t", "r").unwrap().unwrap().lifecycle, + TurnLifecycle::Completed + ); +} + +#[test] +fn settling_needs_a_live_turn() { + let (turns, _) = store(); + let now = "2026-01-01T01:00:00Z"; + assert!( + !turns + .settle_turn("t", "missing", TurnLifecycle::Interrupted, now) + .unwrap() + ); + turns.put(&completed("t", "done", 0)).unwrap(); + assert!( + !turns + .settle_turn("t", "done", TurnLifecycle::Interrupted, now) + .unwrap() + ); + let mut live = started("t", "live", 1); + live.phase = Some(crate::turn_state::TurnPhase::Thinking); + turns.put(&live).unwrap(); + assert!( + turns + .settle_turn("t", "live", TurnLifecycle::Interrupted, now) + .unwrap() + ); + let settled = turns.get_turn("t", "live").unwrap().unwrap(); + assert_eq!(settled.lifecycle, TurnLifecycle::Interrupted); + assert_eq!(settled.phase, None); + assert_eq!(settled.updated_at, now); +} + +#[test] +fn deletes_report_whether_anything_went() { + let (turns, _) = store(); + assert!(!turns.delete("t").unwrap()); + assert!(!turns.delete_turn("t", "r").unwrap()); + turns.put(&started("t", "r1", 0)).unwrap(); + turns.put(&started("t", "r2", 1)).unwrap(); + turns.put(&started("u", "r1", 2)).unwrap(); + assert!(turns.delete("t").unwrap()); + assert!(turns.list_thread("t").unwrap().is_empty()); + assert_eq!(turns.list().unwrap().len(), 1); + assert_eq!(turns.clear_all().unwrap(), 1); + assert_eq!(turns.clear_all().unwrap(), 0); +} + +#[test] +fn list_answers_each_threads_newest_turn() { + let (turns, _) = store(); + turns.put(&started("t", "old", 0)).unwrap(); + turns.put(&started("t", "new", 5)).unwrap(); + turns.put(&started("u", "only", 1)).unwrap(); + let mut latest: Vec<(String, String)> = turns + .list() + .unwrap() + .into_iter() + .map(|turn| (turn.thread_id, turn.request_id)) + .collect(); + latest.sort(); + assert_eq!( + latest, + [ + ("t".to_string(), "new".to_string()), + ("u".to_string(), "only".to_string()) + ] + ); +} + +#[test] +fn ids_with_separators_do_not_collide() { + let (turns, _) = store(); + turns.put(&started("a/b", "c", 0)).unwrap(); + turns.put(&started("a", "b/c", 1)).unwrap(); + assert_eq!(turns.list().unwrap().len(), 2); +} + +#[test] +fn an_unreadable_snapshot_is_an_error() { + let (turns, docs) = store(); + turns.put(&started("t", "r", 0)).unwrap(); + let bridge = Blocking::new().unwrap(); + let raw = Arc::clone(&docs); + bridge + .run(async move { + raw.put( + COLLECTION, + &id("t", "bodiless"), + json!({"thread_id": "t"}), + Precondition::None, + ) + .await + }) + .unwrap() + .unwrap(); + let error = turns.list_thread("t").unwrap_err(); + assert!(error.contains("no body"), "{error}"); + let raw = Arc::clone(&docs); + bridge + .run(async move { + raw.put( + COLLECTION, + &id("t", "bodiless"), + json!({"thread_id": "t", "turn": "nope"}), + Precondition::None, + ) + .await + }) + .unwrap() + .unwrap(); + let error = turns.get_turn("t", "bodiless").unwrap_err(); + assert!(error.contains("unreadable"), "{error}"); +} + +#[test] +fn contention_is_reported_after_the_last_attempt() { + let error = contended("t", "r"); + assert_eq!(error.kind(), ErrorKind::Conflict); + assert!(is_race(&error)); + assert!(!is_race(&StorageError::unavailable("down"))); +} + +#[test] +fn concurrent_settles_and_sweeps_never_lose_a_turn() { + let (turns, _) = store(); + for i in 0..8 { + turns.put(&started("t", &format!("r{i}"), i)).unwrap(); + } + let handles: Vec<_> = (0..4) + .map(|worker| { + let turns = turns.clone(); + std::thread::spawn(move || { + if worker == 0 { + turns.mark_all_interrupted("2026-01-01T02:00:00Z").unwrap(); + } else { + for i in 0..8 { + turns + .settle_turn( + "t", + &format!("r{i}"), + TurnLifecycle::Completed, + "2026-01-01T02:00:00Z", + ) + .unwrap(); + } + } + }) + }) + .collect(); + for handle in handles { + handle.join().unwrap(); + } + let all = turns.list_thread("t").unwrap(); + assert_eq!(all.len(), 8); + assert!(all.iter().all(is_terminal), "every turn ended up terminal"); +} diff --git a/crates/tinyagents-session/src/port/memory/mod.rs b/crates/tinyagents-session/src/port/memory/mod.rs index 6a7d1779a..3bb9ec0bd 100644 --- a/crates/tinyagents-session/src/port/memory/mod.rs +++ b/crates/tinyagents-session/src/port/memory/mod.rs @@ -21,10 +21,10 @@ use crate::transcript::{ }; use crate::turn_state::{TurnLifecycle, TurnState}; -const MAX_GENERATIONS: u32 = 4096; +pub(super) const MAX_GENERATIONS: u32 = 4096; /// Completed turns kept per thread, as the on-disk store keeps them. -const COMPLETED_RETENTION: usize = 20; +pub(super) const COMPLETED_RETENTION: usize = 20; /// A [`SessionStoreProvider`] keeping every agent's stores in memory, each /// agent's apart from every other's. @@ -357,14 +357,14 @@ impl InMemoryTurnStates { } /// Newest first: by `started_at`, then `updated_at`. -fn newest_first(a: &TurnState, b: &TurnState) -> Ordering { +pub(super) fn newest_first(a: &TurnState, b: &TurnState) -> Ordering { compare_rfc3339(&b.started_at, &a.started_at) .then_with(|| compare_rfc3339(&b.updated_at, &a.updated_at)) } /// Completed-turn retention follows completion time, as the durable store /// does. Listing the latest live turn still follows its start time. -fn completed_newest_first(a: &TurnState, b: &TurnState) -> Ordering { +pub(super) fn completed_newest_first(a: &TurnState, b: &TurnState) -> Ordering { compare_rfc3339(&b.updated_at, &a.updated_at) .then_with(|| compare_rfc3339(&b.started_at, &a.started_at)) } @@ -383,7 +383,7 @@ fn key(thread_id: &str, request_id: &str) -> (String, String) { (thread_id.to_string(), request_id.to_string()) } -fn is_terminal(state: &TurnState) -> bool { +pub(super) fn is_terminal(state: &TurnState) -> bool { matches!( state.lifecycle, TurnLifecycle::Interrupted | TurnLifecycle::Completed diff --git a/crates/tinyagents-session/src/port/mod.rs b/crates/tinyagents-session/src/port/mod.rs index 94c3bdd62..0d390e12a 100644 --- a/crates/tinyagents-session/src/port/mod.rs +++ b/crates/tinyagents-session/src/port/mod.rs @@ -34,15 +34,24 @@ //! //! - [`InMemorySessionStores`]: per-agent, process-lifetime storage, for //! tests and for hosts that deliberately keep nothing. +//! - `DriverSessionStores` (feature `storage-drivers`): every agent's stores +//! in one `tinystoragedrivers` backend (SQLite, MongoDB, ...), each agent +//! under its own storage scope. //! - The on-disk layout (`session_raw/`, `tinyagents_store/`, turn-state //! files) is wrapped by the host that owns it, not here: this crate keeps //! the file and SQLite building blocks, and a host crate decides to use them. +#[cfg(feature = "storage-drivers")] +mod drivers; mod memory; mod types; use std::sync::Arc; +#[cfg(feature = "storage-drivers")] +pub use drivers::{ + DriverSessionStores, DriverTranscriptHistory, DriverTranscriptLocator, DriverTurnStates, +}; pub use memory::{InMemorySessionStores, InMemoryTranscriptLocator, InMemoryTurnStates}; pub use types::AgentStores; diff --git a/crates/tinyagents-session/src/transcript.rs b/crates/tinyagents-session/src/transcript.rs index 437f023e5..403d09d08 100644 --- a/crates/tinyagents-session/src/transcript.rs +++ b/crates/tinyagents-session/src/transcript.rs @@ -156,6 +156,8 @@ pub use history::{ FileTranscriptHistory, FileTranscriptLocator, TranscriptHistory, TranscriptLocator, TranscriptPartial, TranscriptRead, TranscriptTurn, TruncateCut, }; +#[cfg(feature = "storage-drivers")] +pub(crate) use jsonl::stamped_rows; pub use legacy_md::read_transcript_legacy_md; pub use migration::{TranscriptLayoutMigration, migrate_layout_if_needed}; pub use paths::{find_latest_transcript, resolve_keyed_transcript_path}; diff --git a/crates/tinyagents-session/src/transcript/jsonl.rs b/crates/tinyagents-session/src/transcript/jsonl.rs index cc935dfdf..2acc4de9c 100644 --- a/crates/tinyagents-session/src/transcript/jsonl.rs +++ b/crates/tinyagents-session/src/transcript/jsonl.rs @@ -616,23 +616,10 @@ pub(super) fn serialise_message_lines( request_id: Option<&str>, buf: &mut String, ) -> Result<()> { - let last_assistant_idx = messages.iter().rposition(|m| m.role == "assistant"); - let turn_stamps = turn_step_stamps(messages, last_assistant_idx, last_assistant_turn_usage); - for (i, msg) in messages.iter().enumerate() { - let turn_usage = if Some(i) == last_assistant_idx { - last_assistant_turn_usage - .cloned() - .or_else(|| msg.turn_usage.clone()) - } else { - msg.turn_usage.clone() - }; - let mut line = build_message_line(msg, turn_usage.as_ref(), request_id, false); - if let Some((iteration, ts)) = turn_stamps.get(&i) { - line.iteration = line.iteration.or(Some(*iteration)); - if line.ts.is_none() && !ts.is_empty() { - line.ts = Some(ts.clone()); - } - } + for (i, line) in stamped_lines(messages, last_assistant_turn_usage, request_id) + .into_iter() + .enumerate() + { let line_json = serde_json::to_string(&line).with_context(|| format!("serialise message line {i}"))?; buf.push_str(&line_json); @@ -641,6 +628,55 @@ pub(super) fn serialise_message_lines( Ok(()) } +/// The lines [`serialise_message_lines`] writes for `messages`: the turn's +/// usage on its last assistant row, `request_id` on every fresh row, and the +/// per-step `(iteration, ts)` stamps. +fn stamped_lines( + messages: &[TranscriptMessage], + last_assistant_turn_usage: Option<&TurnUsage>, + request_id: Option<&str>, +) -> Vec { + let last_assistant_idx = messages.iter().rposition(|m| m.role == "assistant"); + let turn_stamps = turn_step_stamps(messages, last_assistant_idx, last_assistant_turn_usage); + messages + .iter() + .enumerate() + .map(|(i, msg)| { + let turn_usage = if Some(i) == last_assistant_idx { + last_assistant_turn_usage + .cloned() + .or_else(|| msg.turn_usage.clone()) + } else { + msg.turn_usage.clone() + }; + let mut line = build_message_line(msg, turn_usage.as_ref(), request_id, false); + if let Some((iteration, ts)) = turn_stamps.get(&i) { + line.iteration = line.iteration.or(Some(*iteration)); + if line.ts.is_none() && !ts.is_empty() { + line.ts = Some(ts.clone()); + } + } + line + }) + .collect() +} + +/// `messages` exactly as the JSONL writer would record them for one turn and +/// a reader would return them: what a non-file transcript backend stores so +/// its replay carries the same per-turn provenance (usage, request ids, +/// step stamps) as a transcript file. +#[cfg(feature = "storage-drivers")] +pub(crate) fn stamped_rows( + messages: &[TranscriptMessage], + last_assistant_turn_usage: Option<&TurnUsage>, + request_id: Option<&str>, +) -> Vec { + stamped_lines(messages, last_assistant_turn_usage, request_id) + .into_iter() + .map(message_from_line) + .collect() +} + /// Per-step `(iteration, ts)` stamps for the intermediate assistant rows of the /// turn being written. /// diff --git a/docs/modules/session/store-port.md b/docs/modules/session/store-port.md index e92d83e4c..f9d9af91c 100644 --- a/docs/modules/session/store-port.md +++ b/docs/modules/session/store-port.md @@ -41,8 +41,8 @@ not claim isolation. `TranscriptLocator`/`TranscriptHistory` and `TurnStates` are synchronous: the turn path that commits a transcript is a chain of sync methods. An -implementation over an async client bridges at its own boundary (for tokio, -`block_in_place` on a multi-threaded runtime). `Store` and `AppendStore` are +implementation over an async client bridges at its own boundary (`DriverSessionStores` uses a dedicated runtime thread; +`block_in_place` also works, on a multi-threaded tokio runtime only). `Store` and `AppendStore` are the harness's async traits; `Arc` and `Arc` implement them, so code generic over a store accepts injected handles, and `FileStatusStore::over` keeps run status in any `Store`. @@ -59,6 +59,67 @@ implement them, so code generic over a store accepts injected handles, and `openhuman_rpc::session_store`). - `TranscriptLocator::append_interrupted_partial` carries the display-only partial of an interrupted turn without a file path. +- `DriverSessionStores` (feature `storage-drivers`): every agent's stores in + one `tinystoragedrivers` backend, described below. + +## `DriverSessionStores` + +With the `storage-drivers` feature, a host that has opened a +`tinystoragedrivers` backend (SQLite on a desktop, MongoDB in the cloud, +memory in tests) gets a complete provider from it: + +```rust +let backend: Arc = Arc::new(SqliteStorage::open(path)?); +let provider = DriverSessionStores::new(backend)?.recover_on_open(true); +``` + +Each agent id becomes a storage `Scope`. An id that is not a valid scope, +or that starts with the reserved `sha256:` prefix, maps to `sha256:` of +itself, so no raw id can name a hashed scope, and the driver enforces scopes, so the +provider claims isolation and passes `session_store_isolation_conformance`. + +| Store | Backend shape | +| --- | --- | +| transcripts | `session_transcripts`: one index document per stem (thread, agent, sub-agent flag, creation time); `session_transcript_entries`: an append-only log per stem | +| turn states | `session_turn_states`: one document per `(thread, request)` | +| key-value | `session_kv` through the harness `DriverStore` | +| journal | streams `16:session_journal/` (the harness `DriverAppendStore` always writes `:` first) | + +- **Transcript log.** Each write is one entry, numbered from 0 and claimed + with an insert-only write. An ordinary turn stores only its new rows + (`extend`); a first write or a compaction stores the whole set (`set`). Two + writers on one transcript (two processes on one database) serialize: the + loser re-reads and decides again, which is also how a write racing a seal is + refused. As in the JSONL writer, a turn whose `prev` is not what is stored + is refused as stale rather than allowed to replace newer rows, and the + written rows carry the turn's usage, request id and step stamps, built by + the writer's own code (`transcript::stamped_rows`). +- **Index.** Each index document records the log entry it was built from + (`indexed_seq`). A handle with an older replay leaves it alone, and the + update is a compare-and-swap. A failed index refresh never fails the write + that preceded it; the next write retries it. +- **Generations.** `begin_generation` reserves the successor in the index + before sealing the predecessor, so two compactions cannot both open it. A + reservation left unwritten for 30 seconds (its process stopped mid-way) may + be taken over, so a sealed head never strands the conversation. + `begin_generation_from_baseline` checks the baseline inside the seal + itself, so a turn committed after the caller read it fails the compaction + instead of vanishing from the successor. A failed seal releases the + reservation only at the version this call wrote. +- **Turn states.** Conditional writes, settling, the interruption sweep and + retention pruning are all conditional on the document version. +- **Sync seams.** Transcript and turn-state calls run on a + `tinystoragedrivers` `Blocking` bridge (one dedicated runtime thread), so + they work from any caller, inside a runtime or not. +- **Recovery.** A backend cannot list its scopes, so `recover` covers the + agents this provider has opened. `recover_on_open(true)` also interrupts an + agent's in-flight turns the first time it is opened. The agent counts as + recovered only once the sweep succeeds, and a concurrent first open waits + for that sweep. Use it only when a + single process owns the database (the desktop app). +- **Failing closed.** If the backend cannot bind an agent's scope, `for_agent` + returns stores that refuse every call with that error and does not cache + them. `try_for_agent` returns the error itself. ## Conformance @@ -66,8 +127,9 @@ implement them, so code generic over a store accepts injected handles, and whole provider: transcripts (sessions, thread and agent lookups, compaction generations, partials kept out of the replay), turn states (conditional writes, settling, interrupted-marking) and the key-value and journal stores. -`session_store_conformance` runs against `InMemorySessionStores` and the file -building blocks. It checks the common session-store behavior, not provider +`session_store_conformance` runs against `InMemorySessionStores`, the file +building blocks and `DriverSessionStores` over the memory and SQLite drivers. It checks the common session-store behavior, not provider isolation. `session_store_isolation_conformance` is a separate check that two agents cannot see each other's data; it runs only against -`InMemorySessionStores` and hosts whose providers claim isolation. +`InMemorySessionStores`, `DriverSessionStores` and hosts whose providers claim +isolation. diff --git a/vendor/tinystoragedrivers b/vendor/tinystoragedrivers new file mode 160000 index 000000000..77d08547e --- /dev/null +++ b/vendor/tinystoragedrivers @@ -0,0 +1 @@ +Subproject commit 77d08547e01f45da8e0f746705ad26d218600c72