Skip to main content

behest_store/memory/
session.rs

1//! In-memory session store backed by `HashMap`s for sessions and messages.
2
3use 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/// In-memory session store for testing, development, and prototyping.
13///
14/// Sessions are stored in `HashMap<Uuid, Session>` and messages in
15/// `HashMap<Uuid, Vec<MessageRecord>>`, both protected by `RwLock`.
16/// A reverse index `message_id → session_id` enables O(1) message lookups
17/// in [`update_usage`](SessionStore::update_usage). Data is lost when the
18/// process exits. Implements [`SessionStore`].
19#[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    /// Creates an empty in-memory session store.
28    #[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        // Clean up the reverse message index for all messages in this session.
59        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        // Update session timestamp first (acquire sessions lock, release)
90        {
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        // Append message (acquire messages lock, release)
101        self.messages
102            .write()
103            .await
104            .entry(session_id)
105            .or_default()
106            .push(message.clone());
107
108        // Maintain the message_id → session_id reverse index for O(1) update_usage.
109        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        // O(1) lookup via the message_id → session_id reverse index.
124        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}