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<()> {
|
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 = ?",
|
||||||
|
|||||||
Reference in New Issue
Block a user