1use 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
11pub 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 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 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 assert!(
189 loaded.len() <= 4 && !loaded.is_empty(),
190 "expected ~3 messages, got {}",
191 loaded.len()
192 );
193 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}