1use 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
27pub 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 #[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 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 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 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}