feat: persist GooseMode per-session via session DB (#7854)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user