feat(goose-acp): enable parallel sessions with isolated agent state (#6392)

Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
Adrian Cole
2026-01-15 06:19:17 +08:00
committed by GitHub
parent fb0eca2c36
commit 7d4a6bd8ff
86 changed files with 2594 additions and 1938 deletions
+40 -25
View File
@@ -11,7 +11,6 @@ use axum::{
Json, Router,
};
use goose::agents::ExtensionLoadResult;
use goose::config::PermissionManager;
use base64::Engine;
use goose::agents::ExtensionConfig;
@@ -23,7 +22,7 @@ use goose::recipe::Recipe;
use goose::recipe_deeplink;
use goose::session::extension_data::ExtensionState;
use goose::session::session_manager::SessionType;
use goose::session::{EnabledExtensionsState, Session, SessionManager};
use goose::session::{EnabledExtensionsState, Session};
use goose::{
agents::{extension::ToolInfo, extension_manager::get_parameter_names},
config::permission::PermissionLevel,
@@ -208,17 +207,19 @@ async fn start_agent(
let counter = state.session_counter.fetch_add(1, Ordering::SeqCst) + 1;
let name = format!("New session {}", counter);
let mut session =
SessionManager::create_session(PathBuf::from(&working_dir), name, SessionType::User)
.await
.map_err(|err| {
error!("Failed to create session: {}", err);
goose::posthog::emit_error("session_create_failed", &err.to_string());
ErrorResponse {
message: format!("Failed to create session: {}", err),
status: StatusCode::BAD_REQUEST,
}
})?;
let manager = state.session_manager();
let mut session = manager
.create_session(PathBuf::from(&working_dir), name, SessionType::User)
.await
.map_err(|err| {
error!("Failed to create session: {}", err);
goose::posthog::emit_error("session_create_failed", &err.to_string());
ErrorResponse {
message: format!("Failed to create session: {}", err),
status: StatusCode::BAD_REQUEST,
}
})?;
// Initialize session with extensions (either overrides from hub or global defaults)
let extensions_to_use =
@@ -228,7 +229,8 @@ async fn start_agent(
if let Err(e) = extensions_state.to_extension_data(&mut extension_data) {
tracing::warn!("Failed to initialize session with extensions: {}", e);
} else {
SessionManager::update_session(&session.id)
manager
.update(&session.id)
.extension_data(extension_data.clone())
.apply()
.await
@@ -242,7 +244,8 @@ async fn start_agent(
}
if let Some(recipe) = original_recipe {
SessionManager::update_session(&session.id)
manager
.update(&session.id)
.recipe(Some(recipe))
.apply()
.await
@@ -256,7 +259,8 @@ async fn start_agent(
}
// Refetch session to get all updates
session = SessionManager::get_session(&session.id, false)
session = manager
.get_session(&session.id, false)
.await
.map_err(|err| {
error!("Failed to get updated session: {}", err);
@@ -317,7 +321,9 @@ async fn resume_agent(
) -> Result<Json<ResumeAgentResponse>, ErrorResponse> {
goose::posthog::set_session_context("desktop", true);
let session = SessionManager::get_session(&payload.session_id, true)
let session = state
.session_manager()
.get_session(&payload.session_id, true)
.await
.map_err(|err| {
error!("Failed to resume session {}: {}", payload.session_id, err);
@@ -395,7 +401,9 @@ async fn update_from_session(
message: format!("Failed to get agent: {}", status),
status,
})?;
let session = SessionManager::get_session(&payload.session_id, false)
let session = state
.session_manager()
.get_session(&payload.session_id, false)
.await
.map_err(|err| ErrorResponse {
message: format!("Failed to get session: {}", err),
@@ -453,11 +461,12 @@ async fn get_tools(
) -> Result<Json<Vec<ToolInfo>>, StatusCode> {
let config = Config::global();
let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto);
let agent = state.get_agent_for_route(query.session_id).await?;
let permission_manager = PermissionManager::default();
let session_id = query.session_id;
let agent = state.get_agent_for_route(session_id.clone()).await?;
let permission_manager = agent.config.permission_manager.clone();
let mut tools: Vec<ToolInfo> = agent
.list_tools(query.extension_name)
.list_tools(&session_id, query.extension_name)
.await
.into_iter()
.map(|tool| {
@@ -720,7 +729,9 @@ async fn restart_agent(
) -> Result<Json<RestartAgentResponse>, ErrorResponse> {
let session_id = payload.session_id.clone();
let session = SessionManager::get_session(&session_id, false)
let session = state
.session_manager()
.get_session(&session_id, false)
.await
.map_err(|err| {
error!("Failed to get session during restart: {}", err);
@@ -770,7 +781,9 @@ async fn update_working_dir(
}
// Update the session's working directory
SessionManager::update_session(&session_id)
state
.session_manager()
.update(&session_id)
.working_dir(path)
.apply()
.await
@@ -783,7 +796,9 @@ async fn update_working_dir(
})?;
// Get the updated session and restart the agent
let session = SessionManager::get_session(&session_id, false)
let session = state
.session_manager()
.get_session(&session_id, false)
.await
.map_err(|err| {
error!("Failed to get session after working dir update: {}", err);
@@ -901,7 +916,7 @@ async fn call_tool(
let tool_result = agent
.extension_manager
.dispatch_tool_call(tool_call, CancellationToken::default())
.dispatch_tool_call(&payload.session_id, tool_call, CancellationToken::default())
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
+14 -6
View File
@@ -392,14 +392,25 @@ mod tests {
use super::*;
use axum::{body::Body, http::Request};
use tower::ServiceExt;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[tokio::test(flavor = "multi_thread")]
async fn test_transcribe_endpoint_requires_auth() {
let _guard = env_lock::lock_env([("OPENAI_API_KEY", Some("fake-openai-no-keyring"))]);
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/audio/transcriptions"))
.respond_with(ResponseTemplate::new(401))
.mount(&mock_server)
.await;
let _guard = env_lock::lock_env([
("OPENAI_API_KEY", Some("fake-key")),
("OPENAI_HOST", Some(mock_server.uri().as_str())),
]);
let state = AppState::new().await.unwrap();
let app = routes(state);
// Test without auth header
let request = Request::builder()
.uri("/audio/transcribe")
.method("POST")
@@ -414,10 +425,7 @@ mod tests {
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert!(
response.status() == StatusCode::PRECONDITION_FAILED
|| response.status() == StatusCode::UNAUTHORIZED
);
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test(flavor = "multi_thread")]
@@ -557,7 +557,7 @@ pub async fn init_config() -> Result<Json<String>, StatusCode> {
pub async fn upsert_permissions(
Json(query): Json<UpsertPermissionsQuery>,
) -> Result<Json<String>, StatusCode> {
let mut permission_manager = goose::config::PermissionManager::default();
let permission_manager = goose::config::PermissionManager::instance();
for tool_permission in &query.tool_permissions {
permission_manager.update_user_permission(
+1 -1
View File
@@ -23,7 +23,7 @@ use axum::Router;
// Function to configure all routes
pub fn configure(state: Arc<crate::state::AppState>, secret_key: String) -> Router {
Router::new()
.merge(status::routes())
.merge(status::routes(state.clone()))
.merge(reply::routes(state.clone()))
.merge(action_required::routes(state.clone()))
.merge(agent::routes(state.clone()))
+5 -2
View File
@@ -9,7 +9,6 @@ use axum::{extract::State, http::StatusCode, routing::post, Json, Router};
use goose::recipe::local_recipes;
use goose::recipe::validate_recipe::validate_recipe_template_from_content;
use goose::recipe::Recipe;
use goose::session::SessionManager;
use goose::{recipe_deeplink, slash_commands};
use serde::{Deserialize, Serialize};
@@ -168,7 +167,11 @@ async fn create_recipe(
request.session_id
);
let session = match SessionManager::get_session(&request.session_id, true).await {
let session = match state
.session_manager()
.get_session(&request.session_id, true)
.await
{
Ok(session) => session,
Err(e) => {
tracing::error!("Failed to get session: {}", e);
+12 -7
View File
@@ -146,8 +146,9 @@ pub enum MessageEvent {
Ping,
}
async fn get_token_state(session_id: &str) -> TokenState {
SessionManager::get_session(session_id, false)
async fn get_token_state(session_manager: &SessionManager, session_id: &str) -> TokenState {
session_manager
.get_session(session_id, false)
.await
.map(|session| TokenState {
input_tokens: session.input_tokens.unwrap_or(0),
@@ -258,7 +259,7 @@ pub async fn reply(
}
};
let session = match SessionManager::get_session(&session_id, true).await {
let session = match state.session_manager().get_session(&session_id, true).await {
Ok(metadata) => metadata,
Err(e) => {
tracing::error!("Failed to read session for {}: {}", session_id, e);
@@ -284,7 +285,11 @@ pub async fn reply(
let mut all_messages = match conversation_so_far {
Some(history) => {
let conv = Conversation::new_unvalidated(history);
if let Err(e) = SessionManager::replace_conversation(&session_id, &conv).await {
if let Err(e) = state
.session_manager()
.replace_conversation(&session_id, &conv)
.await
{
tracing::warn!(
"Failed to replace session conversation for {}: {}",
session_id,
@@ -339,7 +344,7 @@ pub async fn reply(
all_messages.push(message.clone());
let token_state = get_token_state(&session_id).await;
let token_state = get_token_state(state.session_manager(), &session_id).await;
stream_event(MessageEvent::Message { message, token_state }, &tx, &cancel_token).await;
}
@@ -385,7 +390,7 @@ pub async fn reply(
let session_duration = session_start.elapsed();
if let Ok(session) = SessionManager::get_session(&session_id, true).await {
if let Ok(session) = state.session_manager().get_session(&session_id, true).await {
let total_tokens = session.total_tokens.unwrap_or(0);
tracing::info!(
counter.goose.session_completions = 1,
@@ -433,7 +438,7 @@ pub async fn reply(
);
}
let final_token_state = get_token_state(&session_id).await;
let final_token_state = get_token_state(state.session_manager(), &session_id).await;
let _ = stream_event(
MessageEvent::Finish {
+60 -19
View File
@@ -13,7 +13,7 @@ use goose::agents::ExtensionConfig;
use goose::recipe::Recipe;
use goose::session::extension_data::ExtensionState;
use goose::session::session_manager::SessionInsights;
use goose::session::{EnabledExtensionsState, Session, SessionManager};
use goose::session::{EnabledExtensionsState, Session};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
@@ -91,8 +91,12 @@ const MAX_NAME_LENGTH: usize = 200;
),
tag = "Session Management"
)]
async fn list_sessions() -> Result<Json<SessionListResponse>, StatusCode> {
let sessions = SessionManager::list_sessions()
async fn list_sessions(
State(state): State<Arc<AppState>>,
) -> Result<Json<SessionListResponse>, StatusCode> {
let sessions = state
.session_manager()
.list_sessions()
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
@@ -116,8 +120,13 @@ async fn list_sessions() -> Result<Json<SessionListResponse>, StatusCode> {
),
tag = "Session Management"
)]
async fn get_session(Path(session_id): Path<String>) -> Result<Json<Session>, StatusCode> {
let session = SessionManager::get_session(&session_id, true)
async fn get_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> Result<Json<Session>, StatusCode> {
let session = state
.session_manager()
.get_session(&session_id, true)
.await
.map_err(|_| StatusCode::NOT_FOUND)?;
@@ -136,8 +145,12 @@ async fn get_session(Path(session_id): Path<String>) -> Result<Json<Session>, St
),
tag = "Session Management"
)]
async fn get_session_insights() -> Result<Json<SessionInsights>, StatusCode> {
let insights = SessionManager::get_insights()
async fn get_session_insights(
State(state): State<Arc<AppState>>,
) -> Result<Json<SessionInsights>, StatusCode> {
let insights = state
.session_manager()
.get_insights()
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
Ok(Json(insights))
@@ -163,6 +176,7 @@ async fn get_session_insights() -> Result<Json<SessionInsights>, StatusCode> {
tag = "Session Management"
)]
async fn update_session_name(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
Json(request): Json<UpdateSessionNameRequest>,
) -> Result<StatusCode, StatusCode> {
@@ -174,7 +188,9 @@ async fn update_session_name(
return Err(StatusCode::BAD_REQUEST);
}
SessionManager::update_session(&session_id)
state
.session_manager()
.update(&session_id)
.user_provided_name(name.to_string())
.apply()
.await
@@ -207,7 +223,9 @@ async fn update_session_user_recipe_values(
Path(session_id): Path<String>,
Json(request): Json<UpdateSessionUserRecipeValuesRequest>,
) -> Result<Json<UpdateSessionUserRecipeValuesResponse>, ErrorResponse> {
SessionManager::update_session(&session_id)
state
.session_manager()
.update(&session_id)
.user_recipe_values(Some(request.user_recipe_values))
.apply()
.await
@@ -216,7 +234,9 @@ async fn update_session_user_recipe_values(
status: StatusCode::INTERNAL_SERVER_ERROR,
})?;
let session = SessionManager::get_session(&session_id, false)
let session = state
.session_manager()
.get_session(&session_id, false)
.await
.map_err(|err| ErrorResponse {
message: err.to_string(),
@@ -270,8 +290,13 @@ async fn update_session_user_recipe_values(
),
tag = "Session Management"
)]
async fn delete_session(Path(session_id): Path<String>) -> Result<StatusCode, StatusCode> {
SessionManager::delete_session(&session_id)
async fn delete_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> Result<StatusCode, StatusCode> {
state
.session_manager()
.delete_session(&session_id)
.await
.map_err(|e| {
if e.to_string().contains("not found") {
@@ -301,8 +326,13 @@ async fn delete_session(Path(session_id): Path<String>) -> Result<StatusCode, St
),
tag = "Session Management"
)]
async fn export_session(Path(session_id): Path<String>) -> Result<Json<String>, StatusCode> {
let exported = SessionManager::export_session(&session_id)
async fn export_session(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> Result<Json<String>, StatusCode> {
let exported = state
.session_manager()
.export_session(&session_id)
.await
.map_err(|_| StatusCode::NOT_FOUND)?;
@@ -325,9 +355,12 @@ async fn export_session(Path(session_id): Path<String>) -> Result<Json<String>,
tag = "Session Management"
)]
async fn import_session(
State(state): State<Arc<AppState>>,
Json(request): Json<ImportSessionRequest>,
) -> Result<Json<Session>, StatusCode> {
let session = SessionManager::import_session(&request.json)
let session = state
.session_manager()
.import_session(&request.json)
.await
.map_err(|_| StatusCode::BAD_REQUEST)?;
@@ -354,12 +387,15 @@ async fn import_session(
tag = "Session Management"
)]
async fn edit_message(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
Json(request): Json<EditMessageRequest>,
) -> Result<Json<EditMessageResponse>, StatusCode> {
let manager = state.session_manager();
match request.edit_type {
EditType::Fork => {
let new_session = SessionManager::copy_session(&session_id, "(edited)".to_string())
let new_session = manager
.copy_session(&session_id, "(edited)".to_string())
.await
.map_err(|e| {
tracing::error!("Failed to copy session: {}", e);
@@ -367,7 +403,8 @@ async fn edit_message(
StatusCode::INTERNAL_SERVER_ERROR
})?;
SessionManager::truncate_conversation(&new_session.id, request.timestamp)
manager
.truncate_conversation(&new_session.id, request.timestamp)
.await
.map_err(|e| {
tracing::error!("Failed to truncate conversation: {}", e);
@@ -380,7 +417,8 @@ async fn edit_message(
}))
}
EditType::Edit => {
SessionManager::truncate_conversation(&session_id, request.timestamp)
manager
.truncate_conversation(&session_id, request.timestamp)
.await
.map_err(|e| {
tracing::error!("Failed to truncate conversation: {}", e);
@@ -419,9 +457,12 @@ pub struct SessionExtensionsResponse {
tag = "Session Management"
)]
async fn get_session_extensions(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> Result<Json<SessionExtensionsResponse>, StatusCode> {
let session = SessionManager::get_session(&session_id, false)
let session = state
.session_manager()
.get_session(&session_id, false)
.await
.map_err(|_| StatusCode::NOT_FOUND)?;
+11 -3
View File
@@ -1,8 +1,12 @@
use axum::body::Body;
use axum::extract::State;
use axum::http::HeaderValue;
use axum::response::IntoResponse;
use axum::{extract::Path, http::StatusCode, routing::get, Json, Router};
use goose::session::{generate_diagnostics, get_system_info, SystemInfo};
use std::sync::Arc;
use crate::state::AppState;
#[utoipa::path(get, path = "/status",
responses(
@@ -28,8 +32,11 @@ async fn system_info() -> Json<SystemInfo> {
(status = 500, description = "Failed to generate diagnostics"),
)
)]
async fn diagnostics(Path(session_id): Path<String>) -> impl IntoResponse {
match generate_diagnostics(&session_id).await {
async fn diagnostics(
State(state): State<Arc<AppState>>,
Path(session_id): Path<String>,
) -> impl IntoResponse {
match generate_diagnostics(state.session_manager(), &session_id).await {
Ok(zip_data) => {
let filename = format!("attachment; filename=\"diagnostics_{}.zip\"", session_id);
let headers = [
@@ -48,9 +55,10 @@ async fn diagnostics(Path(session_id): Path<String>) -> impl IntoResponse {
Err(_) => Err(StatusCode::INTERNAL_SERVER_ERROR),
}
}
pub fn routes() -> Router {
pub fn routes(state: Arc<AppState>) -> Router {
Router::new()
.route("/status", get(status))
.route("/system_info", get(system_info))
.route("/diagnostics/{session_id}", get(diagnostics))
.with_state(state)
}
+5
View File
@@ -1,6 +1,7 @@
use axum::http::StatusCode;
use goose::execution::manager::AgentManager;
use goose::scheduler_trait::SchedulerTrait;
use goose::session::SessionManager;
use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::sync::atomic::AtomicUsize;
@@ -81,6 +82,10 @@ impl AppState {
self.agent_manager.scheduler()
}
pub fn session_manager(&self) -> &SessionManager {
self.agent_manager.session_manager()
}
pub async fn set_recipe_file_hash_map(&self, hash_map: HashMap<String, PathBuf>) {
let mut map = self.recipe_file_hash_map.lock().await;
*map = hash_map;