From ec789475d912cb06b8c3bcecbfcb2448410abebf Mon Sep 17 00:00:00 2001 From: Douwe Osinga Date: Thu, 29 Jan 2026 11:09:10 -0500 Subject: [PATCH] Session manager fixes (#6809) Co-authored-by: Douwe Osinga --- crates/goose/src/session/session_manager.rs | 61 +++++++++++++-------- 1 file changed, 37 insertions(+), 24 deletions(-) diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index 12af879b..ea039da6 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -724,7 +724,9 @@ impl SessionStorage { } async fn run_migrations(pool: &Pool) -> 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 { info!( @@ -734,18 +736,19 @@ impl SessionStorage { for version in (current_version + 1)..=CURRENT_SCHEMA_VERSION { info!(" Applying migration v{}...", version); - Self::apply_migration(pool, version).await?; - Self::update_schema_version(pool, version).await?; + Self::apply_migration(&mut tx, version).await?; + Self::update_schema_version(&mut tx, version).await?; info!(" ✓ Migration v{} complete", version); } info!("All migrations complete"); } + tx.commit().await?; Ok(()) } - async fn get_schema_version(pool: &Pool) -> Result { + async fn get_schema_version(tx: &mut sqlx::Transaction<'_, Sqlite>) -> Result { let table_exists = sqlx::query_scalar::<_, bool>( r#" SELECT EXISTS ( @@ -754,7 +757,7 @@ impl SessionStorage { ) "#, ) - .fetch_one(pool) + .fetch_one(&mut **tx) .await?; if !table_exists { @@ -762,22 +765,25 @@ impl SessionStorage { } let version = sqlx::query_scalar::<_, i32>("SELECT MAX(version) FROM schema_version") - .fetch_one(pool) + .fetch_one(&mut **tx) .await?; Ok(version) } - async fn update_schema_version(pool: &Pool, version: i32) -> Result<()> { + async fn update_schema_version( + tx: &mut sqlx::Transaction<'_, Sqlite>, + version: i32, + ) -> Result<()> { sqlx::query("INSERT INTO schema_version (version) VALUES (?)") .bind(version) - .execute(pool) + .execute(&mut **tx) .await?; Ok(()) } #[allow(clippy::too_many_lines)] - async fn apply_migration(pool: &Pool, version: i32) -> Result<()> { + async fn apply_migration(tx: &mut sqlx::Transaction<'_, Sqlite>, version: i32) -> Result<()> { match version { 1 => { sqlx::query( @@ -788,7 +794,7 @@ impl SessionStorage { ) "#, ) - .execute(pool) + .execute(&mut **tx) .await?; } 2 => { @@ -797,7 +803,7 @@ impl SessionStorage { ALTER TABLE sessions ADD COLUMN user_recipe_values_json TEXT "#, ) - .execute(pool) + .execute(&mut **tx) .await?; } 3 => { @@ -806,7 +812,7 @@ impl SessionStorage { ALTER TABLE messages ADD COLUMN metadata_json TEXT "#, ) - .execute(pool) + .execute(&mut **tx) .await?; } 4 => { @@ -815,7 +821,7 @@ impl SessionStorage { ALTER TABLE sessions ADD COLUMN name TEXT DEFAULT '' "#, ) - .execute(pool) + .execute(&mut **tx) .await?; sqlx::query( @@ -823,7 +829,7 @@ impl SessionStorage { ALTER TABLE sessions ADD COLUMN user_set_name BOOLEAN DEFAULT FALSE "#, ) - .execute(pool) + .execute(&mut **tx) .await?; } 5 => { @@ -832,11 +838,11 @@ impl SessionStorage { ALTER TABLE sessions ADD COLUMN session_type TEXT NOT NULL DEFAULT 'user' "#, ) - .execute(pool) + .execute(&mut **tx) .await?; sqlx::query("CREATE INDEX idx_sessions_type ON sessions(session_type)") - .execute(pool) + .execute(&mut **tx) .await?; } 6 => { @@ -845,7 +851,7 @@ impl SessionStorage { ALTER TABLE sessions ADD COLUMN provider_name TEXT "#, ) - .execute(pool) + .execute(&mut **tx) .await?; sqlx::query( @@ -853,7 +859,7 @@ impl SessionStorage { ALTER TABLE sessions ADD COLUMN model_config_json TEXT "#, ) - .execute(pool) + .execute(&mut **tx) .await?; } 7 => { @@ -862,7 +868,7 @@ impl SessionStorage { ALTER TABLE messages ADD COLUMN message_id TEXT "#, ) - .execute(pool) + .execute(&mut **tx) .await?; sqlx::query( @@ -871,11 +877,11 @@ impl SessionStorage { SET message_id = 'msg_' || session_id || '_' || id "#, ) - .execute(pool) + .execute(&mut **tx) .await?; sqlx::query("CREATE INDEX idx_messages_message_id ON messages(message_id)") - .execute(pool) + .execute(&mut **tx) .await?; } _ => { @@ -1158,12 +1164,18 @@ impl SessionStorage { for message in conversation.messages() { 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( r#" - INSERT INTO messages (session_id, role, content_json, created_timestamp, metadata_json) - VALUES (?, ?, ?, ?, ?) + INSERT INTO messages (message_id, session_id, role, content_json, created_timestamp, metadata_json) + VALUES (?, ?, ?, ?, ?, ?) "#, ) + .bind(message_id) .bind(session_id) .bind(role_to_string(&message.role)) .bind(serde_json::to_string(&message.content)?) @@ -1403,7 +1415,8 @@ impl SessionStorage { 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>( "SELECT metadata_json FROM messages WHERE message_id = ? AND session_id = ?",