Skip to main content

klieo_memory_sqlite/
short_term.rs

1//! `SqliteShortTerm` — `ShortTermMemory` over a SQLite table.
2
3use crate::connection::DbHandle;
4use async_trait::async_trait;
5use chrono::Utc;
6use klieo_core::error::MemoryError;
7use klieo_core::ids::ThreadId;
8use klieo_core::llm::{Message, Role, ToolCall};
9use klieo_core::memory::ShortTermMemory;
10
11/// SQLite-backed short-term conversation memory.
12pub struct SqliteShortTerm {
13    db: DbHandle,
14}
15
16impl SqliteShortTerm {
17    pub(crate) fn new(db: DbHandle) -> Self {
18        Self { db }
19    }
20}
21
22fn role_to_str(r: Role) -> &'static str {
23    match r {
24        Role::System => "system",
25        Role::User => "user",
26        Role::Assistant => "assistant",
27        Role::Tool => "tool",
28    }
29}
30
31fn role_from_str(s: &str) -> Result<Role, MemoryError> {
32    match s {
33        "system" => Ok(Role::System),
34        "user" => Ok(Role::User),
35        "assistant" => Ok(Role::Assistant),
36        "tool" => Ok(Role::Tool),
37        other => Err(MemoryError::Serialization(format!("unknown role: {other}"))),
38    }
39}
40
41#[async_trait]
42impl ShortTermMemory for SqliteShortTerm {
43    async fn append(&self, thread: ThreadId, msg: Message) -> Result<(), MemoryError> {
44        let tool_calls_json = serde_json::to_string(&msg.tool_calls)
45            .map_err(|e| MemoryError::Serialization(e.to_string()))?;
46        let role = role_to_str(msg.role);
47        let now = Utc::now().to_rfc3339();
48        self.db
49            .execute(move |conn| {
50                let tx = conn.transaction()?;
51                let next_seq: i64 = tx
52                    .query_row(
53                        "SELECT COALESCE(MAX(seq), 0) + 1 FROM short_term_messages WHERE thread_id = ?1",
54                        rusqlite::params![&thread.0],
55                        |r| r.get(0),
56                    )?;
57                tx.execute(
58                    "INSERT INTO short_term_messages (thread_id, seq, role, content, tool_calls, tool_call_id, created_at) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
59                    rusqlite::params![
60                        &thread.0,
61                        next_seq,
62                        role,
63                        &msg.content,
64                        &tool_calls_json,
65                        &msg.tool_call_id,
66                        &now,
67                    ],
68                )?;
69                tx.commit()?;
70                Ok(())
71            })
72            .await
73    }
74
75    async fn load(&self, thread: ThreadId, max_tokens: usize) -> Result<Vec<Message>, MemoryError> {
76        let rows: Vec<(String, String, String, Option<String>)> = self
77            .db
78            .execute(move |conn| {
79                let mut stmt = conn.prepare(
80                    "SELECT role, content, tool_calls, tool_call_id FROM short_term_messages WHERE thread_id = ?1 ORDER BY seq ASC",
81                )?;
82                let iter = stmt.query_map(rusqlite::params![&thread.0], |row| {
83                    Ok((
84                        row.get::<_, String>(0)?,
85                        row.get::<_, String>(1)?,
86                        row.get::<_, String>(2)?,
87                        row.get::<_, Option<String>>(3)?,
88                    ))
89                })?;
90                iter.collect::<Result<Vec<_>, _>>()
91            })
92            .await?;
93
94        let mut messages: Vec<Message> = rows
95            .into_iter()
96            .map(|(role, content, tool_calls_json, tool_call_id)| {
97                let tool_calls: Vec<ToolCall> = serde_json::from_str(&tool_calls_json)
98                    .map_err(|e| MemoryError::Serialization(e.to_string()))?;
99                Ok(Message {
100                    role: role_from_str(&role)?,
101                    content,
102                    tool_calls,
103                    tool_call_id,
104                })
105            })
106            .collect::<Result<Vec<_>, MemoryError>>()?;
107
108        // Approximate token-budget truncation: ~4 chars per token. Walk
109        // from newest to oldest accumulating cost; split off everything
110        // older than the kept-suffix boundary in O(n).
111        let total: usize = messages.iter().map(|m| (m.content.len() / 4).max(1)).sum();
112        if total > max_tokens {
113            let mut kept = 0usize;
114            let mut keep_from = messages.len();
115            for (idx, msg) in messages.iter().enumerate().rev() {
116                let cost = (msg.content.len() / 4).max(1);
117                if kept + cost > max_tokens {
118                    break;
119                }
120                kept += cost;
121                keep_from = idx;
122            }
123            messages = messages.split_off(keep_from);
124        }
125        Ok(messages)
126    }
127
128    async fn clear(&self, thread: ThreadId) -> Result<(), MemoryError> {
129        self.db
130            .execute(move |conn| {
131                conn.execute(
132                    "DELETE FROM short_term_messages WHERE thread_id = ?1",
133                    rusqlite::params![&thread.0],
134                )?;
135                Ok(())
136            })
137            .await
138    }
139}
140
141#[cfg(test)]
142mod tests {
143    use super::*;
144    use crate::connection::DbHandle;
145
146    fn user(text: &str) -> Message {
147        Message {
148            role: Role::User,
149            content: text.into(),
150            tool_calls: vec![],
151            tool_call_id: None,
152        }
153    }
154
155    async fn fresh() -> SqliteShortTerm {
156        let db = DbHandle::open(":memory:").await.unwrap();
157        SqliteShortTerm::new(db)
158    }
159
160    #[tokio::test]
161    async fn append_then_load_round_trips() {
162        let m = fresh().await;
163        let t = ThreadId::new("t1");
164        m.append(t.clone(), user("hello")).await.unwrap();
165        m.append(t.clone(), user("world")).await.unwrap();
166        let loaded = m.load(t, 10_000).await.unwrap();
167        assert_eq!(loaded.len(), 2);
168        assert_eq!(loaded[0].content, "hello");
169        assert_eq!(loaded[1].content, "world");
170    }
171
172    #[tokio::test]
173    async fn load_truncates_to_token_budget() {
174        let m = fresh().await;
175        let t = ThreadId::new("t1");
176        // Each ~40-char message ~= 10 tokens.
177        for i in 0..20 {
178            m.append(
179                t.clone(),
180                user(&format!("msg-{i:03}-padding-padding-padding")),
181            )
182            .await
183            .unwrap();
184        }
185        let loaded = m.load(t, 30).await.unwrap();
186        // Should keep ~3 messages (30 tokens / 10 each), oldest dropped.
187        assert!(
188            loaded.len() <= 4 && !loaded.is_empty(),
189            "expected ~3 messages, got {}",
190            loaded.len()
191        );
192        // Newest still present.
193        assert!(
194            loaded.last().unwrap().content.contains("msg-019"),
195            "newest message must survive truncation"
196        );
197    }
198
199    #[tokio::test]
200    async fn clear_removes_thread() {
201        let m = fresh().await;
202        let t = ThreadId::new("t1");
203        m.append(t.clone(), user("hello")).await.unwrap();
204        m.clear(t.clone()).await.unwrap();
205        let loaded = m.load(t, 10_000).await.unwrap();
206        assert!(loaded.is_empty());
207    }
208
209    #[tokio::test]
210    async fn threads_are_isolated() {
211        let m = fresh().await;
212        m.append(ThreadId::new("a"), user("a-msg")).await.unwrap();
213        m.append(ThreadId::new("b"), user("b-msg")).await.unwrap();
214        let a = m.load(ThreadId::new("a"), 10_000).await.unwrap();
215        let b = m.load(ThreadId::new("b"), 10_000).await.unwrap();
216        assert_eq!(a.len(), 1);
217        assert_eq!(b.len(), 1);
218        assert_eq!(a[0].content, "a-msg");
219        assert_eq!(b[0].content, "b-msg");
220    }
221
222    #[tokio::test]
223    async fn role_round_trip_for_all_variants() {
224        let m = fresh().await;
225        let t = ThreadId::new("r");
226        for role in [Role::System, Role::User, Role::Assistant, Role::Tool] {
227            m.append(
228                t.clone(),
229                Message {
230                    role,
231                    content: "x".into(),
232                    tool_calls: vec![],
233                    tool_call_id: None,
234                },
235            )
236            .await
237            .unwrap();
238        }
239        let loaded = m.load(t, 10_000).await.unwrap();
240        assert_eq!(loaded.len(), 4);
241        assert_eq!(loaded[0].role, Role::System);
242        assert_eq!(loaded[3].role, Role::Tool);
243    }
244}