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