1use std::collections::{BTreeMap, HashMap};
2use std::path::Path;
3use std::sync::{Arc, Mutex};
4
5use rusqlite::{params, Connection, OptionalExtension, TransactionBehavior};
6use serde::ser::{SerializeMap, SerializeSeq};
7use serde::{Deserialize, Deserializer, Serialize, Serializer};
8use serde_json::Value;
9use sha2::{Digest, Sha256};
10
11use crate::types::{Message, MessageRole, ToolArtifactRef, ToolCall};
12
13mod conformance;
14mod redis_store;
15
16pub use conformance::session_store_conformance;
17pub use redis_store::RedisSessionStore;
18
19pub trait Session: Send + Sync {
20 fn session_id(&self) -> &str;
21 fn get_items(&self, limit: Option<usize>) -> SessionFuture<Vec<SessionItem>>;
22 fn add_items(&self, items: Vec<SessionItem>) -> SessionFuture<()>;
23 fn pop_item(&self) -> SessionFuture<Option<SessionItem>>;
24 fn clear(&self) -> SessionFuture<()>;
25
26 fn supports_add_items_once(&self) -> bool {
27 false
28 }
29
30 fn add_items_once(
31 &self,
32 _commit_id: String,
33 _payload_digest: String,
34 _items: Vec<SessionItem>,
35 ) -> SessionFuture<SessionAppendOutcome> {
36 Box::pin(async {
37 Err("checkpoint_session_idempotency_unsupported: session does not support add_items_once"
38 .to_string())
39 })
40 }
41
42 fn clear_session(&self) -> SessionFuture<()> {
43 self.clear()
44 }
45}
46pub type SessionFuture<T> =
47 std::pin::Pin<Box<dyn std::future::Future<Output = Result<T, String>> + Send>>;
48
49#[derive(Debug, Clone, Copy, PartialEq, Eq)]
50pub enum SessionAppendOutcome {
51 Committed,
52 Replayed,
53}
54
55pub fn checkpoint_session_commit_id(checkpoint_key: &str) -> String {
56 let digest = Sha256::digest(checkpoint_key.as_bytes());
57 format!("vv-agent:checkpoint-v2:session:{digest:x}")
58}
59
60pub fn session_commit_payload_digest(items: &[SessionItem]) -> Result<String, String> {
61 let payload = serde_json::json!({
62 "schema_version": "vv-agent.session-commit.v1",
63 "items": items,
64 });
65 let bytes = crate::checkpoint::canonical_json_bytes(&payload, "session commit payload")
66 .map_err(|error| error.to_string())?;
67 Ok(format!("{:x}", Sha256::digest(bytes)))
68}
69
70fn validate_session_commit(
71 commit_id: &str,
72 payload_digest: &str,
73 items: &[SessionItem],
74) -> Result<(), String> {
75 if commit_id.trim().is_empty() {
76 return Err("session_commit_identity_invalid: commit_id must be non-empty".to_string());
77 }
78 if payload_digest.len() != 64
79 || !payload_digest
80 .bytes()
81 .all(|byte| byte.is_ascii_digit() || (b'a'..=b'f').contains(&byte))
82 {
83 return Err(
84 "session_commit_payload_digest_invalid: payload_digest must be lowercase SHA-256"
85 .to_string(),
86 );
87 }
88 if session_commit_payload_digest(items)? != payload_digest {
89 return Err(
90 "session_commit_payload_digest_mismatch: payload_digest does not match items"
91 .to_string(),
92 );
93 }
94 Ok(())
95}
96
97#[derive(Debug, Clone, PartialEq)]
98pub enum SessionItem {
99 User {
100 content: String,
101 },
102 Assistant {
103 content: String,
104 },
105 System {
106 content: String,
107 },
108 Tool {
109 content: String,
110 tool_call_id: String,
111 },
112 Message {
113 message: Box<Message>,
114 },
115}
116
117impl SessionItem {
118 pub fn to_message(&self) -> Message {
119 match self {
120 Self::User { content } => Message::user(content.clone()),
121 Self::Assistant { content } => Message::assistant(content.clone()),
122 Self::System { content } => Message::system(content.clone()),
123 Self::Tool {
124 content,
125 tool_call_id,
126 } => Message::tool(content.clone(), tool_call_id.clone()),
127 Self::Message { message } => message.as_ref().clone(),
128 }
129 }
130
131 pub fn from_message(message: &Message) -> Option<Self> {
132 let has_unrepresentable_tool_call_id = match message.role {
133 MessageRole::Tool => message.tool_call_id.is_none(),
134 _ => message.tool_call_id.is_some(),
135 };
136 if has_unrepresentable_tool_call_id
137 || message.name.is_some()
138 || !message.tool_calls.is_empty()
139 || message.reasoning_content.is_some()
140 || message.image_url.is_some()
141 || !message.metadata.is_empty()
142 || message.artifact_ref.is_some()
143 {
144 return Some(Self::Message {
145 message: Box::new(message.clone()),
146 });
147 }
148 match message.role {
149 MessageRole::System => Some(Self::System {
150 content: message.content.clone(),
151 }),
152 MessageRole::User => Some(Self::User {
153 content: message.content.clone(),
154 }),
155 MessageRole::Assistant => Some(Self::Assistant {
156 content: message.content.clone(),
157 }),
158 MessageRole::Tool => Some(Self::Tool {
159 content: message.content.clone(),
160 tool_call_id: message.tool_call_id.clone().unwrap_or_default(),
161 }),
162 }
163 }
164}
165
166impl Serialize for SessionItem {
167 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
168 where
169 S: Serializer,
170 {
171 let message = self.to_message();
172 SessionMessageWire(&message).serialize(serializer)
173 }
174}
175
176impl<'de> Deserialize<'de> for SessionItem {
177 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
178 where
179 D: Deserializer<'de>,
180 {
181 let value = Value::deserialize(deserializer)?;
182 let wire = serde_json::from_value::<SessionMessageInput>(value)
183 .map_err(serde::de::Error::custom)?;
184 let message = wire.into_message().map_err(serde::de::Error::custom)?;
185 SessionItem::from_message(&message)
186 .ok_or_else(|| serde::de::Error::custom("unsupported session message role"))
187 }
188}
189
190#[derive(Deserialize)]
191#[serde(deny_unknown_fields)]
192struct SessionMessageInput {
193 role: String,
194 content: String,
195 #[serde(default, deserialize_with = "deserialize_present_string")]
196 name: Option<String>,
197 #[serde(default, deserialize_with = "deserialize_present_string")]
198 tool_call_id: Option<String>,
199 #[serde(default)]
200 tool_calls: Vec<SessionToolCallInput>,
201 #[serde(default, deserialize_with = "deserialize_present_string")]
202 reasoning_content: Option<String>,
203 #[serde(default, deserialize_with = "deserialize_present_string")]
204 image_url: Option<String>,
205 #[serde(default)]
206 metadata: BTreeMap<String, Value>,
207 #[serde(default, deserialize_with = "deserialize_present_artifact_ref")]
208 artifact_ref: Option<ToolArtifactRef>,
209}
210
211impl SessionMessageInput {
212 fn into_message(self) -> Result<Message, String> {
213 let role = match self.role.as_str() {
214 "system" => MessageRole::System,
215 "user" => MessageRole::User,
216 "assistant" => MessageRole::Assistant,
217 "tool" => MessageRole::Tool,
218 value => return Err(format!("unknown message role: {value}")),
219 };
220 let tool_calls = self
221 .tool_calls
222 .into_iter()
223 .map(SessionToolCallInput::into_tool_call)
224 .collect::<Result<Vec<_>, _>>()?;
225 if let Some(artifact_ref) = &self.artifact_ref {
226 artifact_ref.validate()?;
227 }
228 Ok(Message {
229 role,
230 content: self.content,
231 name: self.name,
232 tool_call_id: self.tool_call_id,
233 tool_calls,
234 reasoning_content: self.reasoning_content,
235 image_url: self.image_url,
236 metadata: self.metadata,
237 artifact_ref: self.artifact_ref,
238 })
239 }
240}
241
242#[derive(Deserialize)]
243#[serde(deny_unknown_fields)]
244struct SessionToolCallInput {
245 id: String,
246 #[serde(rename = "type")]
247 kind: String,
248 function: SessionToolFunctionInput,
249 #[serde(default, deserialize_with = "deserialize_present_object")]
250 extra_content: Option<BTreeMap<String, Value>>,
251}
252
253impl SessionToolCallInput {
254 fn into_tool_call(self) -> Result<ToolCall, String> {
255 if self.id.trim().is_empty() {
256 return Err("session tool call id must be non-empty".to_string());
257 }
258 if self.kind != "function" {
259 return Err(format!("unknown session tool call type: {}", self.kind));
260 }
261 if self.function.name.trim().is_empty() {
262 return Err("session tool function name must be non-empty".to_string());
263 }
264 let arguments = serde_json::from_str::<Value>(&self.function.arguments)
265 .map_err(|error| format!("invalid session tool arguments JSON: {error}"))?;
266 let Value::Object(arguments) = arguments else {
267 return Err("session tool arguments must decode to an object".to_string());
268 };
269 Ok(ToolCall {
270 id: self.id,
271 name: self.function.name,
272 arguments: arguments.into_iter().collect(),
273 extra_content: self
274 .extra_content
275 .map(|value| Value::Object(value.into_iter().collect())),
276 })
277 }
278}
279
280#[derive(Deserialize)]
281#[serde(deny_unknown_fields)]
282struct SessionToolFunctionInput {
283 name: String,
284 arguments: String,
285}
286
287fn deserialize_present_string<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
288where
289 D: Deserializer<'de>,
290{
291 String::deserialize(deserializer).map(Some)
292}
293
294fn deserialize_present_object<'de, D>(
295 deserializer: D,
296) -> Result<Option<BTreeMap<String, Value>>, D::Error>
297where
298 D: Deserializer<'de>,
299{
300 BTreeMap::<String, Value>::deserialize(deserializer).map(Some)
301}
302
303fn deserialize_present_artifact_ref<'de, D>(
304 deserializer: D,
305) -> Result<Option<ToolArtifactRef>, D::Error>
306where
307 D: Deserializer<'de>,
308{
309 ToolArtifactRef::deserialize(deserializer).map(Some)
310}
311
312struct SessionMessageWire<'a>(&'a Message);
313
314impl Serialize for SessionMessageWire<'_> {
315 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
316 where
317 S: Serializer,
318 {
319 let message = self.0;
320 let mut field_count = 2;
321 field_count += usize::from(message.name.is_some());
322 field_count += usize::from(message.tool_call_id.is_some());
323 field_count += usize::from(!message.tool_calls.is_empty());
324 field_count += usize::from(message.reasoning_content.is_some());
325 field_count += usize::from(message.image_url.is_some());
326 field_count += usize::from(!message.metadata.is_empty());
327 field_count += usize::from(message.artifact_ref.is_some());
328
329 let mut state = serializer.serialize_map(Some(field_count))?;
330 state.serialize_entry("role", &message.role)?;
331 state.serialize_entry("content", &message.content)?;
332 if let Some(name) = &message.name {
333 state.serialize_entry("name", name)?;
334 }
335 if let Some(tool_call_id) = &message.tool_call_id {
336 state.serialize_entry("tool_call_id", tool_call_id)?;
337 }
338 if !message.tool_calls.is_empty() {
339 state.serialize_entry("tool_calls", &SessionToolCallsWire(&message.tool_calls))?;
340 }
341 if let Some(reasoning_content) = &message.reasoning_content {
342 state.serialize_entry("reasoning_content", reasoning_content)?;
343 }
344 if let Some(image_url) = &message.image_url {
345 state.serialize_entry("image_url", image_url)?;
346 }
347 if !message.metadata.is_empty() {
348 state.serialize_entry("metadata", &message.metadata)?;
349 }
350 if let Some(artifact_ref) = &message.artifact_ref {
351 state.serialize_entry("artifact_ref", artifact_ref)?;
352 }
353 state.end()
354 }
355}
356
357struct SessionToolCallsWire<'a>(&'a [ToolCall]);
358
359impl Serialize for SessionToolCallsWire<'_> {
360 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
361 where
362 S: Serializer,
363 {
364 let mut sequence = serializer.serialize_seq(Some(self.0.len()))?;
365 for tool_call in self.0 {
366 sequence.serialize_element(&SessionToolCallWire(tool_call))?;
367 }
368 sequence.end()
369 }
370}
371
372struct SessionToolCallWire<'a>(&'a ToolCall);
373
374impl Serialize for SessionToolCallWire<'_> {
375 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
376 where
377 S: Serializer,
378 {
379 let tool_call = self.0;
380 let field_count = 3 + usize::from(tool_call.extra_content.is_some());
381 let mut state = serializer.serialize_map(Some(field_count))?;
382 state.serialize_entry("id", &tool_call.id)?;
383 state.serialize_entry("type", "function")?;
384 state.serialize_entry("function", &SessionToolFunctionWire(tool_call))?;
385 if let Some(extra_content) = &tool_call.extra_content {
386 state.serialize_entry("extra_content", extra_content)?;
387 }
388 state.end()
389 }
390}
391
392struct SessionToolFunctionWire<'a>(&'a ToolCall);
393
394impl Serialize for SessionToolFunctionWire<'_> {
395 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
396 where
397 S: Serializer,
398 {
399 let mut state = serializer.serialize_map(Some(2))?;
400 state.serialize_entry("name", &self.0.name)?;
401 let arguments =
402 serde_json::to_string(&self.0.arguments).map_err(serde::ser::Error::custom)?;
403 state.serialize_entry("arguments", &arguments)?;
404 state.end()
405 }
406}
407
408#[derive(Clone)]
409pub struct MemorySession {
410 session_id: Arc<String>,
411 items: Arc<Mutex<Vec<SessionItem>>>,
412 commits: Arc<Mutex<HashMap<String, String>>>,
413}
414
415impl MemorySession {
416 pub fn new(session_id: impl Into<String>) -> Self {
417 Self {
418 session_id: Arc::new(session_id.into()),
419 items: Arc::new(Mutex::new(Vec::new())),
420 commits: Arc::new(Mutex::new(HashMap::new())),
421 }
422 }
423}
424
425impl Session for MemorySession {
426 fn session_id(&self) -> &str {
427 self.session_id.as_str()
428 }
429
430 fn get_items(&self, limit: Option<usize>) -> SessionFuture<Vec<SessionItem>> {
431 let items = self.items.clone();
432 Box::pin(async move {
433 let items = items
434 .lock()
435 .map_err(|_| "session lock poisoned".to_string())?;
436 let values = match limit {
437 Some(limit) => items
438 .iter()
439 .rev()
440 .take(limit)
441 .cloned()
442 .collect::<Vec<_>>()
443 .into_iter()
444 .rev()
445 .collect(),
446 None => items.clone(),
447 };
448 Ok(values)
449 })
450 }
451
452 fn add_items(&self, new_items: Vec<SessionItem>) -> SessionFuture<()> {
453 let items = self.items.clone();
454 Box::pin(async move {
455 items
456 .lock()
457 .map_err(|_| "session lock poisoned".to_string())?
458 .extend(new_items);
459 Ok(())
460 })
461 }
462
463 fn supports_add_items_once(&self) -> bool {
464 true
465 }
466
467 fn add_items_once(
468 &self,
469 commit_id: String,
470 payload_digest: String,
471 new_items: Vec<SessionItem>,
472 ) -> SessionFuture<SessionAppendOutcome> {
473 let items = self.items.clone();
474 let commits = self.commits.clone();
475 Box::pin(async move {
476 validate_session_commit(&commit_id, &payload_digest, &new_items)?;
477 let mut commits = commits
478 .lock()
479 .map_err(|_| "session commit lock poisoned".to_string())?;
480 if let Some(existing) = commits.get(&commit_id) {
481 if existing != &payload_digest {
482 return Err(
483 "session_commit_identity_conflict: commit_id has a different payload"
484 .to_string(),
485 );
486 }
487 return Ok(SessionAppendOutcome::Replayed);
488 }
489 items
490 .lock()
491 .map_err(|_| "session lock poisoned".to_string())?
492 .extend(new_items);
493 commits.insert(commit_id, payload_digest);
494 Ok(SessionAppendOutcome::Committed)
495 })
496 }
497
498 fn pop_item(&self) -> SessionFuture<Option<SessionItem>> {
499 let items = self.items.clone();
500 Box::pin(async move {
501 Ok(items
502 .lock()
503 .map_err(|_| "session lock poisoned".to_string())?
504 .pop())
505 })
506 }
507
508 fn clear(&self) -> SessionFuture<()> {
509 let items = self.items.clone();
510 let commits = self.commits.clone();
511 Box::pin(async move {
512 items
513 .lock()
514 .map_err(|_| "session lock poisoned".to_string())?
515 .clear();
516 commits
517 .lock()
518 .map_err(|_| "session commit lock poisoned".to_string())?
519 .clear();
520 Ok(())
521 })
522 }
523}
524
525pub trait SessionStore: Send + Sync {
526 fn session(&self, session_id: &str) -> Arc<dyn Session>;
527}
528
529#[derive(Clone, Default)]
530pub struct MemorySessionStore {
531 sessions: Arc<Mutex<HashMap<String, Arc<dyn Session>>>>,
532}
533
534impl MemorySessionStore {
535 pub fn new() -> Self {
536 Self::default()
537 }
538
539 pub fn session(&self, session_id: &str) -> Arc<dyn Session> {
540 <Self as SessionStore>::session(self, session_id)
541 }
542}
543
544impl SessionStore for MemorySessionStore {
545 fn session(&self, session_id: &str) -> Arc<dyn Session> {
546 let mut sessions = self
547 .sessions
548 .lock()
549 .expect("memory session store lock poisoned");
550 sessions
551 .entry(session_id.to_string())
552 .or_insert_with(|| Arc::new(MemorySession::new(session_id)))
553 .clone()
554 }
555}
556
557#[derive(Clone)]
558pub struct SqliteSessionStore {
559 connection: Arc<Mutex<Connection>>,
560}
561
562const SQLITE_SESSION_SCHEMA_VERSION: i64 = 2;
563const CANONICAL_SESSION_COLUMNS: [&str; 3] = ["session_id", "item_index", "payload"];
564const CANONICAL_SESSION_COMMIT_COLUMNS: [&str; 3] = ["session_id", "commit_id", "payload_digest"];
565const CREATE_SESSION_ITEMS_TABLE: &str = r#"
566 CREATE TABLE IF NOT EXISTS session_items (
567 session_id TEXT NOT NULL,
568 item_index INTEGER PRIMARY KEY AUTOINCREMENT,
569 payload TEXT NOT NULL
570 )
571"#;
572const CREATE_SESSION_ITEMS_INDEX: &str = r#"
573 CREATE INDEX IF NOT EXISTS idx_session_items_session_id_item_index
574 ON session_items (session_id, item_index)
575"#;
576const CREATE_SESSION_COMMITS_TABLE: &str = r#"
577 CREATE TABLE IF NOT EXISTS session_commits (
578 session_id TEXT NOT NULL,
579 commit_id TEXT NOT NULL,
580 payload_digest TEXT NOT NULL,
581 PRIMARY KEY (session_id, commit_id)
582 )
583"#;
584
585impl SqliteSessionStore {
586 pub fn open_memory() -> Result<Self, String> {
587 Self::open(":memory:")
588 }
589
590 pub fn open(path: impl AsRef<Path>) -> Result<Self, String> {
591 let mut connection = Connection::open(path).map_err(sqlite_error)?;
592 connection
593 .execute_batch(
594 r#"
595 PRAGMA busy_timeout = 5000;
596 PRAGMA journal_mode=WAL;
597 "#,
598 )
599 .map_err(sqlite_error)?;
600 initialize_sqlite_session_schema(&mut connection)?;
601 Ok(Self {
602 connection: Arc::new(Mutex::new(connection)),
603 })
604 }
605
606 pub fn session(&self, session_id: &str) -> Arc<dyn Session> {
607 <Self as SessionStore>::session(self, session_id)
608 }
609}
610
611impl SessionStore for SqliteSessionStore {
612 fn session(&self, session_id: &str) -> Arc<dyn Session> {
613 Arc::new(SqliteSession {
614 session_id: Arc::new(session_id.to_string()),
615 connection: self.connection.clone(),
616 })
617 }
618}
619
620#[derive(Clone)]
621struct SqliteSession {
622 session_id: Arc<String>,
623 connection: Arc<Mutex<Connection>>,
624}
625
626impl Session for SqliteSession {
627 fn session_id(&self) -> &str {
628 self.session_id.as_str()
629 }
630
631 fn get_items(&self, limit: Option<usize>) -> SessionFuture<Vec<SessionItem>> {
632 let session_id = self.session_id.to_string();
633 let connection = self.connection.clone();
634 Box::pin(async move {
635 let connection = connection
636 .lock()
637 .map_err(|_| "sqlite session store lock poisoned".to_string())?;
638 let mut statement = if limit.is_some() {
639 connection
640 .prepare(
641 r#"
642 SELECT payload
643 FROM (
644 SELECT item_index, payload
645 FROM session_items
646 WHERE session_id = ?1
647 ORDER BY item_index DESC
648 LIMIT ?2
649 )
650 ORDER BY item_index ASC
651 "#,
652 )
653 .map_err(sqlite_error)?
654 } else {
655 connection
656 .prepare(
657 r#"
658 SELECT payload
659 FROM session_items
660 WHERE session_id = ?1
661 ORDER BY item_index ASC
662 "#,
663 )
664 .map_err(sqlite_error)?
665 };
666 let mut rows = if let Some(limit) = limit {
667 statement
668 .query(params![
669 session_id,
670 i64::try_from(limit).unwrap_or(i64::MAX)
671 ])
672 .map_err(sqlite_error)?
673 } else {
674 statement.query(params![session_id]).map_err(sqlite_error)?
675 };
676 let mut items = Vec::new();
677 while let Some(row) = rows.next().map_err(sqlite_error)? {
678 let payload: String = row.get(0).map_err(sqlite_error)?;
679 items.push(serde_json::from_str(&payload).map_err(json_error)?);
680 }
681 Ok(items)
682 })
683 }
684
685 fn add_items(&self, items: Vec<SessionItem>) -> SessionFuture<()> {
686 let session_id = self.session_id.to_string();
687 let connection = self.connection.clone();
688 Box::pin(async move {
689 if items.is_empty() {
690 return Ok(());
691 }
692 let mut connection = connection
693 .lock()
694 .map_err(|_| "sqlite session store lock poisoned".to_string())?;
695 let transaction = connection.transaction().map_err(sqlite_error)?;
696 for item in items {
697 let payload = serde_json::to_string(&item).map_err(json_error)?;
698 transaction
699 .execute(
700 "INSERT INTO session_items (session_id, payload) VALUES (?1, ?2)",
701 params![session_id, payload],
702 )
703 .map_err(sqlite_error)?;
704 }
705 transaction.commit().map_err(sqlite_error)?;
706 Ok(())
707 })
708 }
709
710 fn supports_add_items_once(&self) -> bool {
711 true
712 }
713
714 fn add_items_once(
715 &self,
716 commit_id: String,
717 payload_digest: String,
718 items: Vec<SessionItem>,
719 ) -> SessionFuture<SessionAppendOutcome> {
720 let session_id = self.session_id.to_string();
721 let connection = self.connection.clone();
722 Box::pin(async move {
723 validate_session_commit(&commit_id, &payload_digest, &items)?;
724 let payloads = items
725 .iter()
726 .map(serde_json::to_string)
727 .collect::<Result<Vec<_>, _>>()
728 .map_err(json_error)?;
729 let mut connection = connection
730 .lock()
731 .map_err(|_| "sqlite session store lock poisoned".to_string())?;
732 let transaction = connection
733 .transaction_with_behavior(TransactionBehavior::Immediate)
734 .map_err(sqlite_error)?;
735 let existing = transaction
736 .query_row(
737 "SELECT payload_digest FROM session_commits WHERE session_id = ?1 AND commit_id = ?2",
738 params![session_id, commit_id],
739 |row| row.get::<_, String>(0),
740 )
741 .optional()
742 .map_err(sqlite_error)?;
743 if let Some(existing) = existing {
744 if existing != payload_digest {
745 return Err(
746 "session_commit_identity_conflict: commit_id has a different payload"
747 .to_string(),
748 );
749 }
750 transaction.commit().map_err(sqlite_error)?;
751 return Ok(SessionAppendOutcome::Replayed);
752 }
753 for payload in payloads {
754 transaction
755 .execute(
756 "INSERT INTO session_items (session_id, payload) VALUES (?1, ?2)",
757 params![session_id, payload],
758 )
759 .map_err(sqlite_error)?;
760 }
761 transaction
762 .execute(
763 "INSERT INTO session_commits (session_id, commit_id, payload_digest) VALUES (?1, ?2, ?3)",
764 params![session_id, commit_id, payload_digest],
765 )
766 .map_err(sqlite_error)?;
767 transaction.commit().map_err(sqlite_error)?;
768 Ok(SessionAppendOutcome::Committed)
769 })
770 }
771
772 fn pop_item(&self) -> SessionFuture<Option<SessionItem>> {
773 let session_id = self.session_id.to_string();
774 let connection = self.connection.clone();
775 Box::pin(async move {
776 let mut connection = connection
777 .lock()
778 .map_err(|_| "sqlite session store lock poisoned".to_string())?;
779 let transaction = connection.transaction().map_err(sqlite_error)?;
780 let row = transaction
781 .query_row(
782 r#"
783 SELECT item_index, payload
784 FROM session_items
785 WHERE session_id = ?1
786 ORDER BY item_index DESC
787 LIMIT 1
788 "#,
789 params![session_id],
790 |row| Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?)),
791 )
792 .optional()
793 .map_err(sqlite_error)?;
794 let Some((item_index, payload)) = row else {
795 transaction.commit().map_err(sqlite_error)?;
796 return Ok(None);
797 };
798 let item = serde_json::from_str(&payload).map_err(json_error)?;
799 transaction
800 .execute(
801 "DELETE FROM session_items WHERE item_index = ?1",
802 params![item_index],
803 )
804 .map_err(sqlite_error)?;
805 transaction.commit().map_err(sqlite_error)?;
806 Ok(Some(item))
807 })
808 }
809
810 fn clear(&self) -> SessionFuture<()> {
811 let session_id = self.session_id.to_string();
812 let connection = self.connection.clone();
813 Box::pin(async move {
814 let mut connection = connection
815 .lock()
816 .map_err(|_| "sqlite session store lock poisoned".to_string())?;
817 let transaction = connection.transaction().map_err(sqlite_error)?;
818 transaction
819 .execute(
820 "DELETE FROM session_items WHERE session_id = ?1",
821 params![session_id],
822 )
823 .map_err(sqlite_error)?;
824 transaction
825 .execute(
826 "DELETE FROM session_commits WHERE session_id = ?1",
827 params![session_id],
828 )
829 .map_err(sqlite_error)?;
830 transaction.commit().map_err(sqlite_error)?;
831 Ok(())
832 })
833 }
834}
835
836fn initialize_sqlite_session_schema(connection: &mut Connection) -> Result<(), String> {
837 let transaction = connection
838 .transaction_with_behavior(TransactionBehavior::Immediate)
839 .map_err(sqlite_error)?;
840 let version = transaction
841 .query_row("PRAGMA user_version", [], |row| row.get::<_, i64>(0))
842 .map_err(sqlite_error)?;
843 let tables = session_table_names(&transaction)?;
844 if tables.is_empty() {
845 if version != 0 {
846 return Err(format!(
847 "empty session database has unexpected schema version {version}"
848 ));
849 }
850 transaction
851 .execute_batch(CREATE_SESSION_ITEMS_TABLE)
852 .map_err(sqlite_error)?;
853 transaction
854 .execute_batch(CREATE_SESSION_ITEMS_INDEX)
855 .map_err(sqlite_error)?;
856 transaction
857 .execute_batch(CREATE_SESSION_COMMITS_TABLE)
858 .map_err(sqlite_error)?;
859 transaction
860 .execute_batch("PRAGMA user_version = 2;")
861 .map_err(sqlite_error)?;
862 } else {
863 if version != SQLITE_SESSION_SCHEMA_VERSION {
864 return Err(format!(
865 "session schema version {version} does not match required version \
866 {SQLITE_SESSION_SCHEMA_VERSION}"
867 ));
868 }
869 if tables != ["session_commits".to_string(), "session_items".to_string()] {
870 return Err(format!("unsupported session schema tables: {tables:?}"));
871 }
872 let item_columns = session_table_columns(&transaction, "session_items")?;
873 if item_columns != CANONICAL_SESSION_COLUMNS {
874 return Err(format!(
875 "unsupported session_items schema columns: {item_columns:?}"
876 ));
877 }
878 let commit_columns = session_table_columns(&transaction, "session_commits")?;
879 if commit_columns != CANONICAL_SESSION_COMMIT_COLUMNS {
880 return Err(format!(
881 "unsupported session_commits schema columns: {commit_columns:?}"
882 ));
883 }
884 let index_exists = transaction
885 .query_row(
886 "SELECT EXISTS(SELECT 1 FROM sqlite_master \
887 WHERE type = 'index' AND name = 'idx_session_items_session_id_item_index')",
888 [],
889 |row| row.get::<_, i64>(0),
890 )
891 .map_err(sqlite_error)?
892 != 0;
893 if !index_exists {
894 return Err("session schema is missing the canonical session_items index".to_string());
895 }
896 }
897 transaction.commit().map_err(sqlite_error)
898}
899
900fn session_table_names(connection: &Connection) -> Result<Vec<String>, String> {
901 let mut statement = connection
902 .prepare(
903 "SELECT name FROM sqlite_master \
904 WHERE type = 'table' AND name NOT LIKE 'sqlite_%' ORDER BY name",
905 )
906 .map_err(sqlite_error)?;
907 let rows = statement
908 .query_map([], |row| row.get::<_, String>(0))
909 .map_err(sqlite_error)?;
910 rows.collect::<rusqlite::Result<Vec<_>>>()
911 .map_err(sqlite_error)
912}
913
914fn session_table_columns(connection: &Connection, table: &str) -> Result<Vec<String>, String> {
915 let mut statement = connection
916 .prepare(&format!("PRAGMA table_info({table})"))
917 .map_err(sqlite_error)?;
918 let rows = statement
919 .query_map([], |row| row.get::<_, String>(1))
920 .map_err(sqlite_error)?;
921 rows.collect::<rusqlite::Result<Vec<_>>>()
922 .map_err(sqlite_error)
923}
924
925fn sqlite_error(error: rusqlite::Error) -> String {
926 error.to_string()
927}
928
929fn json_error(error: serde_json::Error) -> String {
930 error.to_string()
931}