Skip to main content

vv_agent/
sessions.rs

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}