feat: persist GooseMode per-session via session DB (#7854)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -432,6 +432,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::agent::agent_add_extension,
|
||||
super::routes::agent::agent_remove_extension,
|
||||
super::routes::agent::update_agent_provider,
|
||||
super::routes::agent::update_session,
|
||||
super::routes::action_required::confirm_tool_action,
|
||||
super::routes::reply::reply,
|
||||
super::routes::session::list_sessions,
|
||||
@@ -579,6 +580,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
ModelInfo,
|
||||
ModelConfig,
|
||||
Session,
|
||||
goose::config::goose_mode::GooseMode,
|
||||
SessionInsights,
|
||||
SessionType,
|
||||
SystemInfo,
|
||||
@@ -625,6 +627,7 @@ derive_utoipa!(Icon as IconSchema);
|
||||
goose::agents::types::RetryConfig,
|
||||
goose::agents::types::SuccessCheck,
|
||||
super::routes::agent::UpdateProviderRequest,
|
||||
super::routes::agent::UpdateSessionRequest,
|
||||
super::routes::agent::GetToolsQuery,
|
||||
super::routes::agent::ReadResourceRequest,
|
||||
super::routes::agent::ReadResourceResponse,
|
||||
|
||||
@@ -50,6 +50,12 @@ pub struct UpdateProviderRequest {
|
||||
request_params: Option<std::collections::HashMap<String, serde_json::Value>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, utoipa::ToSchema)]
|
||||
pub struct UpdateSessionRequest {
|
||||
session_id: String,
|
||||
goose_mode: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, utoipa::ToSchema)]
|
||||
pub struct GetToolsQuery {
|
||||
extension_name: Option<String>,
|
||||
@@ -233,9 +239,16 @@ async fn start_agent(
|
||||
let name = "New Chat".to_string();
|
||||
|
||||
let manager = state.session_manager();
|
||||
let config = Config::global();
|
||||
let current_mode = config.get_goose_mode().unwrap_or_default();
|
||||
|
||||
let mut session = manager
|
||||
.create_session(PathBuf::from(&working_dir), name, SessionType::User)
|
||||
.create_session(
|
||||
PathBuf::from(&working_dir),
|
||||
name,
|
||||
SessionType::User,
|
||||
current_mode,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| {
|
||||
error!("Failed to create session: {}", err);
|
||||
@@ -251,6 +264,7 @@ async fn start_agent(
|
||||
.and_then(|r| r.extensions.as_deref());
|
||||
let extensions_to_use =
|
||||
resolve_extensions_for_new_session(recipe_extensions, extension_overrides);
|
||||
|
||||
let mut extension_data = session.extension_data.clone();
|
||||
let extensions_state = EnabledExtensionsState::new(extensions_to_use);
|
||||
if let Err(e) = extensions_state.to_extension_data(&mut extension_data) {
|
||||
@@ -492,10 +506,9 @@ async fn get_tools(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Query(query): Query<GetToolsQuery>,
|
||||
) -> Result<Json<Vec<ToolInfo>>, StatusCode> {
|
||||
let config = Config::global();
|
||||
let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto);
|
||||
let session_id = query.session_id;
|
||||
let agent = state.get_agent_for_route(session_id.clone()).await?;
|
||||
let goose_mode = agent.goose_mode().await;
|
||||
let permission_manager = agent.config.permission_manager.clone();
|
||||
|
||||
let mut tools: Vec<ToolInfo> = agent
|
||||
@@ -594,6 +607,59 @@ async fn update_agent_provider(
|
||||
)
|
||||
})?;
|
||||
|
||||
// Propagate session mode to the new provider
|
||||
let mode = agent.goose_mode().await;
|
||||
agent
|
||||
.update_goose_mode(mode, &payload.session_id)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("Failed to propagate mode to provider: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/agent/update_session",
|
||||
request_body = UpdateSessionRequest,
|
||||
responses(
|
||||
(status = 200, description = "Session updated"),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 500, description = "Internal error")
|
||||
)
|
||||
)]
|
||||
async fn update_session(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(payload): Json<UpdateSessionRequest>,
|
||||
) -> Result<(), (StatusCode, String)> {
|
||||
let agent = state
|
||||
.get_agent_for_route(payload.session_id.clone())
|
||||
.await
|
||||
.map_err(|e| (e, "No agent for session id".to_owned()))?;
|
||||
|
||||
if let Some(mode_str) = payload.goose_mode {
|
||||
let mode: GooseMode = mode_str.parse().map_err(|_| {
|
||||
(
|
||||
StatusCode::BAD_REQUEST,
|
||||
format!("Invalid mode: {}", mode_str),
|
||||
)
|
||||
})?;
|
||||
|
||||
agent
|
||||
.update_goose_mode(mode, &payload.session_id)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("Failed to update mode: {}", e),
|
||||
)
|
||||
})?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1231,6 +1297,7 @@ pub fn routes(state: Arc<AppState>) -> Router {
|
||||
.route("/agent/export_app/{name}", get(export_app))
|
||||
.route("/agent/import_app", post(import_app))
|
||||
.route("/agent/update_provider", post(update_agent_provider))
|
||||
.route("/agent/update_session", post(update_session))
|
||||
.route("/agent/update_from_session", post(update_from_session))
|
||||
.route("/agent/add_extension", post(agent_add_extension))
|
||||
.route("/agent/remove_extension", post(agent_remove_extension))
|
||||
|
||||
Reference in New Issue
Block a user