refactor: remove threads layer, use sessions directly for ACP (#9078)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -19,7 +19,7 @@ use std::sync::{Arc, LazyLock};
|
||||
use tracing::{info, warn};
|
||||
use utoipa::ToSchema;
|
||||
|
||||
pub const CURRENT_SCHEMA_VERSION: i32 = 11;
|
||||
pub const CURRENT_SCHEMA_VERSION: i32 = 12;
|
||||
pub const SESSIONS_FOLDER: &str = "sessions";
|
||||
pub const DB_NAME: &str = "sessions.db";
|
||||
|
||||
@@ -82,7 +82,9 @@ pub struct Session {
|
||||
#[serde(default)]
|
||||
pub goose_mode: GooseMode,
|
||||
#[serde(default)]
|
||||
pub thread_id: Option<String>,
|
||||
pub archived_at: Option<DateTime<Utc>>,
|
||||
#[serde(default)]
|
||||
pub project_id: Option<String>,
|
||||
}
|
||||
|
||||
pub struct SessionUpdateBuilder<'a> {
|
||||
@@ -105,7 +107,9 @@ pub struct SessionUpdateBuilder<'a> {
|
||||
provider_name: Option<Option<String>>,
|
||||
model_config: Option<Option<ModelConfig>>,
|
||||
goose_mode: Option<GooseMode>,
|
||||
thread_id: Option<Option<String>>,
|
||||
archived_at: Option<Option<DateTime<Utc>>>,
|
||||
|
||||
project_id: Option<Option<String>>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, ToSchema, Debug)]
|
||||
@@ -137,7 +141,8 @@ impl<'a> SessionUpdateBuilder<'a> {
|
||||
provider_name: None,
|
||||
model_config: None,
|
||||
goose_mode: None,
|
||||
thread_id: None,
|
||||
archived_at: None,
|
||||
project_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -246,8 +251,13 @@ impl<'a> SessionUpdateBuilder<'a> {
|
||||
self
|
||||
}
|
||||
|
||||
pub fn thread_id(mut self, thread_id: Option<String>) -> Self {
|
||||
self.thread_id = Some(thread_id);
|
||||
pub fn archived_at(mut self, archived_at: Option<DateTime<Utc>>) -> Self {
|
||||
self.archived_at = Some(archived_at);
|
||||
self
|
||||
}
|
||||
|
||||
pub fn project_id(mut self, project_id: Option<String>) -> Self {
|
||||
self.project_id = Some(project_id);
|
||||
self
|
||||
}
|
||||
}
|
||||
@@ -258,7 +268,11 @@ pub struct SessionManager {
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SessionNameUpdate {
|
||||
pub thread: super::thread_manager::Thread,
|
||||
pub session_id: String,
|
||||
pub name: String,
|
||||
pub updated_at: chrono::DateTime<chrono::Utc>,
|
||||
pub message_count: usize,
|
||||
pub user_set_name: bool,
|
||||
}
|
||||
|
||||
impl SessionManager {
|
||||
@@ -384,21 +398,16 @@ impl SessionManager {
|
||||
.apply()
|
||||
.await?;
|
||||
|
||||
// Also update the thread name so ACP clients see it via session/list.
|
||||
if let Some(ref thread_id) = session.thread_id {
|
||||
let thread_mgr = super::thread_manager::ThreadManager::new(self.storage.clone());
|
||||
let thread = thread_mgr.get_thread(thread_id).await?;
|
||||
if !thread.user_set_name && thread.name != name {
|
||||
let thread = thread_mgr
|
||||
.update_thread(thread_id, Some(name), Some(false), None)
|
||||
.await?;
|
||||
return Ok(Some(SessionNameUpdate { thread }));
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(None)
|
||||
let session = self.get_session(id, false).await?;
|
||||
return Ok(Some(SessionNameUpdate {
|
||||
session_id: id.to_string(),
|
||||
name,
|
||||
updated_at: session.updated_at,
|
||||
message_count: session.message_count,
|
||||
user_set_name: session.user_set_name,
|
||||
}));
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
pub async fn search_chat_history(
|
||||
@@ -433,6 +442,22 @@ impl SessionManager {
|
||||
.update_message_metadata(id, message_id, f)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Patch `tool_meta` on a specific `ToolRequest` within a stored message.
|
||||
/// Used to persist LLM-generated tool titles and chain summaries so they
|
||||
/// survive session reload. Merge-based: existing keys not in `patch` are
|
||||
/// preserved. No-op if the message or tool_call_id is not found.
|
||||
pub async fn update_tool_request_meta(
|
||||
&self,
|
||||
session_id: &str,
|
||||
message_id: &str,
|
||||
tool_call_id: &str,
|
||||
patch: serde_json::Value,
|
||||
) -> Result<()> {
|
||||
self.storage
|
||||
.update_tool_request_meta(session_id, message_id, tool_call_id, patch)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
pub struct SessionStorage {
|
||||
@@ -473,7 +498,8 @@ impl Default for Session {
|
||||
provider_name: None,
|
||||
model_config: None,
|
||||
goose_mode: GooseMode::default(),
|
||||
thread_id: None,
|
||||
archived_at: None,
|
||||
project_id: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -543,7 +569,8 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session {
|
||||
.ok()
|
||||
.and_then(|s| s.parse().ok())
|
||||
.unwrap_or_default(),
|
||||
thread_id: row.try_get("thread_id").ok().flatten(),
|
||||
archived_at: row.try_get("archived_at").ok(),
|
||||
project_id: row.try_get("project_id").ok().flatten(),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -644,7 +671,8 @@ impl SessionStorage {
|
||||
provider_name TEXT,
|
||||
model_config_json TEXT,
|
||||
goose_mode TEXT NOT NULL DEFAULT 'auto',
|
||||
thread_id TEXT
|
||||
archived_at TIMESTAMP,
|
||||
project_id TEXT
|
||||
)
|
||||
"#,
|
||||
)
|
||||
@@ -684,49 +712,6 @@ impl SessionStorage {
|
||||
sqlx::query("CREATE INDEX idx_sessions_type ON sessions(session_type)")
|
||||
.execute(pool)
|
||||
.await?;
|
||||
sqlx::query("CREATE INDEX IF NOT EXISTS idx_sessions_thread ON sessions(thread_id)")
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
"CREATE TABLE IF NOT EXISTS threads (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL DEFAULT 'New Chat',
|
||||
user_set_name BOOLEAN DEFAULT FALSE,
|
||||
working_dir TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
archived_at TIMESTAMP,
|
||||
metadata_json TEXT DEFAULT '{}'
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
"CREATE TABLE IF NOT EXISTS thread_messages (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
thread_id TEXT NOT NULL REFERENCES threads(id),
|
||||
session_id TEXT,
|
||||
message_id TEXT,
|
||||
role TEXT NOT NULL,
|
||||
content_json TEXT NOT NULL,
|
||||
created_timestamp INTEGER NOT NULL,
|
||||
metadata_json TEXT DEFAULT '{}'
|
||||
)",
|
||||
)
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
"CREATE INDEX IF NOT EXISTS idx_thread_messages_thread ON thread_messages(thread_id)",
|
||||
)
|
||||
.execute(pool)
|
||||
.await?;
|
||||
sqlx::query("CREATE INDEX IF NOT EXISTS idx_thread_messages_message_id ON thread_messages(message_id)")
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
crate::providers::inventory::create_tables(pool).await?;
|
||||
|
||||
Ok(())
|
||||
@@ -1075,6 +1060,46 @@ impl SessionStorage {
|
||||
11 => {
|
||||
crate::providers::inventory::create_tables_in_tx(tx).await?;
|
||||
}
|
||||
12 => {
|
||||
// Add archived_at, project_id columns to sessions.
|
||||
let has_archived_at = sqlx::query_scalar::<_, i32>(
|
||||
"SELECT COUNT(*) FROM pragma_table_info('sessions') WHERE name = 'archived_at'",
|
||||
)
|
||||
.fetch_one(&mut **tx)
|
||||
.await?
|
||||
> 0;
|
||||
if !has_archived_at {
|
||||
sqlx::query("ALTER TABLE sessions ADD COLUMN archived_at TIMESTAMP")
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
}
|
||||
|
||||
let has_project_id = sqlx::query_scalar::<_, i32>(
|
||||
"SELECT COUNT(*) FROM pragma_table_info('sessions') WHERE name = 'project_id'",
|
||||
)
|
||||
.fetch_one(&mut **tx)
|
||||
.await?
|
||||
> 0;
|
||||
if !has_project_id {
|
||||
sqlx::query("ALTER TABLE sessions ADD COLUMN project_id TEXT")
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
}
|
||||
|
||||
// Drop thread tables and thread_id column (no longer used).
|
||||
sqlx::query("DROP TABLE IF EXISTS thread_messages")
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
sqlx::query("DROP TABLE IF EXISTS threads")
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
sqlx::query("DROP INDEX IF EXISTS idx_sessions_thread")
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
sqlx::query("ALTER TABLE sessions DROP COLUMN thread_id")
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
}
|
||||
_ => {
|
||||
anyhow::bail!("Unknown migration version: {}", version);
|
||||
}
|
||||
@@ -1136,7 +1161,8 @@ impl SessionStorage {
|
||||
total_tokens, input_tokens, output_tokens,
|
||||
accumulated_total_tokens, accumulated_input_tokens, accumulated_output_tokens,
|
||||
schedule_id, recipe_json, user_recipe_values_json,
|
||||
provider_name, model_config_json, goose_mode, thread_id
|
||||
provider_name, model_config_json, goose_mode,
|
||||
archived_at, project_id
|
||||
FROM sessions
|
||||
WHERE id = ?
|
||||
"#,
|
||||
@@ -1200,7 +1226,9 @@ impl SessionStorage {
|
||||
add_update!(builder.provider_name, "provider_name");
|
||||
add_update!(builder.model_config, "model_config_json");
|
||||
add_update!(builder.goose_mode, "goose_mode");
|
||||
add_update!(builder.thread_id, "thread_id");
|
||||
add_update!(builder.archived_at, "archived_at");
|
||||
|
||||
add_update!(builder.project_id, "project_id");
|
||||
|
||||
if updates.is_empty() {
|
||||
return Ok(());
|
||||
@@ -1269,14 +1297,22 @@ impl SessionStorage {
|
||||
if let Some(goose_mode) = builder.goose_mode {
|
||||
q = q.bind(goose_mode.to_string());
|
||||
}
|
||||
if let Some(thread_id) = builder.thread_id {
|
||||
q = q.bind(thread_id);
|
||||
if let Some(ref archived_at) = builder.archived_at {
|
||||
q = q.bind(archived_at.as_ref());
|
||||
}
|
||||
|
||||
if let Some(ref project_id) = builder.project_id {
|
||||
q = q.bind(project_id.as_ref());
|
||||
}
|
||||
|
||||
let pool = self.pool().await?;
|
||||
let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?;
|
||||
q = q.bind(&builder.session_id);
|
||||
q.execute(&mut *tx).await?;
|
||||
let result = q.execute(&mut *tx).await?;
|
||||
|
||||
if result.rows_affected() == 0 {
|
||||
return Err(anyhow::anyhow!("Session not found: {}", builder.session_id));
|
||||
}
|
||||
|
||||
tx.commit().await?;
|
||||
Ok(())
|
||||
@@ -1423,7 +1459,8 @@ impl SessionStorage {
|
||||
s.total_tokens, s.input_tokens, s.output_tokens,
|
||||
s.accumulated_total_tokens, s.accumulated_input_tokens, s.accumulated_output_tokens,
|
||||
s.schedule_id, s.recipe_json, s.user_recipe_values_json,
|
||||
s.provider_name, s.model_config_json, s.goose_mode, s.thread_id,
|
||||
s.provider_name, s.model_config_json, s.goose_mode,
|
||||
s.archived_at, s.project_id,
|
||||
COUNT(m.id) as message_count
|
||||
FROM sessions s
|
||||
LEFT JOIN messages m ON s.id = m.session_id
|
||||
@@ -1582,15 +1619,15 @@ impl SessionStorage {
|
||||
.recipe(original_session.recipe)
|
||||
.user_recipe_values(original_session.user_recipe_values);
|
||||
|
||||
// Preserve provider, model config, and goose_mode from original session
|
||||
if let Some(project_id) = original_session.project_id {
|
||||
builder = builder.project_id(Some(project_id));
|
||||
}
|
||||
if let Some(provider_name) = original_session.provider_name {
|
||||
builder = builder.provider_name(provider_name);
|
||||
}
|
||||
|
||||
if let Some(model_config) = original_session.model_config {
|
||||
builder = builder.model_config(model_config);
|
||||
}
|
||||
|
||||
builder = builder.goose_mode(original_session.goose_mode);
|
||||
|
||||
builder.apply().await?;
|
||||
@@ -1680,6 +1717,81 @@ impl SessionStorage {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Patch `tool_meta` on a specific `ToolRequest` within a stored message's
|
||||
/// `content_json`. Finds the row(s) with matching `message_id`, scans each
|
||||
/// row's content for a `ToolRequest` with the given `tool_call_id`, and
|
||||
/// merges `patch` into its `tool_meta`. Uses `BEGIN IMMEDIATE` so
|
||||
/// concurrent writers serialize correctly.
|
||||
async fn update_tool_request_meta(
|
||||
&self,
|
||||
session_id: &str,
|
||||
message_id: &str,
|
||||
tool_call_id: &str,
|
||||
patch: serde_json::Value,
|
||||
) -> Result<()> {
|
||||
use crate::conversation::message::MessageContent;
|
||||
|
||||
let pool = self.pool().await?;
|
||||
let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?;
|
||||
|
||||
let rows = sqlx::query_as::<_, (i64, String)>(
|
||||
"SELECT id, content_json FROM messages \
|
||||
WHERE session_id = ? AND message_id = ? \
|
||||
ORDER BY id ASC",
|
||||
)
|
||||
.bind(session_id)
|
||||
.bind(message_id)
|
||||
.fetch_all(&mut *tx)
|
||||
.await?;
|
||||
|
||||
for (row_id, content_json) in rows {
|
||||
let mut content: Vec<MessageContent> = serde_json::from_str(&content_json)?;
|
||||
let mut found = false;
|
||||
for block in &mut content {
|
||||
if let MessageContent::ToolRequest(tr) = block {
|
||||
if tr.id == tool_call_id {
|
||||
tr.tool_meta = Some(merge_tool_meta(tr.tool_meta.take(), &patch));
|
||||
found = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
continue;
|
||||
}
|
||||
|
||||
let updated_json = serde_json::to_string(&content)?;
|
||||
sqlx::query("UPDATE messages SET content_json = ? WHERE id = ?")
|
||||
.bind(updated_json)
|
||||
.bind(row_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
tx.commit().await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge a JSON object `patch` into an existing optional object value,
|
||||
/// preserving keys not present in the patch.
|
||||
fn merge_tool_meta(
|
||||
existing: Option<serde_json::Value>,
|
||||
patch: &serde_json::Value,
|
||||
) -> serde_json::Value {
|
||||
let mut base = match existing {
|
||||
Some(serde_json::Value::Object(map)) => map,
|
||||
_ => serde_json::Map::new(),
|
||||
};
|
||||
if let serde_json::Value::Object(patch_map) = patch {
|
||||
for (k, v) in patch_map {
|
||||
base.insert(k.clone(), v.clone());
|
||||
}
|
||||
}
|
||||
serde_json::Value::Object(base)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
Reference in New Issue
Block a user