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