Session manager fixes (#6809)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -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 = ?",
|
||||
|
||||
Reference in New Issue
Block a user