Skip to main content

behest_runtime/
memory.rs

1//! In-memory implementations of runtime stores.
2//!
3//! Provides [`MemoryRunStore`], an ephemeral [`RunStore`] backed by
4//! `tokio::sync::RwLock`-guarded [`HashMap`]s. Useful for testing and
5//! single-process deployments where persistence is not required.
6//!
7//! # Architecture
8//!
9//! Each run is stored as three in-memory maps:
10//! - `runs` — [`RunRecord`] keyed by [`Uuid`].
11//! - `events` — [`RunEventRecord`] vectors keyed by run UUID.
12//! - `projections` — Materialised [`RunState`] projections updated
13//!   transactionally on every event append.
14
15use std::collections::HashMap;
16use std::sync::atomic::{AtomicU64, Ordering};
17
18use async_trait::async_trait;
19use tokio::sync::RwLock;
20use uuid::Uuid;
21
22use super::error::{RuntimeError, RuntimeResult};
23use super::run::{RunId, RunRecord, RunStatus};
24use super::state::RunState;
25use super::store::{RunEventRecord, RunStore};
26
27/// In-memory [`RunStore`] implementation backed by `RwLock`-protected hash maps.
28///
29/// Stores run records, event streams, and materialised [`RunState`] projections
30/// in process memory. All data is lost when the store is dropped.
31pub struct MemoryRunStore {
32    runs: RwLock<HashMap<Uuid, RunRecord>>,
33    events: RwLock<HashMap<Uuid, Vec<RunEventRecord>>>,
34    projections: RwLock<HashMap<Uuid, RunState>>,
35    sequence: AtomicU64,
36}
37
38impl MemoryRunStore {
39    /// Creates a new in-memory run store.
40    #[must_use]
41    pub fn new() -> Self {
42        Self {
43            runs: RwLock::new(HashMap::new()),
44            events: RwLock::new(HashMap::new()),
45            projections: RwLock::new(HashMap::new()),
46            sequence: AtomicU64::new(0),
47        }
48    }
49}
50
51impl Default for MemoryRunStore {
52    fn default() -> Self {
53        Self::new()
54    }
55}
56
57#[async_trait]
58impl RunStore for MemoryRunStore {
59    async fn create_run(&self, record: RunRecord) -> RuntimeResult<()> {
60        let id = *record.id.as_uuid();
61        let initial_state = RunState::create(&record, &[]);
62        self.runs.write().await.insert(id, record);
63        self.projections.write().await.insert(id, initial_state);
64        Ok(())
65    }
66
67    async fn get_run(&self, run_id: RunId) -> RuntimeResult<Option<RunRecord>> {
68        Ok(self.runs.read().await.get(run_id.as_uuid()).cloned())
69    }
70
71    async fn get_run_state(&self, run_id: RunId) -> RuntimeResult<Option<RunState>> {
72        Ok(self.projections.read().await.get(run_id.as_uuid()).cloned())
73    }
74
75    async fn update_run_status(&self, run_id: RunId, status: RunStatus) -> RuntimeResult<()> {
76        let mut runs = self.runs.write().await;
77        let record = runs
78            .get_mut(run_id.as_uuid())
79            .ok_or(RuntimeError::RunNotFound(run_id))?;
80        record.update_status(status);
81
82        let mut projections = self.projections.write().await;
83        if let Some(state) = projections.get_mut(run_id.as_uuid()) {
84            state.status = status;
85            state.updated_at = chrono::Utc::now();
86        }
87        Ok(())
88    }
89
90    async fn append_event(&self, mut record: RunEventRecord) -> RuntimeResult<()> {
91        record.sequence = self.sequence.fetch_add(1, Ordering::SeqCst);
92        let run_id_uuid = *record.run_id.as_uuid();
93
94        let mut events = self.events.write().await;
95        let mut projections = self.projections.write().await;
96
97        events.entry(run_id_uuid).or_default().push(record.clone());
98
99        if let Some(state) = projections.get_mut(&run_id_uuid) {
100            state.apply(&record.event);
101            state.event_count += 1;
102            state.updated_at = record.timestamp;
103        } else {
104            let runs = self.runs.read().await;
105            if let Some(record_val) = runs.get(&run_id_uuid) {
106                let mut state = RunState::create(record_val, &[]);
107                state.apply(&record.event);
108                state.event_count = 1;
109                state.updated_at = record.timestamp;
110                projections.insert(run_id_uuid, state);
111            }
112        }
113        Ok(())
114    }
115
116    async fn list_events(&self, run_id: RunId) -> RuntimeResult<Vec<RunEventRecord>> {
117        Ok(self
118            .events
119            .read()
120            .await
121            .get(run_id.as_uuid())
122            .cloned()
123            .unwrap_or_default())
124    }
125
126    async fn list_runs(&self, session_id: Uuid) -> RuntimeResult<Vec<RunRecord>> {
127        Ok(self
128            .runs
129            .read()
130            .await
131            .values()
132            .filter(|r| r.session_id == session_id)
133            .cloned()
134            .collect())
135    }
136
137    async fn list_runs_filtered(
138        &self,
139        session_id: Option<Uuid>,
140        status: Option<RunStatus>,
141        limit: usize,
142        offset: usize,
143    ) -> RuntimeResult<Vec<RunRecord>> {
144        let runs = self.runs.read().await;
145        let mut result: Vec<RunRecord> = runs
146            .values()
147            .filter(|r| {
148                if let Some(sid) = session_id
149                    && r.session_id != sid
150                {
151                    return false;
152                }
153                if let Some(s) = &status
154                    && r.status != *s
155                {
156                    return false;
157                }
158                true
159            })
160            .cloned()
161            .collect();
162        result.sort_by_key(|r| std::cmp::Reverse(r.created_at));
163        Ok(result
164            .into_iter()
165            .skip(offset)
166            .take(limit.clamp(1, 1000))
167            .collect())
168    }
169
170    async fn delete_run(&self, run_id: RunId) -> RuntimeResult<()> {
171        self.runs.write().await.remove(run_id.as_uuid());
172        self.events.write().await.remove(run_id.as_uuid());
173        self.projections.write().await.remove(run_id.as_uuid());
174        Ok(())
175    }
176
177    async fn health_check(&self) -> RuntimeResult<()> {
178        Ok(())
179    }
180}
181
182#[cfg(test)]
183#[allow(clippy::unwrap_used, clippy::expect_used)]
184mod tests {
185    use super::*;
186    use crate::event::{AgentEvent, RunStarted as RunStartedEvent, UsageRecorded};
187    use behest_provider::{ModelName, ProviderId};
188    use serde_json::Value;
189
190    fn make_run(session_id: Uuid, provider: &str, model: &str) -> RunRecord {
191        RunRecord::new(
192            RunId::new(),
193            session_id,
194            ProviderId::new(provider),
195            ModelName::new(model),
196            Value::Null,
197            None,
198        )
199    }
200
201    #[tokio::test]
202    async fn memory_run_store_should_create_and_get() {
203        let store = MemoryRunStore::new();
204        let session_id = Uuid::new_v4();
205        let record = make_run(session_id, "test", "gpt-4");
206        let run_id = record.id;
207
208        store.create_run(record).await.unwrap();
209        let fetched = store.get_run(run_id).await.unwrap();
210        assert!(fetched.is_some());
211    }
212
213    #[tokio::test]
214    async fn memory_run_store_should_update_status() {
215        let store = MemoryRunStore::new();
216        let session_id = Uuid::new_v4();
217        let record = make_run(session_id, "test", "gpt-4");
218        let run_id = record.id;
219
220        store.create_run(record).await.unwrap();
221        store
222            .update_run_status(run_id, RunStatus::Completed)
223            .await
224            .unwrap();
225
226        let fetched = store.get_run(run_id).await.unwrap().unwrap();
227        assert_eq!(fetched.status, RunStatus::Completed);
228    }
229
230    #[tokio::test]
231    async fn memory_run_store_should_append_and_list_events() {
232        let store = MemoryRunStore::new();
233        let session_id = Uuid::new_v4();
234        let record = make_run(session_id, "test", "gpt-4");
235        let run_id = record.id;
236        store.create_run(record).await.unwrap();
237
238        let event = RunEventRecord::new(
239            0,
240            run_id,
241            AgentEvent::RunStarted(RunStartedEvent {
242                run_id,
243                session_id,
244                provider: ProviderId::new("test"),
245                model: ModelName::new("gpt-4"),
246                timestamp: chrono::Utc::now(),
247            }),
248        );
249        store.append_event(event).await.unwrap();
250
251        let events = store.list_events(run_id).await.unwrap();
252        assert_eq!(events.len(), 1);
253    }
254
255    #[tokio::test]
256    async fn memory_run_store_should_list_by_session() {
257        let store = MemoryRunStore::new();
258        let session_id = Uuid::new_v4();
259
260        let r1 = make_run(session_id, "a", "m1");
261        let r2 = make_run(session_id, "b", "m2");
262        let r3 = make_run(Uuid::new_v4(), "c", "m3");
263
264        store.create_run(r1).await.unwrap();
265        store.create_run(r2).await.unwrap();
266        store.create_run(r3).await.unwrap();
267
268        let runs = store.list_runs(session_id).await.unwrap();
269        assert_eq!(runs.len(), 2);
270    }
271
272    #[tokio::test]
273    async fn memory_run_store_should_delete() {
274        let store = MemoryRunStore::new();
275        let session_id = Uuid::new_v4();
276        let record = make_run(session_id, "test", "m");
277        let run_id = record.id;
278
279        store.create_run(record).await.unwrap();
280        store.delete_run(run_id).await.unwrap();
281
282        let fetched = store.get_run(run_id).await.unwrap();
283        assert!(fetched.is_none());
284    }
285
286    #[tokio::test]
287    async fn memory_run_store_should_maintain_transactional_projection() {
288        let store = MemoryRunStore::new();
289        let session_id = Uuid::new_v4();
290        let record = make_run(session_id, "test", "gpt-4");
291        let run_id = record.id;
292        store.create_run(record).await.unwrap();
293
294        // 1. Check initial projection is Pending
295        let state = store.get_run_state(run_id).await.unwrap().unwrap();
296        assert_eq!(state.status, RunStatus::Pending);
297        assert_eq!(state.event_count, 0);
298
299        // 2. Append RunStarted event and check projection
300        let event1 = RunEventRecord::new(
301            0,
302            run_id,
303            AgentEvent::RunStarted(RunStartedEvent {
304                run_id,
305                session_id,
306                provider: ProviderId::new("test"),
307                model: ModelName::new("gpt-4"),
308                timestamp: chrono::Utc::now(),
309            }),
310        );
311        store.append_event(event1).await.unwrap();
312
313        let state = store.get_run_state(run_id).await.unwrap().unwrap();
314        assert_eq!(state.status, RunStatus::SessionLoaded);
315        assert_eq!(state.event_count, 1);
316
317        // 3. Append UsageRecorded event and check projection accumulates usage
318        let event2 = RunEventRecord::new(
319            0,
320            run_id,
321            AgentEvent::UsageRecorded(UsageRecorded {
322                run_id,
323                usage: behest_provider::TokenUsage::new(100, 200),
324                timestamp: chrono::Utc::now(),
325            }),
326        );
327        store.append_event(event2).await.unwrap();
328
329        let state = store.get_run_state(run_id).await.unwrap().unwrap();
330        assert_eq!(state.total_usage.input_tokens, 100);
331        assert_eq!(state.total_usage.output_tokens, 200);
332        assert_eq!(state.event_count, 2);
333    }
334}