1use std::str::FromStr;
4
5use async_trait::async_trait;
6use sqlx::sqlite::{SqliteConnectOptions, SqlitePoolOptions};
7use sqlx::SqlitePool;
8
9use crate::ledger::{LedgerError, LedgerStore, UsageEntry};
10
11pub struct SqliteLedger {
12 pool: SqlitePool,
13}
14
15impl SqliteLedger {
16 pub async fn connect(dsn: &str) -> Result<Self, LedgerError> {
24 let opts = SqliteConnectOptions::from_str(dsn)
25 .map_err(|e| LedgerError::Backend(e.to_string()))?
26 .create_if_missing(true);
27
28 let pool = SqlitePoolOptions::new()
29 .max_connections(1)
30 .connect_with(opts)
31 .await
32 .map_err(|e| LedgerError::Backend(e.to_string()))?;
33
34 sqlx::raw_sql(include_str!("../../migrations/0001_usage_events.sql"))
37 .execute(&pool)
38 .await
39 .map_err(|e| LedgerError::Backend(e.to_string()))?;
40
41 let _ = sqlx::query("ALTER TABLE usage_events ADD COLUMN user_id TEXT")
44 .execute(&pool)
45 .await;
46 let _ = sqlx::query("ALTER TABLE usage_events ADD COLUMN thread_id TEXT")
47 .execute(&pool)
48 .await;
49 let _ = sqlx::query("ALTER TABLE usage_events ADD COLUMN message_id TEXT")
50 .execute(&pool)
51 .await;
52 let _ = sqlx::query("ALTER TABLE usage_events ADD COLUMN user_task_type TEXT")
53 .execute(&pool)
54 .await;
55 let _ = sqlx::query("ALTER TABLE usage_events ADD COLUMN ai_task_type TEXT")
56 .execute(&pool)
57 .await;
58
59 Ok(Self { pool })
60 }
61}
62
63#[async_trait]
64impl LedgerStore for SqliteLedger {
65 async fn record(&self, e: &UsageEntry) -> Result<(), LedgerError> {
66 sqlx::query(
67 "INSERT INTO usage_events \
68 (ts, tenant, workspace, user_id, thread_id, message_id, route, provider, model, lane, \
69 input_tokens, output_tokens, cost_usd, request_id, status, user_task_type, ai_task_type) \
70 VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)",
71 )
72 .bind(e.ts.to_rfc3339())
73 .bind(&e.tenant)
74 .bind(&e.workspace)
75 .bind(&e.user)
76 .bind(&e.thread)
77 .bind(&e.message)
78 .bind(&e.route)
79 .bind(&e.provider)
80 .bind(&e.model)
81 .bind(&e.lane)
82 .bind(e.input_tokens as i64)
83 .bind(e.output_tokens as i64)
84 .bind(e.cost_usd)
85 .bind(&e.request_id)
86 .bind(&e.status)
87 .bind(&e.user_task_type)
88 .bind(&e.ai_task_type)
89 .execute(&self.pool)
90 .await
91 .map_err(|e| LedgerError::Backend(e.to_string()))?;
92 Ok(())
93 }
94}
95
96#[cfg(test)]
97mod tests {
98 use super::*;
99 use chrono::Utc;
100 use sqlx::Row;
101
102 fn entry() -> UsageEntry {
103 UsageEntry {
104 ts: Utc::now(),
105 tenant: "acme".into(),
106 workspace: None,
107 user: None,
108 thread: None,
109 message: None,
110 route: "fast".into(),
111 provider: "vertex".into(),
112 model: "gemini-3-flash".into(),
113 lane: "standard".into(),
114 input_tokens: 3,
115 output_tokens: 5,
116 cost_usd: 0.001,
117 request_id: "req-1".into(),
118 status: "ok".into(),
119 op: "chat".into(),
120 user_task_type: None,
121 ai_task_type: "simple".into(),
122 }
123 }
124
125 async fn stored_user_task_type(e: &UsageEntry) -> Option<String> {
126 stored_column(e, "user_task_type").await
127 }
128
129 async fn stored_column(e: &UsageEntry, column: &str) -> Option<String> {
130 let store = SqliteLedger::connect("sqlite::memory:").await.unwrap();
131 store.record(e).await.unwrap();
132 sqlx::query(&format!("SELECT {column} FROM usage_events"))
133 .fetch_one(&store.pool)
134 .await
135 .unwrap()
136 .get(column)
137 }
138
139 #[tokio::test]
140 async fn persists_ai_task_type_column() {
141 let e = UsageEntry {
142 ai_task_type: "conversation".into(),
143 ..entry()
144 };
145 assert_eq!(
146 stored_column(&e, "ai_task_type").await,
147 Some("conversation".to_string())
148 );
149 }
150
151 #[tokio::test]
152 async fn persists_user_task_type_column() {
153 let e = UsageEntry {
154 user_task_type: Some("summarisation".into()),
155 ..entry()
156 };
157 assert_eq!(
158 stored_user_task_type(&e).await,
159 Some("summarisation".to_string())
160 );
161 }
162
163 #[tokio::test]
164 async fn stores_null_user_task_type_when_absent() {
165 assert_eq!(stored_user_task_type(&entry()).await, None);
166 }
167
168 #[tokio::test]
171 async fn backfills_column_on_a_pre_existing_database() {
172 let path = std::env::temp_dir().join(format!("synapse-{}.db", uuid::Uuid::new_v4()));
173 let dsn = format!("sqlite://{}?mode=rwc", path.display());
174
175 let legacy = SqlitePoolOptions::new()
176 .max_connections(1)
177 .connect_with(
178 SqliteConnectOptions::from_str(&dsn)
179 .unwrap()
180 .create_if_missing(true),
181 )
182 .await
183 .unwrap();
184 sqlx::raw_sql(
185 "CREATE TABLE usage_events (\
186 id INTEGER PRIMARY KEY AUTOINCREMENT, ts TEXT NOT NULL, tenant TEXT NOT NULL, \
187 workspace TEXT, route TEXT NOT NULL, provider TEXT NOT NULL, model TEXT NOT NULL, \
188 lane TEXT NOT NULL, input_tokens INTEGER NOT NULL, output_tokens INTEGER NOT NULL, \
189 cost_usd REAL NOT NULL, request_id TEXT NOT NULL, status TEXT NOT NULL)",
190 )
191 .execute(&legacy)
192 .await
193 .unwrap();
194 legacy.close().await;
195
196 let store = SqliteLedger::connect(&dsn).await.unwrap();
197 let e = UsageEntry {
198 user_task_type: Some("summarisation".into()),
199 ai_task_type: "conversation".into(),
200 ..entry()
201 };
202 store.record(&e).await.unwrap();
203 let row = sqlx::query("SELECT user_task_type, ai_task_type FROM usage_events")
204 .fetch_one(&store.pool)
205 .await
206 .unwrap();
207 assert_eq!(
208 row.get::<Option<String>, _>("user_task_type"),
209 Some("summarisation".to_string())
210 );
211 assert_eq!(
212 row.get::<Option<String>, _>("ai_task_type"),
213 Some("conversation".to_string())
214 );
215
216 let _ = std::fs::remove_file(&path);
217 }
218}