feat: persist GooseMode per-session via session DB (#7854)

Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
Adrian Cole
2026-03-16 19:37:21 +08:00
committed by GitHub
parent 2631095f20
commit 94fdcdd07a
42 changed files with 953 additions and 147 deletions
+131 -11
View File
@@ -1,4 +1,5 @@
use crate::config::paths::Paths;
use crate::config::GooseMode;
use crate::conversation::message::Message;
use crate::conversation::Conversation;
use crate::model::ModelConfig;
@@ -18,7 +19,7 @@ use std::sync::{Arc, LazyLock};
use tracing::{info, warn};
use utoipa::ToSchema;
pub const CURRENT_SCHEMA_VERSION: i32 = 7;
pub const CURRENT_SCHEMA_VERSION: i32 = 8;
pub const SESSIONS_FOLDER: &str = "sessions";
pub const DB_NAME: &str = "sessions.db";
@@ -93,6 +94,8 @@ pub struct Session {
pub message_count: usize,
pub provider_name: Option<String>,
pub model_config: Option<ModelConfig>,
#[serde(default)]
pub goose_mode: GooseMode,
}
pub struct SessionUpdateBuilder<'a> {
@@ -114,6 +117,7 @@ pub struct SessionUpdateBuilder<'a> {
user_recipe_values: Option<Option<HashMap<String, String>>>,
provider_name: Option<Option<String>>,
model_config: Option<Option<ModelConfig>>,
goose_mode: Option<GooseMode>,
}
#[derive(Serialize, ToSchema, Debug)]
@@ -144,6 +148,7 @@ impl<'a> SessionUpdateBuilder<'a> {
user_recipe_values: None,
provider_name: None,
model_config: None,
goose_mode: None,
}
}
@@ -241,6 +246,11 @@ impl<'a> SessionUpdateBuilder<'a> {
self.model_config = Some(Some(model_config));
self
}
pub fn goose_mode(mut self, mode: GooseMode) -> Self {
self.goose_mode = Some(mode);
self
}
}
pub struct SessionManager {
@@ -269,9 +279,10 @@ impl SessionManager {
working_dir: PathBuf,
name: String,
session_type: SessionType,
goose_mode: GooseMode,
) -> Result<Session> {
self.storage
.create_session(working_dir, name, session_type)
.create_session(working_dir, name, session_type, goose_mode)
.await
}
@@ -417,6 +428,7 @@ impl Default for Session {
message_count: 0,
provider_name: None,
model_config: None,
goose_mode: GooseMode::default(),
}
}
}
@@ -481,6 +493,11 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session {
message_count: row.try_get("message_count").unwrap_or(0) as usize,
provider_name: row.try_get("provider_name").ok().flatten(),
model_config,
goose_mode: row
.try_get::<String, _>("goose_mode")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or_default(),
})
}
}
@@ -579,7 +596,8 @@ impl SessionStorage {
recipe_json TEXT,
user_recipe_values_json TEXT,
provider_name TEXT,
model_config_json TEXT
model_config_json TEXT,
goose_mode TEXT NOT NULL DEFAULT 'auto'
)
"#,
)
@@ -692,8 +710,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
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
provider_name, model_config_json, goose_mode
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&session.id)
@@ -715,6 +733,7 @@ impl SessionStorage {
.bind(user_recipe_values_json)
.bind(&session.provider_name)
.bind(model_config_json)
.bind(session.goose_mode.to_string())
.execute(&mut *tx)
.await?;
@@ -887,6 +906,15 @@ impl SessionStorage {
.execute(&mut **tx)
.await?;
}
8 => {
sqlx::query(
r#"
ALTER TABLE sessions ADD COLUMN goose_mode TEXT NOT NULL DEFAULT 'auto'
"#,
)
.execute(&mut **tx)
.await?;
}
_ => {
anyhow::bail!("Unknown migration version: {}", version);
}
@@ -900,6 +928,7 @@ impl SessionStorage {
working_dir: PathBuf,
name: String,
session_type: SessionType,
goose_mode: GooseMode,
) -> Result<Session> {
let pool = self.pool().await?;
let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?;
@@ -907,7 +936,7 @@ impl SessionStorage {
let today = chrono::Utc::now().format("%Y%m%d").to_string();
let session = sqlx::query_as(
r#"
INSERT INTO sessions (id, name, user_set_name, session_type, working_dir, extension_data)
INSERT INTO sessions (id, name, user_set_name, session_type, working_dir, extension_data, goose_mode)
VALUES (
? || '_' || CAST(COALESCE((
SELECT MAX(CAST(SUBSTR(id, 10) AS INTEGER))
@@ -918,7 +947,8 @@ impl SessionStorage {
FALSE,
?,
?,
'{}'
'{}',
?
)
RETURNING *
"#,
@@ -928,6 +958,7 @@ impl SessionStorage {
.bind(&name)
.bind(session_type.to_string())
.bind(&*working_dir.to_string_lossy())
.bind(goose_mode.to_string())
.fetch_one(&mut *tx)
.await?;
@@ -944,7 +975,7 @@ 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
provider_name, model_config_json, goose_mode
FROM sessions
WHERE id = ?
"#,
@@ -1007,6 +1038,7 @@ impl SessionStorage {
add_update!(builder.user_recipe_values, "user_recipe_values_json");
add_update!(builder.provider_name, "provider_name");
add_update!(builder.model_config, "model_config_json");
add_update!(builder.goose_mode, "goose_mode");
if updates.is_empty() {
return Ok(());
@@ -1072,6 +1104,9 @@ impl SessionStorage {
.transpose()?;
q = q.bind(model_config_json);
}
if let Some(goose_mode) = builder.goose_mode {
q = q.bind(goose_mode.to_string());
}
let pool = self.pool().await?;
let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?;
@@ -1216,7 +1251,7 @@ 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.provider_name, s.model_config_json, s.goose_mode,
COUNT(m.id) as message_count
FROM sessions s
INNER JOIN messages m ON s.id = m.session_id
@@ -1304,6 +1339,7 @@ impl SessionStorage {
import.working_dir.clone(),
import.name.clone(),
import.session_type,
import.goose_mode,
)
.await?;
@@ -1347,6 +1383,7 @@ impl SessionStorage {
original_session.working_dir.clone(),
new_name,
original_session.session_type,
original_session.goose_mode,
)
.await?;
@@ -1357,7 +1394,7 @@ impl SessionStorage {
.recipe(original_session.recipe)
.user_recipe_values(original_session.user_recipe_values);
// Preserve provider and model config from original session
// Preserve provider, model config, and goose_mode from original session
if let Some(provider_name) = original_session.provider_name {
builder = builder.provider_name(provider_name);
}
@@ -1366,6 +1403,8 @@ impl SessionStorage {
builder = builder.model_config(model_config);
}
builder = builder.goose_mode(original_session.goose_mode);
builder.apply().await?;
if let Some(conversation) = original_session.conversation {
@@ -1458,6 +1497,7 @@ mod tests {
use super::*;
use crate::conversation::message::{Message, MessageContent};
use tempfile::TempDir;
use test_case::test_case;
const NUM_CONCURRENT_SESSIONS: i32 = 10;
@@ -1529,6 +1569,7 @@ mod tests {
PathBuf::from("/tmp/lock-upgrade-test"),
"Lock Upgrade Session".to_string(),
SessionType::User,
GooseMode::default(),
)
.await
.unwrap();
@@ -1566,7 +1607,12 @@ mod tests {
let description = format!("Test session {}", i);
let session = sm
.create_session(working_dir.clone(), description, SessionType::User)
.create_session(
working_dir.clone(),
description,
SessionType::User,
GooseMode::default(),
)
.await
.unwrap();
@@ -1654,6 +1700,7 @@ mod tests {
PathBuf::from("/tmp/test"),
DESCRIPTION.to_string(),
SessionType::User,
GooseMode::default(),
)
.await
.unwrap();
@@ -1733,4 +1780,77 @@ mod tests {
assert!(imported.user_set_name);
assert_eq!(imported.working_dir, PathBuf::from("/tmp/test"));
}
#[test_case(GooseMode::Approve)]
#[test_case(GooseMode::SmartApprove)]
#[test_case(GooseMode::Chat)]
#[tokio::test]
async fn test_goose_mode_persists(mode: GooseMode) {
let temp_dir = TempDir::new().unwrap();
let sm = SessionManager::new(temp_dir.path().to_path_buf());
let session = sm
.create_session(
temp_dir.path().to_path_buf(),
"test".into(),
SessionType::User,
mode,
)
.await
.unwrap();
let reloaded = sm.get_session(&session.id, false).await.unwrap();
assert_eq!(reloaded.goose_mode, mode);
}
#[tokio::test]
async fn test_goose_mode_update() {
let temp_dir = TempDir::new().unwrap();
let sm = SessionManager::new(temp_dir.path().to_path_buf());
let session = sm
.create_session(
temp_dir.path().to_path_buf(),
"test".into(),
SessionType::User,
GooseMode::default(),
)
.await
.unwrap();
sm.update(&session.id)
.goose_mode(GooseMode::Approve)
.apply()
.await
.unwrap();
let reloaded = sm.get_session(&session.id, false).await.unwrap();
assert_eq!(reloaded.goose_mode, GooseMode::Approve);
}
#[tokio::test]
async fn test_goose_mode_malformed_defaults_to_auto() {
let temp_dir = TempDir::new().unwrap();
let sm = SessionManager::new(temp_dir.path().to_path_buf());
let session = sm
.create_session(
temp_dir.path().to_path_buf(),
"test".into(),
SessionType::User,
GooseMode::Approve,
)
.await
.unwrap();
let pool = &sm.storage().pool;
sqlx::query("UPDATE sessions SET goose_mode = 'garbage' WHERE id = ?")
.bind(&session.id)
.execute(pool)
.await
.unwrap();
let reloaded = sm.get_session(&session.id, false).await.unwrap();
assert_eq!(reloaded.goose_mode, GooseMode::default());
}
}