Session manager fixes (#6809)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2026-01-29 11:09:10 -05:00
committed by GitHub
parent 5bffc1c8ef
commit ec789475d9
+37 -24
View File
@@ -724,7 +724,9 @@ impl SessionStorage {
} }
async fn run_migrations(pool: &Pool<Sqlite>) -> Result<()> { async fn run_migrations(pool: &Pool<Sqlite>) -> Result<()> {
let current_version = Self::get_schema_version(pool).await?; let mut tx = pool.begin().await?;
let current_version = Self::get_schema_version(&mut tx).await?;
if current_version < CURRENT_SCHEMA_VERSION { if current_version < CURRENT_SCHEMA_VERSION {
info!( info!(
@@ -734,18 +736,19 @@ impl SessionStorage {
for version in (current_version + 1)..=CURRENT_SCHEMA_VERSION { for version in (current_version + 1)..=CURRENT_SCHEMA_VERSION {
info!(" Applying migration v{}...", version); info!(" Applying migration v{}...", version);
Self::apply_migration(pool, version).await?; Self::apply_migration(&mut tx, version).await?;
Self::update_schema_version(pool, version).await?; Self::update_schema_version(&mut tx, version).await?;
info!(" ✓ Migration v{} complete", version); info!(" ✓ Migration v{} complete", version);
} }
info!("All migrations complete"); info!("All migrations complete");
} }
tx.commit().await?;
Ok(()) Ok(())
} }
async fn get_schema_version(pool: &Pool<Sqlite>) -> Result<i32> { async fn get_schema_version(tx: &mut sqlx::Transaction<'_, Sqlite>) -> Result<i32> {
let table_exists = sqlx::query_scalar::<_, bool>( let table_exists = sqlx::query_scalar::<_, bool>(
r#" r#"
SELECT EXISTS ( SELECT EXISTS (
@@ -754,7 +757,7 @@ impl SessionStorage {
) )
"#, "#,
) )
.fetch_one(pool) .fetch_one(&mut **tx)
.await?; .await?;
if !table_exists { if !table_exists {
@@ -762,22 +765,25 @@ impl SessionStorage {
} }
let version = sqlx::query_scalar::<_, i32>("SELECT MAX(version) FROM schema_version") let version = sqlx::query_scalar::<_, i32>("SELECT MAX(version) FROM schema_version")
.fetch_one(pool) .fetch_one(&mut **tx)
.await?; .await?;
Ok(version) Ok(version)
} }
async fn update_schema_version(pool: &Pool<Sqlite>, version: i32) -> Result<()> { async fn update_schema_version(
tx: &mut sqlx::Transaction<'_, Sqlite>,
version: i32,
) -> Result<()> {
sqlx::query("INSERT INTO schema_version (version) VALUES (?)") sqlx::query("INSERT INTO schema_version (version) VALUES (?)")
.bind(version) .bind(version)
.execute(pool) .execute(&mut **tx)
.await?; .await?;
Ok(()) Ok(())
} }
#[allow(clippy::too_many_lines)] #[allow(clippy::too_many_lines)]
async fn apply_migration(pool: &Pool<Sqlite>, version: i32) -> Result<()> { async fn apply_migration(tx: &mut sqlx::Transaction<'_, Sqlite>, version: i32) -> Result<()> {
match version { match version {
1 => { 1 => {
sqlx::query( sqlx::query(
@@ -788,7 +794,7 @@ impl SessionStorage {
) )
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
2 => { 2 => {
@@ -797,7 +803,7 @@ impl SessionStorage {
ALTER TABLE sessions ADD COLUMN user_recipe_values_json TEXT ALTER TABLE sessions ADD COLUMN user_recipe_values_json TEXT
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
3 => { 3 => {
@@ -806,7 +812,7 @@ impl SessionStorage {
ALTER TABLE messages ADD COLUMN metadata_json TEXT ALTER TABLE messages ADD COLUMN metadata_json TEXT
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
4 => { 4 => {
@@ -815,7 +821,7 @@ impl SessionStorage {
ALTER TABLE sessions ADD COLUMN name TEXT DEFAULT '' ALTER TABLE sessions ADD COLUMN name TEXT DEFAULT ''
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
sqlx::query( sqlx::query(
@@ -823,7 +829,7 @@ impl SessionStorage {
ALTER TABLE sessions ADD COLUMN user_set_name BOOLEAN DEFAULT FALSE ALTER TABLE sessions ADD COLUMN user_set_name BOOLEAN DEFAULT FALSE
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
5 => { 5 => {
@@ -832,11 +838,11 @@ impl SessionStorage {
ALTER TABLE sessions ADD COLUMN session_type TEXT NOT NULL DEFAULT 'user' ALTER TABLE sessions ADD COLUMN session_type TEXT NOT NULL DEFAULT 'user'
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
sqlx::query("CREATE INDEX idx_sessions_type ON sessions(session_type)") sqlx::query("CREATE INDEX idx_sessions_type ON sessions(session_type)")
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
6 => { 6 => {
@@ -845,7 +851,7 @@ impl SessionStorage {
ALTER TABLE sessions ADD COLUMN provider_name TEXT ALTER TABLE sessions ADD COLUMN provider_name TEXT
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
sqlx::query( sqlx::query(
@@ -853,7 +859,7 @@ impl SessionStorage {
ALTER TABLE sessions ADD COLUMN model_config_json TEXT ALTER TABLE sessions ADD COLUMN model_config_json TEXT
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
7 => { 7 => {
@@ -862,7 +868,7 @@ impl SessionStorage {
ALTER TABLE messages ADD COLUMN message_id TEXT ALTER TABLE messages ADD COLUMN message_id TEXT
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
sqlx::query( sqlx::query(
@@ -871,11 +877,11 @@ impl SessionStorage {
SET message_id = 'msg_' || session_id || '_' || id SET message_id = 'msg_' || session_id || '_' || id
"#, "#,
) )
.execute(pool) .execute(&mut **tx)
.await?; .await?;
sqlx::query("CREATE INDEX idx_messages_message_id ON messages(message_id)") sqlx::query("CREATE INDEX idx_messages_message_id ON messages(message_id)")
.execute(pool) .execute(&mut **tx)
.await?; .await?;
} }
_ => { _ => {
@@ -1158,12 +1164,18 @@ impl SessionStorage {
for message in conversation.messages() { for message in conversation.messages() {
let metadata_json = serde_json::to_string(&message.metadata)?; let metadata_json = serde_json::to_string(&message.metadata)?;
let message_id = message
.id
.clone()
.unwrap_or_else(|| format!("msg_{}_{}", session_id, uuid::Uuid::new_v4()));
sqlx::query( sqlx::query(
r#" r#"
INSERT INTO messages (session_id, role, content_json, created_timestamp, metadata_json) INSERT INTO messages (message_id, session_id, role, content_json, created_timestamp, metadata_json)
VALUES (?, ?, ?, ?, ?) VALUES (?, ?, ?, ?, ?, ?)
"#, "#,
) )
.bind(message_id)
.bind(session_id) .bind(session_id)
.bind(role_to_string(&message.role)) .bind(role_to_string(&message.role))
.bind(serde_json::to_string(&message.content)?) .bind(serde_json::to_string(&message.content)?)
@@ -1403,7 +1415,8 @@ impl SessionStorage {
crate::conversation::message::MessageMetadata, crate::conversation::message::MessageMetadata,
) -> crate::conversation::message::MessageMetadata, ) -> crate::conversation::message::MessageMetadata,
{ {
let mut tx = self.pool.begin().await?; let pool = self.pool().await?;
let mut tx = pool.begin().await?;
let current_metadata_json = sqlx::query_scalar::<_, String>( let current_metadata_json = sqlx::query_scalar::<_, String>(
"SELECT metadata_json FROM messages WHERE message_id = ? AND session_id = ?", "SELECT metadata_json FROM messages WHERE message_id = ? AND session_id = ?",