behest_store/memory/
session.rs1use std::collections::HashMap;
4
5use async_trait::async_trait;
6use tokio::sync::RwLock;
7use uuid::Uuid;
8
9use crate::{MessageRecord, Session, SessionStore, StoreResult};
10use behest_provider::{ModelName, TokenUsage};
11
12#[derive(Default)]
20pub struct MemorySessionStore {
21 sessions: RwLock<HashMap<Uuid, Session>>,
22 messages: RwLock<HashMap<Uuid, Vec<MessageRecord>>>,
23 message_index: RwLock<HashMap<Uuid, Uuid>>,
24}
25
26impl MemorySessionStore {
27 #[must_use]
29 pub fn new() -> Self {
30 Self::default()
31 }
32}
33
34#[async_trait]
35impl SessionStore for MemorySessionStore {
36 async fn create_session(&self, session: Session) -> StoreResult<Session> {
37 let mut sessions = self.sessions.write().await;
38 let id = session.id;
39 sessions.insert(id, session.clone());
40 self.messages.write().await.insert(id, Vec::new());
41 Ok(session)
42 }
43
44 async fn list_sessions(&self) -> StoreResult<Vec<Session>> {
45 let sessions = self.sessions.read().await;
46 let mut result: Vec<Session> = sessions.values().cloned().collect();
47 result.sort_by_key(|s| std::cmp::Reverse(s.updated_at));
48 Ok(result)
49 }
50
51 async fn get_session(&self, id: &Uuid) -> StoreResult<Option<Session>> {
52 let sessions = self.sessions.read().await;
53 Ok(sessions.get(id).cloned())
54 }
55
56 async fn delete_session(&self, id: &Uuid) -> StoreResult<()> {
57 self.sessions.write().await.remove(id);
58 if let Some(records) = self.messages.write().await.remove(id) {
60 let mut index = self.message_index.write().await;
61 for record in &records {
62 index.remove(&record.id);
63 }
64 }
65 Ok(())
66 }
67
68 async fn update_session(
69 &self,
70 id: &Uuid,
71 title: &str,
72 model: Option<&ModelName>,
73 ) -> StoreResult<Session> {
74 let mut sessions = self.sessions.write().await;
75 let session = sessions
76 .get_mut(id)
77 .ok_or_else(|| behest_core::error::StorageError::NotFound { id: id.to_string() })?;
78 title.clone_into(&mut session.title);
79 session.updated_at = chrono::Utc::now();
80 if let Some(m) = model {
81 session.model = m.clone();
82 }
83 Ok(session.clone())
84 }
85
86 async fn append_message(&self, message: MessageRecord) -> StoreResult<MessageRecord> {
87 let session_id = message.session_id;
88
89 {
91 let mut sessions = self.sessions.write().await;
92 let session = sessions.get_mut(&session_id).ok_or_else(|| {
93 behest_core::error::StorageError::NotFound {
94 id: session_id.to_string(),
95 }
96 })?;
97 session.updated_at = chrono::Utc::now();
98 }
99
100 self.messages
102 .write()
103 .await
104 .entry(session_id)
105 .or_default()
106 .push(message.clone());
107
108 self.message_index
110 .write()
111 .await
112 .insert(message.id, session_id);
113
114 Ok(message)
115 }
116
117 async fn list_messages(&self, session_id: &Uuid) -> StoreResult<Vec<MessageRecord>> {
118 let messages = self.messages.read().await;
119 Ok(messages.get(session_id).cloned().unwrap_or_default())
120 }
121
122 async fn update_usage(&self, message_id: &Uuid, usage: TokenUsage) -> StoreResult<()> {
123 let session_id = {
125 let index = self.message_index.read().await;
126 index.get(message_id).copied().ok_or_else(|| {
127 behest_core::error::StorageError::NotFound {
128 id: message_id.to_string(),
129 }
130 })?
131 };
132
133 let mut messages = self.messages.write().await;
134 if let Some(records) = messages.get_mut(&session_id) {
135 for record in records.iter_mut() {
136 if record.id == *message_id {
137 record.usage = Some(usage);
138 return Ok(());
139 }
140 }
141 }
142 Err(behest_core::error::StorageError::NotFound {
143 id: message_id.to_string(),
144 })
145 }
146}
147
148#[cfg(test)]
149#[allow(clippy::unwrap_used)]
150mod tests {
151 use super::*;
152 use behest_provider::{ContentPart, ModelName};
153
154 fn test_session() -> Session {
155 Session::new("Test Chat", ModelName::new("gpt-4"))
156 }
157
158 #[tokio::test]
159 async fn memory_session_store_should_create_and_get_session() {
160 let store = MemorySessionStore::new();
161 let session = test_session();
162 let id = session.id;
163
164 store.create_session(session).await.unwrap();
165 let loaded = store.get_session(&id).await.unwrap();
166
167 assert!(loaded.is_some());
168 assert_eq!(loaded.unwrap().title, "Test Chat");
169 }
170
171 #[tokio::test]
172 async fn memory_session_store_should_list_sessions_by_updated_at() {
173 let store = MemorySessionStore::new();
174
175 let s1 = test_session();
176 let s2 = test_session();
177 store.create_session(s1).await.unwrap();
178 store.create_session(s2).await.unwrap();
179
180 let sessions = store.list_sessions().await.unwrap();
181 assert_eq!(sessions.len(), 2);
182 }
183
184 #[tokio::test]
185 async fn memory_session_store_should_delete_session_and_messages() {
186 let store = MemorySessionStore::new();
187 let session = test_session();
188 let id = session.id;
189
190 store.create_session(session).await.unwrap();
191 store.delete_session(&id).await.unwrap();
192
193 assert!(store.get_session(&id).await.unwrap().is_none());
194 assert!(store.list_messages(&id).await.unwrap().is_empty());
195 }
196
197 #[tokio::test]
198 async fn memory_session_store_should_append_and_list_messages() {
199 let store = MemorySessionStore::new();
200 let session = test_session();
201 let session_id = session.id;
202 store.create_session(session).await.unwrap();
203
204 let msg1 = MessageRecord::new(
205 session_id,
206 crate::MessageRole::User,
207 vec![ContentPart::text("Hello")],
208 );
209 let msg2 = MessageRecord::new(
210 session_id,
211 crate::MessageRole::Assistant,
212 vec![ContentPart::text("Hi there!")],
213 );
214
215 store.append_message(msg1).await.unwrap();
216 store.append_message(msg2).await.unwrap();
217
218 let messages = store.list_messages(&session_id).await.unwrap();
219 assert_eq!(messages.len(), 2);
220 assert_eq!(messages[0].role, crate::MessageRole::User);
221 assert_eq!(messages[1].role, crate::MessageRole::Assistant);
222 }
223
224 #[tokio::test]
225 async fn memory_session_store_should_reject_message_for_missing_session() {
226 let store = MemorySessionStore::new();
227 let msg = MessageRecord::new(
228 Uuid::now_v7(),
229 crate::MessageRole::User,
230 vec![ContentPart::text("Hello")],
231 );
232
233 let result = store.append_message(msg).await;
234 assert!(result.is_err());
235 }
236
237 #[tokio::test]
238 async fn memory_session_store_should_update_usage() {
239 let store = MemorySessionStore::new();
240 let session = test_session();
241 let session_id = session.id;
242 store.create_session(session).await.unwrap();
243
244 let msg = MessageRecord::new(
245 session_id,
246 crate::MessageRole::Assistant,
247 vec![ContentPart::text("response")],
248 );
249 let msg = store.append_message(msg).await.unwrap();
250
251 let usage = TokenUsage::new(10, 20);
252 store.update_usage(&msg.id, usage).await.unwrap();
253
254 let messages = store.list_messages(&session_id).await.unwrap();
255 assert_eq!(messages[0].usage.unwrap().input_tokens, 10);
256 assert_eq!(messages[0].usage.unwrap().output_tokens, 20);
257 }
258
259 #[tokio::test]
260 async fn memory_session_store_should_return_not_found_for_unknown_usage() {
261 let store = MemorySessionStore::new();
262 let result = store
263 .update_usage(&Uuid::now_v7(), TokenUsage::new(1, 1))
264 .await;
265 assert!(result.is_err());
266 }
267
268 #[tokio::test]
269 async fn memory_session_store_should_return_none_for_unknown_session() {
270 let store = MemorySessionStore::new();
271 let result = store.get_session(&Uuid::now_v7()).await.unwrap();
272 assert!(result.is_none());
273 }
274}