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<()> {
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<Sqlite>) -> Result<i32> {
async fn get_schema_version(tx: &mut sqlx::Transaction<'_, Sqlite>) -> Result<i32> {
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<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 (?)")
.bind(version)
.execute(pool)
.execute(&mut **tx)
.await?;
Ok(())
}
#[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 {
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 = ?",