diff --git a/crates/goose/src/acp/tkmind/agent.rs b/crates/goose/src/acp/tkmind/agent.rs index 7b3631893..3d14ca5d4 100644 --- a/crates/goose/src/acp/tkmind/agent.rs +++ b/crates/goose/src/acp/tkmind/agent.rs @@ -13,6 +13,7 @@ use tracing::error; use crate::agents::extension::ToolInfo; use crate::agents::extension_manager::get_parameter_names; +use crate::agents::ExtensionLoadResult; use crate::agents::platform_extensions::{chatrecall, projectmemory, PLATFORM_EXTENSIONS}; use crate::agents::ExtensionConfig; use crate::config::permission::PermissionLevel; @@ -70,6 +71,29 @@ pub struct RestartAgentRequest { session_id: String, } +#[derive(Deserialize)] +pub struct ResumeAgentRequest { + session_id: String, + #[serde(default = "default_load_model_and_extensions")] + load_model_and_extensions: bool, +} + +fn default_load_model_and_extensions() -> bool { + true +} + +#[derive(serde::Serialize)] +pub struct ResumeAgentResponse { + pub session: Session, + #[serde(skip_serializing_if = "Option::is_none")] + pub extension_results: Option>, +} + +#[derive(serde::Serialize)] +pub struct RestartAgentResponse { + pub extension_results: Vec, +} + #[derive(Deserialize)] pub struct StartAgentRequest { working_dir: String, @@ -415,26 +439,129 @@ async fn update_working_dir( Ok(()) } +async fn restart_agent_internal( + state: &Arc, + session_id: &str, + session: &Session, +) -> Result, ErrorResponse> { + let _ = state.agent_manager.remove_session(session_id).await; + + let agent = state + .get_agent_for_route(session_id.to_string()) + .await + .map_err(|code| ErrorResponse { + message: "Failed to create new agent during restart".into(), + status: code, + })?; + + let provider_future = agent.restore_provider_from_session(session); + let extensions_future = agent.load_extensions_from_session(session); + + let (provider_result, extension_results) = tokio::join!(provider_future, extensions_future); + provider_result.map_err(|err| ErrorResponse { + message: err.to_string(), + status: StatusCode::INTERNAL_SERVER_ERROR, + })?; + + Ok(extension_results) +} + async fn restart_agent( State(state): State>, Json(payload): Json, -) -> Result<(), (StatusCode, String)> { +) -> Result, ErrorResponse> { let session = state .session_manager() .get_session(&payload.session_id, false) .await - .map_err(|_| (StatusCode::NOT_FOUND, "Session not found".to_string()))?; - let agent = state - .get_agent_for_route(payload.session_id.clone()) + .map_err(|err| { + error!("Failed to get session during restart: {err}"); + ErrorResponse::not_found(format!("Failed to get session: {err}")) + })?; + + let extension_results = + restart_agent_internal(&state, &payload.session_id, &session).await?; + + Ok(Json(RestartAgentResponse { extension_results })) +} + +async fn resume_agent( + State(state): State>, + Json(payload): Json, +) -> Result, ErrorResponse> { + let session = state + .session_manager() + .get_session(&payload.session_id, true) .await - .map_err(|status| (status, "No agent for session id".to_owned()))?; - agent.load_extensions_from_session(&session).await; - Ok(()) + .map_err(|err| { + error!("Failed to resume session {}: {err}", payload.session_id); + ErrorResponse::not_found(format!("Failed to resume session: {err}")) + })?; + + let (extension_results, session) = if payload.load_model_and_extensions { + let agent = state + .get_agent_for_route(payload.session_id.clone()) + .await + .map_err(|code| ErrorResponse { + message: "Failed to get agent for route".into(), + status: code, + })?; + + let provider_changed = agent + .restore_provider_from_session(&session) + .await + .map_err(|err| ErrorResponse { + message: err.to_string(), + status: StatusCode::INTERNAL_SERVER_ERROR, + })?; + + let session = if provider_changed { + state + .session_manager() + .get_session(&payload.session_id, true) + .await + .map_err(|err| ErrorResponse { + message: format!("Failed to re-fetch session: {err}"), + status: StatusCode::INTERNAL_SERVER_ERROR, + })? + } else { + session + }; + + let extension_results = + if let Ok(Some(results)) = state.take_extension_loading_task(&payload.session_id).await + { + tracing::debug!( + "Using background extension loading results for session {}", + payload.session_id + ); + state + .remove_extension_loading_task(&payload.session_id) + .await; + results + } else { + tracing::debug!( + "No background task found, loading extensions for session {}", + payload.session_id + ); + agent.load_extensions_from_session(&session).await + }; + + (Some(extension_results), session) + } else { + (None, session) + }; + + Ok(Json(ResumeAgentResponse { + session, + extension_results, + })) } pub fn routes(state: Arc) -> Router { Router::new() .route("/agent/start", post(start_agent)) + .route("/agent/resume", post(resume_agent)) .route("/agent/tools", get(get_tools)) .route("/agent/update_provider", post(update_agent_provider)) .route("/agent/update_session", post(update_session)) diff --git a/crates/goose/src/acp/tkmind/compat.rs b/crates/goose/src/acp/tkmind/compat.rs index 2f4cc03aa..85443df1a 100644 --- a/crates/goose/src/acp/tkmind/compat.rs +++ b/crates/goose/src/acp/tkmind/compat.rs @@ -86,7 +86,7 @@ async fn get_session( ) -> Result, StatusCode> { state .session_manager() - .get_session(&session_id, false) + .get_session(&session_id, true) .await .map(Json) .map_err(|_| StatusCode::NOT_FOUND)