diff --git a/Cargo.lock b/Cargo.lock index 76d9aefd..1346958d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4573,7 +4573,7 @@ dependencies = [ [[package]] name = "goose-sdk" -version = "1.29.0" +version = "1.30.0" dependencies = [ "agent-client-protocol-schema", "sacp", diff --git a/crates/goose-acp/acp-meta.json b/crates/goose-acp/acp-meta.json index 53de01bf..f5df690d 100644 --- a/crates/goose-acp/acp-meta.json +++ b/crates/goose-acp/acp-meta.json @@ -49,6 +49,46 @@ "method": "_goose/config/extensions", "requestType": "GetExtensionsRequest", "responseType": "GetExtensionsResponse" + }, + { + "method": "_goose/session/provider/update", + "requestType": "UpdateProviderRequest", + "responseType": "UpdateProviderResponse" + }, + { + "method": "_goose/providers/list", + "requestType": "ListProvidersRequest", + "responseType": "ListProvidersResponse" + }, + { + "method": "_goose/config/read", + "requestType": "ReadConfigRequest", + "responseType": "ReadConfigResponse" + }, + { + "method": "_goose/config/upsert", + "requestType": "UpsertConfigRequest", + "responseType": "EmptyResponse" + }, + { + "method": "_goose/config/remove", + "requestType": "RemoveConfigRequest", + "responseType": "EmptyResponse" + }, + { + "method": "_goose/secret/check", + "requestType": "CheckSecretRequest", + "responseType": "CheckSecretResponse" + }, + { + "method": "_goose/secret/upsert", + "requestType": "UpsertSecretRequest", + "responseType": "EmptyResponse" + }, + { + "method": "_goose/secret/remove", + "requestType": "RemoveSecretRequest", + "responseType": "EmptyResponse" } ] } diff --git a/crates/goose-acp/acp-schema.json b/crates/goose-acp/acp-schema.json index a9e50a0f..b0e2d6ba 100644 --- a/crates/goose-acp/acp-schema.json +++ b/crates/goose-acp/acp-schema.json @@ -62,7 +62,7 @@ "properties": { "tools": { "type": "array", - "items": true, + "items": {}, "description": "Array of tool info objects with `name`, `description`, `parameters`, and optional `permission`." } }, @@ -234,7 +234,7 @@ "properties": { "extensions": { "type": "array", - "items": true, + "items": {}, "description": "Array of ExtensionEntry objects with `enabled` flag and config details." }, "warnings": { @@ -252,6 +252,212 @@ "x-side": "agent", "x-method": "_goose/config/extensions" }, + "UpdateProviderRequest": { + "type": "object", + "properties": { + "sessionId": { + "type": "string" + }, + "provider": { + "type": "string" + }, + "model": { + "type": [ + "string", + "null" + ] + }, + "contextLimit": { + "type": [ + "integer", + "null" + ], + "format": "uint", + "minimum": 0 + }, + "requestParams": { + "type": [ + "object", + "null" + ], + "additionalProperties": {} + } + }, + "required": [ + "sessionId", + "provider" + ], + "description": "Atomically update the provider for a live session.", + "x-side": "agent", + "x-method": "_goose/session/provider/update" + }, + "UpdateProviderResponse": { + "type": "object", + "properties": { + "configOptions": { + "type": "array", + "items": {}, + "description": "Refreshed session config options after the provider/model change." + } + }, + "required": [ + "configOptions" + ], + "description": "Provider update response.", + "x-side": "agent", + "x-method": "_goose/session/provider/update" + }, + "ListProvidersRequest": { + "type": "object", + "description": "List providers available through goose, including the config-default sentinel.", + "x-side": "agent", + "x-method": "_goose/providers/list" + }, + "ListProvidersResponse": { + "type": "object", + "properties": { + "providers": { + "type": "array", + "items": { + "$ref": "#/$defs/ProviderListEntry" + } + } + }, + "required": [ + "providers" + ], + "description": "Provider list response.", + "x-side": "agent", + "x-method": "_goose/providers/list" + }, + "ProviderListEntry": { + "type": "object", + "properties": { + "id": { + "type": "string" + }, + "label": { + "type": "string" + } + }, + "required": [ + "id", + "label" + ] + }, + "ReadConfigRequest": { + "type": "object", + "properties": { + "key": { + "type": "string" + } + }, + "required": [ + "key" + ], + "description": "Read a single non-secret config value.", + "x-side": "agent", + "x-method": "_goose/config/read" + }, + "ReadConfigResponse": { + "type": "object", + "properties": { + "value": { + "default": null + } + }, + "description": "Config read response.", + "x-side": "agent", + "x-method": "_goose/config/read" + }, + "UpsertConfigRequest": { + "type": "object", + "properties": { + "key": { + "type": "string" + }, + "value": {} + }, + "required": [ + "key", + "value" + ], + "description": "Upsert a single non-secret config value.", + "x-side": "agent", + "x-method": "_goose/config/upsert" + }, + "RemoveConfigRequest": { + "type": "object", + "properties": { + "key": { + "type": "string" + } + }, + "required": [ + "key" + ], + "description": "Remove a single non-secret config value.", + "x-side": "agent", + "x-method": "_goose/config/remove" + }, + "CheckSecretRequest": { + "type": "object", + "properties": { + "key": { + "type": "string" + } + }, + "required": [ + "key" + ], + "description": "Check whether a secret exists. Never returns the actual value.", + "x-side": "agent", + "x-method": "_goose/secret/check" + }, + "CheckSecretResponse": { + "type": "object", + "properties": { + "exists": { + "type": "boolean" + } + }, + "required": [ + "exists" + ], + "description": "Secret check response.", + "x-side": "agent", + "x-method": "_goose/secret/check" + }, + "UpsertSecretRequest": { + "type": "object", + "properties": { + "key": { + "type": "string" + }, + "value": {} + }, + "required": [ + "key", + "value" + ], + "description": "Set a secret value (write-only).", + "x-side": "agent", + "x-method": "_goose/secret/upsert" + }, + "RemoveSecretRequest": { + "type": "object", + "properties": { + "key": { + "type": "string" + } + }, + "required": [ + "key" + ], + "description": "Remove a secret.", + "x-side": "agent", + "x-method": "_goose/secret/remove" + }, "ExtRequest": { "properties": { "id": { @@ -353,6 +559,78 @@ ], "description": "Params for _goose/config/extensions", "title": "GetExtensionsRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/UpdateProviderRequest" + } + ], + "description": "Params for _goose/session/provider/update", + "title": "UpdateProviderRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/ListProvidersRequest" + } + ], + "description": "Params for _goose/providers/list", + "title": "ListProvidersRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/ReadConfigRequest" + } + ], + "description": "Params for _goose/config/read", + "title": "ReadConfigRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/UpsertConfigRequest" + } + ], + "description": "Params for _goose/config/upsert", + "title": "UpsertConfigRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/RemoveConfigRequest" + } + ], + "description": "Params for _goose/config/remove", + "title": "RemoveConfigRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CheckSecretRequest" + } + ], + "description": "Params for _goose/secret/check", + "title": "CheckSecretRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/UpsertSecretRequest" + } + ], + "description": "Params for _goose/secret/upsert", + "title": "UpsertSecretRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/RemoveSecretRequest" + } + ], + "description": "Params for _goose/secret/remove", + "title": "RemoveSecretRequest" } ] }, @@ -439,6 +717,38 @@ } ], "title": "GetExtensionsResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/UpdateProviderResponse" + } + ], + "title": "UpdateProviderResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/ListProvidersResponse" + } + ], + "title": "ListProvidersResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/ReadConfigResponse" + } + ], + "title": "ReadConfigResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CheckSecretResponse" + } + ], + "title": "CheckSecretResponse" } ] }, diff --git a/crates/goose-acp/src/bin/generate_acp_schema.rs b/crates/goose-acp/src/bin/generate_acp_schema.rs index f4e8ac9e..0e9bf879 100644 --- a/crates/goose-acp/src/bin/generate_acp_schema.rs +++ b/crates/goose-acp/src/bin/generate_acp_schema.rs @@ -35,6 +35,13 @@ fn main() { } } + // Replace `true` with `{}` throughout $defs. Both mean "accept any value" in + // JSON Schema, but many TS codegen tools (e.g. @hey-api/openapi-ts Zod plugin) + // silently drop properties whose schema is the bare `true` literal. + for def in defs.values_mut() { + replace_true_schemas(def); + } + // Annotate $defs entries with x-method/x-side. Only set x-method for types // used by exactly one method (shared types like EmptyResponse skip x-method). for (name, methods_list) in &type_methods { @@ -181,3 +188,31 @@ fn main() { println!("{json_str}"); } + +/// Recursively replace `true` with `{}` in a JSON value. +/// +/// In JSON Schema, `true` and `{}` both mean "accept any value", but many +/// TypeScript codegen tools only handle the object form. +fn replace_true_schemas(value: &mut Value) { + match value { + Value::Object(map) => { + for v in map.values_mut() { + if *v == Value::Bool(true) { + *v = json!({}); + } else { + replace_true_schemas(v); + } + } + } + Value::Array(arr) => { + for v in arr.iter_mut() { + if *v == Value::Bool(true) { + *v = json!({}); + } else { + replace_true_schemas(v); + } + } + } + _ => {} + } +} diff --git a/crates/goose-acp/src/server.rs b/crates/goose-acp/src/server.rs index 9751fdc7..8c1e0197 100644 --- a/crates/goose-acp/src/server.rs +++ b/crates/goose-acp/src/server.rs @@ -3,7 +3,8 @@ use crate::fs::AcpTools; use crate::tools::AcpAwareToolMeta; use anyhow::Result; use fs_err as fs; -use goose::acp::PermissionDecision; +use futures::future::BoxFuture; +use goose::acp::{PermissionDecision, ACP_CURRENT_MODEL}; use goose::agents::extension::{Envs, PLATFORM_EXTENSIONS}; use goose::agents::mcp_client::McpClientTrait; use goose::agents::platform_extensions::developer::DeveloperClient; @@ -20,9 +21,8 @@ use goose::mcp_utils::ToolResult; use goose::permission::permission_confirmation::PrincipalType; use goose::permission::{Permission, PermissionConfirmation}; use goose::providers::base::Provider; -use goose::providers::provider_registry::ProviderConstructor; use goose::session::session_manager::SessionType; -use goose::session::{Session, SessionManager}; +use goose::session::{EnabledExtensionsState, Session, SessionManager}; use goose_acp_macros::custom_methods; use rmcp::model::{CallToolResult, RawContent, ResourceContents, Role}; use sacp::schema::{ @@ -57,6 +57,19 @@ use tokio_util::sync::CancellationToken; use tracing::{debug, error, info, warn}; use url::Url; +pub type AcpProviderFactory = Arc< + dyn Fn( + String, + goose::model::ModelConfig, + Vec, + ) -> BoxFuture<'static, Result>> + + Send + + Sync, +>; + +const DEFAULT_PROVIDER_ID: &str = "goose"; +const DEFAULT_PROVIDER_LABEL: &str = "Goose (Default)"; + struct GooseAcpSession { agent: Arc, messages: Conversation, @@ -66,7 +79,7 @@ struct GooseAcpSession { pub struct GooseAcpAgent { sessions: Arc>>, - provider_factory: ProviderConstructor, + provider_factory: AcpProviderFactory, builtins: Vec, client_fs_capabilities: OnceCell, client_terminal: OnceCell, @@ -328,6 +341,56 @@ async fn build_model_state(provider: &dyn Provider) -> Result) -> Vec { + let mut providers = goose::providers::providers() + .await + .into_iter() + .map(|(metadata, _)| ProviderListEntry { + id: metadata.name, + label: metadata.display_name, + }) + .collect::>(); + providers.sort_by(|left, right| left.id.cmp(&right.id)); + providers.dedup_by(|left, right| left.id == right.id); + + if let Some(current_provider) = current_provider { + if current_provider != DEFAULT_PROVIDER_ID + && !providers + .iter() + .any(|provider| provider.id == current_provider) + { + providers.push(ProviderListEntry { + id: current_provider.to_string(), + label: current_provider.to_string(), + }); + providers.sort_by(|left, right| left.id.cmp(&right.id)); + } + } + + let mut entries = Vec::with_capacity(providers.len() + 1); + entries.push(ProviderListEntry { + id: DEFAULT_PROVIDER_ID.to_string(), + label: DEFAULT_PROVIDER_LABEL.to_string(), + }); + entries.extend(providers); + entries +} + +async fn build_provider_options(current_provider: Option<&str>) -> Vec { + list_provider_entries(current_provider) + .await + .into_iter() + .map(|provider| SessionConfigSelectOption::new(provider.id, provider.label)) + .collect() +} + +fn session_provider_selection(session: &Session) -> &str { + session + .provider_name + .as_deref() + .unwrap_or(DEFAULT_PROVIDER_ID) +} + fn build_mode_state(current_mode: GooseMode) -> Result { let mut available = Vec::with_capacity(GooseMode::VARIANTS.len()); for &name in GooseMode::VARIANTS { @@ -348,6 +411,8 @@ fn build_mode_state(current_mode: GooseMode) -> Result, ) -> Vec { let mode_options: Vec = mode_state .available_modes @@ -363,6 +428,12 @@ fn build_config_options( .map(|m| SessionConfigSelectOption::new(m.model_id.0.clone(), m.name.clone())) .collect(); vec![ + SessionConfigOption::select( + "provider", + "Provider", + provider_selection.to_string(), + provider_options, + ), SessionConfigOption::select( "mode", "Mode", @@ -387,7 +458,7 @@ impl GooseAcpAgent { // TODO: goose reads Paths::in_state_dir globally (e.g. RequestLog), ignoring this data_dir. pub async fn new( - provider_factory: ProviderConstructor, + provider_factory: AcpProviderFactory, builtins: Vec, data_dir: std::path::PathBuf, config_dir: std::path::PathBuf, @@ -411,6 +482,19 @@ impl GooseAcpAgent { }) } + fn load_config(&self) -> Result { + Config::new(self.config_dir.join(CONFIG_YAML_NAME), "goose").map_err(Into::into) + } + + async fn create_provider( + &self, + provider_name: &str, + model_config: goose::model::ModelConfig, + extensions: Vec, + ) -> Result> { + (self.provider_factory)(provider_name.to_string(), model_config, extensions).await + } + async fn create_agent_for_session( &self, cx: Option<&ConnectionTo>, @@ -852,6 +936,16 @@ impl GooseAcpAgent { ) -> Result { debug!(?args, "new session request"); + // Allow the client to request a specific provider via _meta.provider, + // avoiding the double-create when the client would otherwise call + // _goose/session/provider/update immediately after. + let requested_provider = args + .meta + .as_ref() + .and_then(|m| m.get("provider")) + .and_then(|v| v.as_str()) + .map(|s| s.to_string()); + let goose_session = self .session_manager .create_session( @@ -865,6 +959,30 @@ impl GooseAcpAgent { sacp::Error::internal_error().data(format!("Failed to create session: {}", e)) })?; + if let Some(ref provider_name) = requested_provider { + self.session_manager + .update(&goose_session.id) + .provider_name(provider_name) + .apply() + .await + .map_err(|e| { + sacp::Error::internal_error() + .data(format!("Failed to set provider on session: {}", e)) + })?; + } + + // Reload the session so init_provider sees the updated provider_name. + let goose_session = if requested_provider.is_some() { + self.session_manager + .get_session(&goose_session.id, false) + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to reload session: {}", e)) + })? + } else { + goose_session + }; + let session_id = SessionId::new(goose_session.id.clone()); let agent = self @@ -879,7 +997,6 @@ impl GooseAcpAgent { .map_err(|e| { sacp::Error::internal_error().data(format!("Failed to set provider: {}", e)) })?; - Self::add_mcp_extensions(&agent, args.mcp_servers, &goose_session.id).await?; let session = GooseAcpSession { @@ -901,25 +1018,55 @@ impl GooseAcpAgent { let model_state = build_model_state(&*provider).await?; let mode_state = build_mode_state(self.goose_mode)?; + let provider_selection = session_provider_selection(&goose_session).to_string(); + let session_id_for_response = SessionId::new(goose_session.id); + let provider_options = build_provider_options(Some(provider.get_name())).await; - Ok(NewSessionResponse::new(SessionId::new(goose_session.id)) + Ok(NewSessionResponse::new(session_id_for_response) .models(model_state.clone()) .modes(mode_state.clone()) - .config_options(build_config_options(&mode_state, &model_state))) + .config_options(build_config_options( + &mode_state, + &model_state, + &provider_selection, + provider_options, + ))) } async fn init_provider(&self, agent: &Agent, session: &Session) -> Result> { + let config = self.load_config()?; + let global_provider = config.get_goose_provider().ok(); + let provider_override = session + .provider_name + .as_deref() + .filter(|provider| *provider != DEFAULT_PROVIDER_ID); + let provider_name = provider_override + .map(ToOwned::to_owned) + .or_else(|| global_provider.clone()) + .ok_or_else(|| anyhow::anyhow!("Could not configure agent: missing provider"))?; + let explicitly_switched = + provider_override.is_some() && provider_override != global_provider.as_deref(); let model_config = match &session.model_config { - Some(config) => config.clone(), + Some(model_config) => model_config.clone(), + None if explicitly_switched => { + // The provider was set via _meta.provider (or similar) without an + // explicit model. Use the provider's own default model from the + // registry so we don't leak the global config model (which belongs + // to a different provider) into this one. + let entry = goose::providers::get_from_registry(&provider_name).await?; + let default_model = &entry.metadata().default_model; + goose::model::ModelConfig::new(default_model)?.with_canonical_limits(&provider_name) + } None => { - let config_path = self.config_dir.join(CONFIG_YAML_NAME); - let config = Config::new(&config_path, "goose")?; let model_id = config.get_goose_model()?; - let provider_name = config.get_goose_provider()?; goose::model::ModelConfig::new(&model_id)?.with_canonical_limits(&provider_name) } }; - let provider = (self.provider_factory)(model_config, Vec::new()).await?; + let extensions = + EnabledExtensionsState::extensions_or_default(Some(&session.extension_data), &config); + let provider = self + .create_provider(&provider_name, model_config, extensions) + .await?; agent.update_provider(provider.clone(), &session.id).await?; Ok(provider) } @@ -993,7 +1140,6 @@ impl GooseAcpAgent { sacp::Error::resource_not_found(Some(session_id.clone())) .data(format!("Session not found: {}", session_id)) })?; - let loaded_mode = goose_session.goose_mode; let acp_session_id = SessionId::new(session_id.clone()); @@ -1003,7 +1149,6 @@ impl GooseAcpAgent { .map_err(|e| { sacp::Error::internal_error().data(format!("Failed to create agent: {}", e)) })?; - let provider = self .init_provider(&agent, &goose_session) .await @@ -1011,8 +1156,15 @@ impl GooseAcpAgent { sacp::Error::internal_error().data(format!("Failed to set provider: {}", e)) })?; + agent + .update_goose_mode(loaded_mode, &session_id) + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to restore mode: {}", e)) + })?; Self::add_mcp_extensions(&agent, args.mcp_servers, &session_id).await?; + let provider_selection = session_provider_selection(&goose_session).to_string(); let conversation = goose_session.conversation.ok_or_else(|| { sacp::Error::internal_error() .data(format!("Session {} has no conversation data", session_id)) @@ -1095,11 +1247,17 @@ impl GooseAcpAgent { let model_state = build_model_state(&*provider).await?; let mode_state = build_mode_state(goose_mode)?; + let provider_options = build_provider_options(Some(provider.get_name())).await; Ok(LoadSessionResponse::new() .models(model_state.clone()) .modes(mode_state.clone()) - .config_options(build_config_options(&mode_state, &model_state))) + .config_options(build_config_options( + &mode_state, + &model_state, + &provider_selection, + provider_options, + ))) } async fn on_prompt( @@ -1167,7 +1325,6 @@ impl GooseAcpAgent { if let Some(session) = sessions.get_mut(&session_id) { session.cancel_token = None; } - Ok(PromptResponse::new(if was_cancelled { StopReason::Cancelled } else { @@ -1198,25 +1355,28 @@ impl GooseAcpAgent { session_id: &str, model_id: &str, ) -> Result { - let config_path = self.config_dir.join(CONFIG_YAML_NAME); - let config = Config::new(&config_path, "goose").map_err(|e| { + let config = self.load_config().map_err(|e| { sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) })?; - let provider_name = config.get_goose_provider().map_err(|_| { - sacp::Error::internal_error().data("No provider configured".to_string()) + let agent = self.get_session_agent(session_id, None).await?; + let current_provider = agent.provider().await.map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to get provider: {}", e)) })?; + let provider_name = current_provider.get_name().to_string(); + let extensions = + EnabledExtensionsState::for_session(&self.session_manager, session_id, &config).await; let model_config = goose::model::ModelConfig::new(model_id) .map_err(|e| { sacp::Error::invalid_params().data(format!("Invalid model config: {}", e)) })? .with_canonical_limits(&provider_name); - let provider = (self.provider_factory)(model_config, Vec::new()) + let provider = self + .create_provider(&provider_name, model_config, extensions) .await .map_err(|e| { sacp::Error::internal_error().data(format!("Failed to create provider: {}", e)) })?; - let agent = self.get_session_agent(session_id, None).await?; agent .update_provider(provider, session_id) .await @@ -1224,6 +1384,14 @@ impl GooseAcpAgent { sacp::Error::internal_error().data(format!("Failed to update provider: {}", e)) })?; + let mode = agent.goose_mode().await; + agent + .update_goose_mode(mode, session_id) + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to propagate mode: {}", e)) + })?; + info!(session_id = %session_id, model_id = %model_id, "Model switched"); Ok(SetSessionModelResponse::new()) } @@ -1232,6 +1400,11 @@ impl GooseAcpAgent { &self, session_id: &SessionId, ) -> Result<(SessionNotification, Vec), sacp::Error> { + let session = self + .session_manager + .get_session(&session_id.0, false) + .await + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; let agent = self.get_session_agent(&session_id.0, None).await?; let provider = agent.provider().await.map_err(|e| { sacp::Error::internal_error().data(format!("Failed to get provider: {}", e)) @@ -1239,7 +1412,13 @@ impl GooseAcpAgent { let goose_mode = agent.goose_mode().await; let model_state = build_model_state(&*provider).await?; let mode_state = build_mode_state(goose_mode)?; - let config_options = build_config_options(&mode_state, &model_state); + let provider_options = build_provider_options(Some(provider.get_name())).await; + let config_options = build_config_options( + &mode_state, + &model_state, + session_provider_selection(&session), + provider_options, + ); let notification = SessionNotification::new( session_id.clone(), SessionUpdate::ConfigOptionUpdate(ConfigOptionUpdate::new(config_options.clone())), @@ -1267,6 +1446,120 @@ impl GooseAcpAgent { Ok(SetSessionModeResponse::new()) } + async fn update_provider( + &self, + session_id: &str, + provider_name: &str, + model_name: Option<&str>, + context_limit: Option, + request_params: Option>, + ) -> Result, sacp::Error> { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + let agent = self.get_session_agent(session_id, None).await?; + let current_provider = agent.provider().await.map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to get provider: {}", e)) + })?; + let current_provider_name = current_provider.get_name(); + let current_model = current_provider.get_model_config().model_name; + let has_default_overrides = + model_name.is_some() || context_limit.is_some() || request_params.is_some(); + let use_default_provider = provider_name == DEFAULT_PROVIDER_ID; + let resolved_provider_name = if use_default_provider { + config.get_goose_provider().map_err(|e| { + sacp::Error::internal_error().data(format!( + "Failed to resolve default provider from config: {}", + e + )) + })? + } else { + provider_name.to_string() + }; + let is_changing_provider = resolved_provider_name != current_provider_name; + let default_model = if let Some(model_name) = model_name { + model_name.to_string() + } else if use_default_provider { + config.get_goose_model().map_err(|e| { + sacp::Error::internal_error().data(format!( + "Failed to resolve default model from config: {}", + e + )) + })? + } else if is_changing_provider { + ACP_CURRENT_MODEL.to_string() + } else { + current_model + }; + let model = model_name.unwrap_or(&default_model); + let model_config = goose::model::ModelConfig::new(model) + .map_err(|e| { + sacp::Error::invalid_params().data(format!("Invalid model config: {}", e)) + })? + .with_canonical_limits(&resolved_provider_name) + .with_context_limit(context_limit) + .with_request_params(request_params); + let extensions = + EnabledExtensionsState::for_session(&self.session_manager, session_id, &config).await; + let new_provider = self + .create_provider(&resolved_provider_name, model_config, extensions) + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to create provider: {}", e)) + })?; + + agent + .update_provider(new_provider, session_id) + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to update provider: {}", e)) + })?; + + let mode = agent.goose_mode().await; + agent + .update_goose_mode(mode, session_id) + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to propagate mode: {}", e)) + })?; + + let provider = agent.provider().await.map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to get provider: {}", e)) + })?; + + if use_default_provider { + let update = self + .session_manager + .update(session_id) + .provider_name(DEFAULT_PROVIDER_ID); + if has_default_overrides { + let provider_model_config = provider.get_model_config(); + update + .model_config(provider_model_config) + .apply() + .await + .map_err(|e| { + sacp::Error::internal_error().data(format!( + "Failed to persist default provider selection overrides: {}", + e + )) + })?; + } else { + update.clear_model_config().apply().await.map_err(|e| { + sacp::Error::internal_error().data(format!( + "Failed to persist default provider selection: {}", + e + )) + })?; + } + } + + let (_, config_options) = self + .build_config_update(&SessionId::new(session_id.to_string())) + .await?; + Ok(config_options) + } + async fn on_list_sessions(&self) -> Result { let sessions = self .session_manager @@ -1468,6 +1761,124 @@ impl GooseAcpAgent { warnings, }) } + + #[custom_method(UpdateProviderRequest)] + async fn on_update_provider( + &self, + req: UpdateProviderRequest, + ) -> Result { + let config_options = self + .update_provider( + &req.session_id, + &req.provider, + req.model.as_deref(), + req.context_limit, + req.request_params, + ) + .await?; + let config_options = config_options + .into_iter() + .map(|option| serde_json::to_value(&option)) + .collect::, _>>() + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + Ok(UpdateProviderResponse { config_options }) + } + + #[custom_method(ListProvidersRequest)] + async fn on_list_providers( + &self, + _req: ListProvidersRequest, + ) -> Result { + Ok(ListProvidersResponse { + providers: list_provider_entries(None).await, + }) + } + + #[custom_method(ReadConfigRequest)] + async fn on_read_config( + &self, + req: ReadConfigRequest, + ) -> Result { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + let response = match config.get_param::(&req.key) { + Ok(value) => ReadConfigResponse { value }, + Err(goose::config::ConfigError::NotFound(_)) => ReadConfigResponse { + value: serde_json::Value::Null, + }, + Err(e) => return Err(sacp::Error::internal_error().data(e.to_string())), + }; + Ok(response) + } + + #[custom_method(UpsertConfigRequest)] + async fn on_upsert_config( + &self, + req: UpsertConfigRequest, + ) -> Result { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + config + .set_param(&req.key, &req.value) + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + Ok(EmptyResponse {}) + } + + #[custom_method(RemoveConfigRequest)] + async fn on_remove_config( + &self, + req: RemoveConfigRequest, + ) -> Result { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + config + .delete(&req.key) + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + Ok(EmptyResponse {}) + } + + #[custom_method(CheckSecretRequest)] + async fn on_check_secret( + &self, + req: CheckSecretRequest, + ) -> Result { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + let exists = config.get_secret::(&req.key).is_ok(); + Ok(CheckSecretResponse { exists }) + } + + #[custom_method(UpsertSecretRequest)] + async fn on_upsert_secret( + &self, + req: UpsertSecretRequest, + ) -> Result { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + config + .set_secret(&req.key, &req.value) + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + Ok(EmptyResponse {}) + } + + #[custom_method(RemoveSecretRequest)] + async fn on_remove_secret( + &self, + req: RemoveSecretRequest, + ) -> Result { + let config = self.load_config().map_err(|e| { + sacp::Error::internal_error().data(format!("Failed to read config: {}", e)) + })?; + config + .delete_secret(&req.key) + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + Ok(EmptyResponse {}) + } } pub struct GooseAcpHandler { @@ -1545,6 +1956,12 @@ impl HandleDispatchFrom for GooseAcpHandler { .clone(); let session_id = req.session_id.clone(); match req.config_id.0.as_ref() { + "provider" => { + match agent.update_provider(&session_id.0, &value_id.0, None, None, None).await { + Ok(_) => {} + Err(e) => { responder.respond_with_error(e)?; return Ok(()); } + } + } "mode" => { match agent.on_set_mode(&session_id.0, &value_id.0).await { Ok(_) => {} @@ -2089,11 +2506,23 @@ print(\"hello, world\") #[test_case( build_mode_state(GooseMode::Auto).unwrap(), + "openai", + vec![ + SessionConfigSelectOption::new("anthropic", "anthropic"), + SessionConfigSelectOption::new("openai", "openai"), + ], SessionModelState::new( ModelId::new("gpt-4"), vec![ModelInfo::new(ModelId::new("gpt-4"), "gpt-4"), ModelInfo::new(ModelId::new("gpt-3.5"), "gpt-3.5")], ) => vec![ + SessionConfigOption::select( + "provider", "Provider", "openai", + vec![ + SessionConfigSelectOption::new("anthropic", "anthropic"), + SessionConfigSelectOption::new("openai", "openai"), + ], + ), SessionConfigOption::select( "mode", "Mode", "auto", vec![ @@ -2115,8 +2544,14 @@ print(\"hello, world\") )] #[test_case( build_mode_state(GooseMode::Approve).unwrap(), + "openai", + vec![SessionConfigSelectOption::new("openai", "openai")], SessionModelState::new(ModelId::new("only-model"), vec![ModelInfo::new(ModelId::new("only-model"), "only-model")]) => vec![ + SessionConfigOption::select( + "provider", "Provider", "openai", + vec![SessionConfigSelectOption::new("openai", "openai")], + ), SessionConfigOption::select( "mode", "Mode", "approve", vec![ @@ -2135,8 +2570,10 @@ print(\"hello, world\") )] fn test_build_config_options( mode_state: SessionModeState, + provider_name: &'static str, + provider_options: Vec, model_state: SessionModelState, ) -> Vec { - build_config_options(&mode_state, &model_state) + build_config_options(&mode_state, &model_state, provider_name, provider_options) } } diff --git a/crates/goose-acp/src/server_factory.rs b/crates/goose-acp/src/server_factory.rs index b2810d85..5c63d017 100644 --- a/crates/goose-acp/src/server_factory.rs +++ b/crates/goose-acp/src/server_factory.rs @@ -1,9 +1,8 @@ use anyhow::Result; -use goose::providers::provider_registry::ProviderConstructor; use std::sync::Arc; use tracing::info; -use crate::server::GooseAcpAgent; +use crate::server::{AcpProviderFactory, GooseAcpAgent}; pub struct AcpServerFactoryConfig { pub builtins: Vec, @@ -32,16 +31,12 @@ impl AcpServer { .unwrap_or(goose::config::GooseMode::Auto); let disable_session_naming = config.get_goose_disable_session_naming().unwrap_or(false); - let config_dir = self.config.config_dir.clone(); - let provider_factory: ProviderConstructor = Arc::new(move |model_config, extensions| { - let config_dir = config_dir.clone(); - Box::pin(async move { - let config_path = config_dir.join(goose::config::base::CONFIG_YAML_NAME); - let config = goose::config::Config::new(&config_path, "goose")?; - let provider_name = config.get_goose_provider()?; - goose::providers::create(&provider_name, model_config, extensions).await - }) - }); + let provider_factory: AcpProviderFactory = + Arc::new(move |provider_name, model_config, extensions| { + Box::pin(async move { + goose::providers::create(&provider_name, model_config, extensions).await + }) + }); let agent = GooseAcpAgent::new( provider_factory, diff --git a/crates/goose-acp/tests/common_tests/mod.rs b/crates/goose-acp/tests/common_tests/mod.rs index a3076b00..e965a926 100644 --- a/crates/goose-acp/tests/common_tests/mod.rs +++ b/crates/goose-acp/tests/common_tests/mod.rs @@ -11,7 +11,7 @@ use fixtures::{ use fs_err as fs; use goose::config::base::CONFIG_YAML_NAME; use goose::config::GooseMode; -use goose::providers::provider_registry::ProviderConstructor; +use goose_acp::server::AcpProviderFactory; use goose_test_support::{McpFixture, FAKE_CODE, TEST_IMAGE_B64, TEST_MODEL}; use sacp::schema::{ ListSessionsResponse, McpServer, McpServerHttp, ModelId, SessionInfo, SessionModeId, @@ -332,8 +332,8 @@ pub async fn run_fs_write_text_file_true() { } pub async fn run_initialize_doesnt_hit_provider() { - let provider_factory: ProviderConstructor = - Arc::new(|_, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) })); + let provider_factory: AcpProviderFactory = + Arc::new(|_, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) })); let openai = OpenAiFixture::new(vec![], C::expected_session_id()).await; let config = TestConnectionConfig { @@ -341,12 +341,7 @@ pub async fn run_initialize_doesnt_hit_provider() { ..Default::default() }; - let conn = C::new(config, openai).await; - assert!(!conn.auth_methods().is_empty()); - assert!(conn - .auth_methods() - .iter() - .any(|m| m.id().0.as_ref() == "goose-provider")); + let _conn = C::new(config, openai).await; } pub async fn run_load_mode() { @@ -740,6 +735,24 @@ pub async fn run_model_list() { assert_eq!(models.current_model_id, ModelId::new(TEST_MODEL)); } +#[allow(dead_code)] +pub async fn run_new_session_returns_initial_config() { + let expected_session_id = C::expected_session_id(); + let openai = OpenAiFixture::new(vec![], expected_session_id.clone()).await; + + let mut conn = C::new(TestConnectionConfig::default(), openai).await; + let SessionData { + session, + models, + modes, + } = conn.new_session().await.unwrap(); + expected_session_id.set(&session.session_id().0); + + assert!(modes.is_some()); + let models = models.expect("new_session should return models inline"); + assert!(!models.available_models.is_empty()); +} + pub async fn run_config_option_model_set() { run_model_set_impl::(SetModelVia::ConfigOption).await; } @@ -808,13 +821,16 @@ async fn run_model_set_impl(via: SetModelVia) { .unwrap(); assert_eq!(output.text, "2"); - // Both model paths emit ConfigOption (no CurrentModelUpdate in the schema). + // Some connections emit a ConfigOption update immediately on model change, + // while the stripped legacy provider path only updates local state before + // the next prompt. let prompt_notifs = session_b.notifications(); let mut all = set_model_notifs; all.extend(prompt_notifs); - assert_eq!( - all, - vec![Notification::ConfigOption, Notification::AgentMessage], + assert!( + all == vec![Notification::AgentMessage] + || all == vec![Notification::ConfigOption, Notification::AgentMessage], + "unexpected notifications after model change: {all:?}" ); // Prompt A: expects default TEST_MODEL (proves sessions are independent) diff --git a/crates/goose-acp/tests/custom_requests_test.rs b/crates/goose-acp/tests/custom_requests_test.rs index c54b96cc..29f983be 100644 --- a/crates/goose-acp/tests/custom_requests_test.rs +++ b/crates/goose-acp/tests/custom_requests_test.rs @@ -5,11 +5,66 @@ use common_tests::fixtures::server::AcpServerConnection; use common_tests::fixtures::{ run_test, send_custom, Connection, Session, SessionData, TestConnectionConfig, }; +use goose::model::ModelConfig; +use goose::providers::base::{MessageStream, Provider}; +use goose::providers::errors::ProviderError; +use goose_acp::server::AcpProviderFactory; use goose_test_support::EnforceSessionId; use std::sync::Arc; use common_tests::fixtures::OpenAiFixture; +struct MockProvider { + name: String, + model_config: ModelConfig, + recommended_models: Vec, +} + +#[async_trait::async_trait] +impl Provider for MockProvider { + fn get_name(&self) -> &str { + &self.name + } + + async fn stream( + &self, + _model_config: &ModelConfig, + _session_id: &str, + _system: &str, + _messages: &[goose::conversation::message::Message], + _tools: &[rmcp::model::Tool], + ) -> Result { + unimplemented!() + } + + fn get_model_config(&self) -> ModelConfig { + self.model_config.clone() + } + + async fn fetch_recommended_models(&self) -> Result, ProviderError> { + Ok(self.recommended_models.clone()) + } +} + +fn mock_provider_factory() -> AcpProviderFactory { + Arc::new(|provider_name, model_config, _extensions| { + Box::pin(async move { + let recommended_models = match provider_name.as_str() { + "anthropic" => vec![ + "claude-3-7-sonnet-latest".to_string(), + "claude-3-5-haiku-latest".to_string(), + ], + _ => vec!["gpt-4o".to_string(), "o4-mini".to_string()], + }; + Ok(Arc::new(MockProvider { + name: provider_name, + model_config, + recommended_models, + }) as Arc) + }) + }) +} + #[test] fn test_custom_session_get() { run_test(async { @@ -83,6 +138,209 @@ fn test_custom_get_extensions() { }); } +#[test] +fn test_custom_list_providers() { + run_test(async { + let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; + let conn = AcpServerConnection::new(TestConnectionConfig::default(), openai).await; + + let response = send_custom(conn.cx(), "_goose/providers/list", serde_json::json!({})) + .await + .expect("provider list should succeed"); + let providers = response + .get("providers") + .and_then(|value| value.as_array()) + .expect("missing providers array"); + + assert!( + providers.iter().any(|provider| { + provider.get("id") == Some(&serde_json::json!("goose")) + && provider.get("label") == Some(&serde_json::json!("Goose (Default)")) + }), + "expected Goose default provider sentinel" + ); + assert!( + providers + .iter() + .any(|provider| provider.get("id") == Some(&serde_json::json!("openai"))), + "expected at least one concrete provider from the goose registry" + ); + }); +} + +#[test] +fn test_custom_config_crud() { + run_test(async { + let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; + let conn = AcpServerConnection::new(TestConnectionConfig::default(), openai).await; + + send_custom( + conn.cx(), + "_goose/config/upsert", + serde_json::json!({ + "key": "GOOSE_PROVIDER", + "value": "anthropic", + }), + ) + .await + .expect("config upsert should succeed"); + + let response = send_custom( + conn.cx(), + "_goose/config/read", + serde_json::json!({ + "key": "GOOSE_PROVIDER", + }), + ) + .await + .expect("config read should succeed"); + assert_eq!(response.get("value"), Some(&serde_json::json!("anthropic"))); + + send_custom( + conn.cx(), + "_goose/config/remove", + serde_json::json!({ + "key": "GOOSE_PROVIDER", + }), + ) + .await + .expect("config remove should succeed"); + + let response = send_custom( + conn.cx(), + "_goose/config/read", + serde_json::json!({ + "key": "GOOSE_PROVIDER", + }), + ) + .await + .expect("config read after remove should succeed"); + assert_eq!(response.get("value"), Some(&serde_json::Value::Null)); + }); +} + +#[test] +fn test_provider_switching_updates_session_state() { + run_test(async { + let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; + let config = TestConnectionConfig { + provider_factory: Some(mock_provider_factory()), + current_model: "gpt-4o".to_string(), + ..Default::default() + }; + let mut conn = AcpServerConnection::new(config, openai).await; + + let SessionData { session, .. } = conn.new_session().await.unwrap(); + let session_id = session.session_id().0.clone(); + + conn.set_config_option(&session_id, "provider", "anthropic") + .await + .expect("provider config option should succeed"); + + let response = send_custom( + conn.cx(), + "session/get", + serde_json::json!({ + "sessionId": session_id, + }), + ) + .await + .expect("session/get should succeed"); + let session_value = response.get("session").expect("missing session"); + assert_eq!( + session_value.get("provider_name"), + Some(&serde_json::json!("anthropic")) + ); + assert_eq!( + session_value + .get("model_config") + .and_then(|value| value.get("model_name")), + Some(&serde_json::json!("current")) + ); + + let response = send_custom( + conn.cx(), + "_goose/session/provider/update", + serde_json::json!({ + "sessionId": session_id, + "provider": "openai", + "model": "o4-mini", + }), + ) + .await + .expect("provider update should succeed"); + let config_options = response + .get("configOptions") + .and_then(|value| value.as_array()) + .expect("missing config options"); + assert!( + !config_options.is_empty(), + "expected refreshed config options" + ); + + let response = send_custom( + conn.cx(), + "session/get", + serde_json::json!({ + "sessionId": session_id, + }), + ) + .await + .expect("session/get after provider update should succeed"); + let session_value = response.get("session").expect("missing session"); + assert_eq!( + session_value.get("provider_name"), + Some(&serde_json::json!("openai")) + ); + assert_eq!( + session_value + .get("model_config") + .and_then(|value| value.get("model_name")), + Some(&serde_json::json!("o4-mini")) + ); + + let response = send_custom( + conn.cx(), + "_goose/session/provider/update", + serde_json::json!({ + "sessionId": session_id, + "provider": "goose", + }), + ) + .await + .expect("provider reset to goose should succeed"); + let config_options = response + .get("configOptions") + .and_then(|value| value.as_array()) + .expect("missing config options after reset"); + assert!( + config_options + .iter() + .any(|option| option.get("id") == Some(&serde_json::json!("provider"))), + "missing provider config option after reset" + ); + + let response = send_custom( + conn.cx(), + "session/get", + serde_json::json!({ + "sessionId": session_id, + }), + ) + .await + .expect("session/get after provider reset should succeed"); + let session_value = response.get("session").expect("missing session"); + assert_eq!( + session_value.get("provider_name"), + Some(&serde_json::json!("goose")) + ); + assert_eq!( + session_value.get("model_config"), + Some(&serde_json::Value::Null) + ); + }); +} + #[test] fn test_custom_unknown_method() { run_test(async { diff --git a/crates/goose-acp/tests/fixtures/mod.rs b/crates/goose-acp/tests/fixtures/mod.rs index e7fb12de..4fc0c970 100644 --- a/crates/goose-acp/tests/fixtures/mod.rs +++ b/crates/goose-acp/tests/fixtures/mod.rs @@ -3,19 +3,18 @@ use async_trait::async_trait; use fs_err as fs; -pub use goose::acp::{map_permission_response, PermissionDecision, PermissionMapping}; +pub use goose::acp::{map_permission_response, PermissionDecision}; use goose::builtin_extension::register_builtin_extensions; use goose::config::paths::Paths; use goose::config::{GooseMode, PermissionManager}; use goose::providers::api_client::{ApiClient, AuthMethod as ApiAuthMethod}; use goose::providers::base::Provider; use goose::providers::openai::OpenAiProvider; -use goose::providers::provider_registry::ProviderConstructor; use goose::session_context::SESSION_ID_HEADER; -use goose_acp::server::{serve, GooseAcpAgent}; +use goose_acp::server::{serve, AcpProviderFactory, GooseAcpAgent}; use goose_test_support::{ExpectedSessionId, TEST_MODEL}; use sacp::schema::{ - AuthMethod, CreateTerminalResponse, KillTerminalResponse, ListSessionsResponse, McpServer, + CreateTerminalResponse, KillTerminalResponse, ListSessionsResponse, McpServer, ReadTextFileRequest, ReadTextFileResponse, ReleaseTerminalResponse, SessionModeState, SessionModelState, SessionUpdate, TerminalExitStatus, TerminalId, TerminalOutputResponse, ToolCallContent, ToolCallStatus, ToolKind, WaitForTerminalExitResponse, WriteTextFileRequest, @@ -155,7 +154,7 @@ pub async fn spawn_acp_server_in_process( builtins: &[String], data_root: &std::path::Path, goose_mode: GooseMode, - provider_factory: Option, + provider_factory: Option, current_model: &str, ) -> (DuplexTransport, JoinHandle<()>, Arc) { fs::create_dir_all(data_root).unwrap(); @@ -171,7 +170,7 @@ pub async fn spawn_acp_server_in_process( } let provider_factory = provider_factory.unwrap_or_else(|| { let base_url = openai_base_url.to_string(); - Arc::new(move |model_config, _extensions| { + Arc::new(move |_provider_name, model_config, _extensions| { let base_url = base_url.clone(); Box::pin(async move { let api_client = @@ -472,7 +471,7 @@ pub struct TestConnectionConfig { pub goose_mode: GooseMode, pub cwd: Option, pub data_root: PathBuf, - pub provider_factory: Option, + pub provider_factory: Option, pub read_text_file: Option, pub write_text_file: Option, pub terminal: Option>, @@ -524,7 +523,6 @@ pub trait Connection: Sized { config_id: &str, value: &str, ) -> anyhow::Result<()>; - fn auth_methods(&self) -> &[AuthMethod]; fn data_root(&self) -> std::path::PathBuf; fn reset_openai(&self); fn reset_permissions(&self); diff --git a/crates/goose-acp/tests/fixtures/provider.rs b/crates/goose-acp/tests/fixtures/provider.rs index d9a38d1e..6f5b658a 100644 --- a/crates/goose-acp/tests/fixtures/provider.rs +++ b/crates/goose-acp/tests/fixtures/provider.rs @@ -4,7 +4,7 @@ use super::{ }; use async_trait::async_trait; use futures::StreamExt; -use goose::acp::{AcpProvider, AcpProviderConfig, PermissionMapping}; +use goose::acp::{AcpProvider, AcpProviderConfig}; use goose::config::{GooseMode, PermissionManager}; use goose::conversation::message::{ActionRequiredData, Message, MessageContent}; use goose::model::ModelConfig; @@ -12,9 +12,12 @@ use goose::permission::permission_confirmation::PrincipalType; use goose::permission::{Permission, PermissionConfirmation}; use goose::providers::base::Provider; use goose_test_support::{ExpectedSessionId, IgnoreSessionId, TEST_MODEL}; -use sacp::schema::{AuthMethod, ListSessionsResponse, McpServer, SessionUpdate, ToolCallStatus}; +use sacp::schema::{ + ListSessionsResponse, McpServer, ModelId, ModelInfo, SessionModelState, SessionUpdate, + ToolCallStatus, +}; use sacp::{Channel, Client, ConnectTo, DynConnectTo}; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::str::FromStr; use std::sync::Arc; use strum::VariantNames; @@ -28,10 +31,10 @@ pub struct AcpProviderConnection { /// Option so close_session can trigger session/close via Drop. provider: Arc>>, permission_manager: Arc, - auth_methods: Vec, session_counter: usize, notification_sink: NotificationSink, session_models: SessionModels, + strip_config_options: bool, work_dir: std::path::PathBuf, data_root: std::path::PathBuf, _openai: OpenAiFixture, @@ -186,7 +189,6 @@ impl Connection for AcpProviderConnection { (mode, mode.to_string()) }) .collect(), - permission_mapping: PermissionMapping::default(), notification_callback: Some(Arc::new(move |n| { sink_clone.lock().unwrap().push(n.update.clone()); })), @@ -208,15 +210,13 @@ impl Connection for AcpProviderConnection { .await .unwrap(); - let auth_methods = provider.auth_methods().to_vec(); - Self { provider: Arc::new(Mutex::new(Some(provider))), permission_manager, - auth_methods, session_counter: 0, notification_sink, session_models, + strip_config_options: config.strip_config_options, work_dir: cwd_path, data_root, _openai: openai, @@ -226,18 +226,23 @@ impl Connection for AcpProviderConnection { } async fn new_session(&mut self) -> anyhow::Result> { - // Tests like run_model_set call new_session() multiple times on the same - // connection, so each needs a distinct key to avoid returning a cached session. self.session_counter += 1; let goose_id = format!("test-session-{}", self.session_counter); - let response = self - .provider - .lock() - .await - .as_ref() - .unwrap() - .ensure_session(Some(&goose_id)) - .await?; + + let models = if self.strip_config_options { + None + } else { + let provider = self.provider.lock().await; + let provider = provider.as_ref().unwrap(); + let available_models = provider.fetch_supported_models().await?; + Some(SessionModelState::new( + ModelId::new(provider.get_model_config().model_name.clone()), + available_models + .into_iter() + .map(|model_id| ModelInfo::new(ModelId::new(model_id.clone()), model_id)) + .collect(), + )) + }; let session = AcpProviderSession { provider: Arc::clone(&self.provider), @@ -246,10 +251,11 @@ impl Connection for AcpProviderConnection { session_models: self.session_models.clone(), work_dir: self.work_dir.clone(), }; + self.notification_sink.lock().unwrap().clear(); Ok(SessionData { session, - models: response.models, - modes: response.modes, + models, + modes: None, }) } @@ -264,13 +270,7 @@ impl Connection for AcpProviderConnection { } async fn list_sessions(&self) -> anyhow::Result { - self.provider - .lock() - .await - .as_ref() - .unwrap() - .list_sessions() - .await + Err(anyhow::anyhow!("not implemented for AcpProviderConnection")) } async fn close_session(&self, _session_id: &str) -> anyhow::Result<()> { @@ -279,89 +279,29 @@ impl Connection for AcpProviderConnection { Ok(()) } - async fn delete_session(&self, session_id: &str) -> anyhow::Result<()> { - self.provider - .lock() - .await - .as_ref() - .unwrap() - .delete_session(session_id) - .await + async fn delete_session(&self, _session_id: &str) -> anyhow::Result<()> { + Err(anyhow::anyhow!("not implemented for AcpProviderConnection")) } fn data_root(&self) -> std::path::PathBuf { self.data_root.clone() } - async fn set_mode(&self, session_id: &str, mode_id: &str) -> anyhow::Result<()> { - let mode = GooseMode::from_str(mode_id) - .map_err(|_| sacp::Error::invalid_params().data(format!("Invalid mode: {mode_id}")))?; - let guard = self.provider.lock().await; - let provider = guard.as_ref().unwrap(); - if !provider.has_session(session_id).await { - return Err( - sacp::Error::resource_not_found(Some(session_id.to_string())) - .data(format!("Session not found: {session_id}")) - .into(), - ); - } - provider - .update_mode(session_id, mode) - .await - .map_err(|e| anyhow::anyhow!("{e}")) + async fn set_mode(&self, _session_id: &str, _mode_id: &str) -> anyhow::Result<()> { + Err(anyhow::anyhow!("not implemented for AcpProviderConnection")) } - async fn set_model(&self, session_id: &str, model_id: &str) -> anyhow::Result<()> { - let config = ModelConfig::new(model_id).map_err(|e| anyhow::anyhow!("{e}"))?; - self.session_models - .lock() - .unwrap() - .insert(session_id.to_string(), config); - Ok(()) + async fn set_model(&self, _session_id: &str, _model_id: &str) -> anyhow::Result<()> { + Err(anyhow::anyhow!("not implemented for AcpProviderConnection")) } async fn set_config_option( &self, - session_id: &str, - config_id: &str, - value: &str, + _session_id: &str, + _config_id: &str, + _value: &str, ) -> anyhow::Result<()> { - // Check up front because the "model" branch doesn't go through the provider. - let guard = self.provider.lock().await; - let provider = guard.as_ref().unwrap(); - if !provider.has_session(session_id).await { - return Err( - sacp::Error::resource_not_found(Some(session_id.to_string())) - .data(format!("Session not found: {session_id}")) - .into(), - ); - } - match config_id { - "mode" => { - let mode = GooseMode::from_str(value).map_err(|_| { - sacp::Error::invalid_params().data(format!("Invalid mode: {value}")) - })?; - provider - .update_mode(session_id, mode) - .await - .map_err(|e| anyhow::anyhow!("{e}")) - } - "model" => { - let config = ModelConfig::new(value).map_err(|e| anyhow::anyhow!("{e}"))?; - self.session_models - .lock() - .unwrap() - .insert(session_id.to_string(), config); - Ok(()) - } - other => Err(sacp::Error::invalid_params() - .data(format!("Unsupported config option: {other}")) - .into()), - } - } - - fn auth_methods(&self) -> &[AuthMethod] { - &self.auth_methods + Err(anyhow::anyhow!("not implemented for AcpProviderConnection")) } fn reset_openai(&self) { @@ -425,6 +365,8 @@ fn strip_config_options(transport: DuplexTransport) -> Channel { }); tokio::spawn(async move { + let mut stripped_initial_config = HashSet::new(); + let goose_to_server = async { let mut from_goose = filter.rx; while let Some(msg) = from_goose.next().await { @@ -437,19 +379,67 @@ fn strip_config_options(transport: DuplexTransport) -> Channel { let server_to_goose = async { let mut from_server = server.rx; while let Some(msg) = from_server.next().await { - let msg = msg.map(|m| match m { - sacp::jsonrpcmsg::Message::Response(mut resp) => { - if let Some(ref mut result) = resp.result { - if let Some(obj) = result.as_object_mut() { - obj.remove("configOptions"); + let msg = match msg { + Ok(m) => match m { + sacp::jsonrpcmsg::Message::Response(mut resp) => { + if let Some(ref mut result) = resp.result { + if let Some(obj) = result.as_object_mut() { + obj.remove("configOptions"); + } + } + Ok(Some(sacp::jsonrpcmsg::Message::Response(resp))) + } + sacp::jsonrpcmsg::Message::Request(req) + if req.id.is_none() + && req.method == "session/update" + && req + .params + .as_ref() + .and_then(|params| match params { + sacp::jsonrpcmsg::Params::Object(obj) => Some(obj), + _ => None, + }) + .and_then(|obj| obj.get("update")) + .and_then(|update| update.get("sessionUpdate")) + .and_then(|session_update| session_update.as_str()) + == Some("config_option_update") => + { + let session_id = req + .params + .as_ref() + .and_then(|params| match params { + sacp::jsonrpcmsg::Params::Object(obj) => Some(obj), + _ => None, + }) + .and_then(|obj| obj.get("sessionId")) + .and_then(|session_id| session_id.as_str()) + .map(str::to_owned); + if let Some(session_id) = session_id { + if stripped_initial_config.insert(session_id) { + Ok(None) + } else { + Ok(Some(sacp::jsonrpcmsg::Message::Request(req))) + } + } else { + Ok(Some(sacp::jsonrpcmsg::Message::Request(req))) } } - sacp::jsonrpcmsg::Message::Response(resp) + other => Ok(Some(other)), + }, + Err(err) => Err(err), + }; + match msg { + Ok(Some(msg)) => { + if filter.tx.unbounded_send(Ok(msg)).is_err() { + break; + } + } + Ok(None) => continue, + Err(err) => { + if filter.tx.unbounded_send(Err(err)).is_err() { + break; + } } - other => other, - }); - if filter.tx.unbounded_send(msg).is_err() { - break; } } }; diff --git a/crates/goose-acp/tests/fixtures/server.rs b/crates/goose-acp/tests/fixtures/server.rs index 6a6ed244..20fb1c09 100644 --- a/crates/goose-acp/tests/fixtures/server.rs +++ b/crates/goose-acp/tests/fixtures/server.rs @@ -1,19 +1,20 @@ use super::{ - map_permission_response, spawn_acp_server_in_process, Connection, PermissionDecision, - PermissionMapping, Session, SessionData, TestConnectionConfig, TestOutput, + map_permission_response, spawn_acp_server_in_process, Connection, PermissionDecision, Session, + SessionData, TestConnectionConfig, TestOutput, }; use async_trait::async_trait; use goose::config::PermissionManager; use goose_test_support::{EnforceSessionId, ExpectedSessionId}; use sacp::schema::{ - AuthMethod, ClientCapabilities, CloseSessionRequest, ContentBlock, CreateTerminalRequest, + ClientCapabilities, CloseSessionRequest, ContentBlock, CreateTerminalRequest, FileSystemCapabilities, ImageContent, InitializeRequest, KillTerminalRequest, - ListSessionsRequest, ListSessionsResponse, LoadSessionRequest, McpServer, NewSessionRequest, - PromptRequest, ProtocolVersion, ReadTextFileRequest, ReleaseTerminalRequest, - RequestPermissionRequest, SessionConfigOptionValue, SessionId, SessionModeId, - SessionNotification, SessionUpdate, SetSessionConfigOptionRequest, SetSessionModeRequest, - SetSessionModelRequest, StopReason, TerminalOutputRequest, TextContent, ToolCallStatus, - WaitForTerminalExitRequest, WriteTextFileRequest, + ListSessionsRequest, ListSessionsResponse, LoadSessionRequest, McpServer, ModelId, ModelInfo, + NewSessionRequest, PromptRequest, ProtocolVersion, ReadTextFileRequest, ReleaseTerminalRequest, + RequestPermissionRequest, SessionConfigKind, SessionConfigOptionCategory, + SessionConfigOptionValue, SessionId, SessionModeId, SessionModelState, SessionNotification, + SessionUpdate, SetSessionConfigOptionRequest, SetSessionModeRequest, SetSessionModelRequest, + StopReason, TerminalOutputRequest, TextContent, ToolCallStatus, WaitForTerminalExitRequest, + WriteTextFileRequest, }; use sacp::{Agent, Client, ConnectionTo}; use std::sync::{Arc, Mutex}; @@ -30,7 +31,6 @@ pub struct AcpServerConnection { permission: Arc>, notify: Arc, permission_manager: Arc, - auth_methods: Vec, _openai: super::OpenAiFixture, _temp_dir: Option, } @@ -133,7 +133,7 @@ impl Connection for AcpServerConnection { fs_cap = fs_cap.write_text_file(true); } - let (cx, auth_methods) = { + let cx = { let updates_clone = updates.clone(); let notify_clone = notify.clone(); let permission_clone = permission.clone(); @@ -143,14 +143,10 @@ impl Connection for AcpServerConnection { let cx_holder: Arc>>> = Arc::new(Mutex::new(None)); let cx_holder_clone = cx_holder.clone(); - let auth_holder: Arc>> = Arc::new(Mutex::new(Vec::new())); - let auth_holder_clone = auth_holder.clone(); let (ready_tx, ready_rx) = tokio::sync::oneshot::channel(); tokio::spawn(async move { - let permission_mapping = PermissionMapping::default(); - let result = Client .builder() .on_receive_notification( @@ -170,9 +166,7 @@ impl Connection for AcpServerConnection { let permission = permission_clone.clone(); async move |req: RequestPermissionRequest, responder, _connection_cx| { let decision = *permission.lock().unwrap(); - let response = - map_permission_response(&permission_mapping, &req, decision); - responder.respond(response) + responder.respond(map_permission_response(&req, decision)) } }, sacp::on_receive_request!(), @@ -261,9 +255,8 @@ impl Connection for AcpServerConnection { ) .connect_with(transport, { let cx_holder = cx_holder_clone; - let auth_holder = auth_holder_clone; async move |cx: ConnectionTo| { - let resp = cx + let _resp = cx .send_request( InitializeRequest::new(ProtocolVersion::LATEST) .client_capabilities( @@ -276,7 +269,6 @@ impl Connection for AcpServerConnection { .await .unwrap(); - *auth_holder.lock().unwrap() = resp.auth_methods; *cx_holder.lock().unwrap() = Some(cx.clone()); let _ = ready_tx.send(()); @@ -292,8 +284,7 @@ impl Connection for AcpServerConnection { ready_rx.await.unwrap(); let cx = cx_holder.lock().unwrap().take().unwrap(); - let auth = std::mem::take(&mut *auth_holder.lock().unwrap()); - (cx, auth) + cx }; Self { @@ -305,7 +296,6 @@ impl Connection for AcpServerConnection { permission, notify, permission_manager, - auth_methods, _openai: openai, _temp_dir: temp_dir, } @@ -330,9 +320,13 @@ impl Connection for AcpServerConnection { notify: self.notify.clone(), _work_dir: work_dir, }; + let models = response.models.or_else(|| { + extract_model_state_from_config_options(response.config_options.as_deref()) + }); + self.updates.lock().unwrap().clear(); Ok(SessionData { session, - models: response.models, + models, modes: response.modes, }) } @@ -438,10 +432,6 @@ impl Connection for AcpServerConnection { .map_err(|e| e.into()) } - fn auth_methods(&self) -> &[AuthMethod] { - &self.auth_methods - } - fn data_root(&self) -> std::path::PathBuf { self.data_root.clone() } @@ -504,6 +494,46 @@ impl Session for AcpServerSession { } } +fn extract_model_state_from_config_options( + config_options: Option<&[sacp::schema::SessionConfigOption]>, +) -> Option { + let option = config_options? + .iter() + .find(|option| option.category.as_ref() == Some(&SessionConfigOptionCategory::Model))?; + let SessionConfigKind::Select(select) = &option.kind else { + return None; + }; + + let available_models = match &select.options { + sacp::schema::SessionConfigSelectOptions::Ungrouped(options) => options + .iter() + .map(|option| { + ModelInfo::new( + ModelId::new(option.value.0.to_string()), + option.name.clone(), + ) + }) + .collect(), + sacp::schema::SessionConfigSelectOptions::Grouped(groups) => groups + .iter() + .flat_map(|group| { + group.options.iter().map(|option| { + ModelInfo::new( + ModelId::new(option.value.0.to_string()), + option.name.clone(), + ) + }) + }) + .collect(), + _ => Vec::new(), + }; + + Some(SessionModelState::new( + ModelId::new(select.current_value.0.to_string()), + available_models, + )) +} + fn collect_agent_text(updates: &Arc>>) -> String { let guard = updates.lock().unwrap(); let mut text = String::new(); diff --git a/crates/goose-acp/tests/provider_test.rs b/crates/goose-acp/tests/provider_test.rs index a4fed959..9a1b1baf 100644 --- a/crates/goose-acp/tests/provider_test.rs +++ b/crates/goose-acp/tests/provider_test.rs @@ -1,47 +1,28 @@ #![recursion_limit = "256"] +#[allow(dead_code)] mod common_tests; use common_tests::fixtures::provider::AcpProviderConnection; use common_tests::fixtures::run_test; use common_tests::{ - run_close_session, run_config_mcp, run_config_option_mode_set, run_config_option_model_set, - run_delete_session, run_fs_read_text_file_true, run_fs_write_text_file_false, - run_fs_write_text_file_true, run_initialize_doesnt_hit_provider, run_list_sessions, - run_load_mode, run_load_model, run_load_session_error, run_load_session_mcp, run_mode_set, - run_model_list, run_model_set, run_model_set_error_session_not_found, - run_permission_persistence, run_prompt_basic, run_prompt_codemode, run_prompt_error, - run_prompt_image, run_prompt_image_attachment, run_prompt_mcp, run_prompt_model_mismatch, - run_prompt_skill, run_shell_terminal_false, run_shell_terminal_true, + run_close_session, run_config_mcp, run_delete_session, run_fs_read_text_file_true, + run_fs_write_text_file_false, run_fs_write_text_file_true, run_load_mode, run_load_model, + run_load_session_error, run_load_session_mcp, run_model_list, run_permission_persistence, + run_prompt_basic, run_prompt_codemode, run_prompt_error, run_prompt_image, + run_prompt_image_attachment, run_prompt_mcp, run_prompt_model_mismatch, run_prompt_skill, + run_shell_terminal_false, run_shell_terminal_true, }; -tests_config_option_set_error!(AcpProviderConnection); -tests_mode_set_error!(AcpProviderConnection); - #[test] fn test_config_mcp() { run_test(async { run_config_mcp::().await }); } -#[test] -fn test_config_option_mode_set() { - run_test(async { run_config_option_mode_set::().await }); -} - -#[test] -fn test_list_sessions() { - run_test(async { run_list_sessions::().await }); -} - #[test] fn test_close_session() { run_test(async { run_close_session::().await }); } -#[test] -fn test_config_option_model_set() { - run_test(async { run_config_option_model_set::().await }); -} - #[test] #[ignore = "delete is a server-side custom method not routed through the provider"] fn test_delete_session() { @@ -65,11 +46,6 @@ fn test_fs_write_text_file_true() { run_test(async { run_fs_write_text_file_true::().await }); } -#[test] -fn test_initialize_doesnt_hit_provider() { - run_test(async { run_initialize_doesnt_hit_provider::().await }); -} - #[test] #[ignore = "TODO: implement load_session in ACP provider"] fn test_load_mode() { @@ -94,27 +70,11 @@ fn test_load_session_mcp() { run_test(async { run_load_session_mcp::().await }); } -#[test] -fn test_mode_set() { - run_test(async { run_mode_set::().await }); -} - #[test] fn test_model_list() { run_test(async { run_model_list::().await }); } -#[test] -fn test_model_set() { - run_test(async { run_model_set::().await }); -} - -#[test] -#[ignore = "ensure_session lazy-creates sessions so deleted ones reappear"] -fn test_model_set_error_session_not_found() { - run_test(async { run_model_set_error_session_not_found::().await }); -} - #[test] fn test_permission_persistence() { run_test(async { run_permission_persistence::().await }); diff --git a/crates/goose-acp/tests/server_test.rs b/crates/goose-acp/tests/server_test.rs index e2a278b5..a57aa40e 100644 --- a/crates/goose-acp/tests/server_test.rs +++ b/crates/goose-acp/tests/server_test.rs @@ -1,3 +1,4 @@ +#[allow(dead_code)] mod common_tests; use common_tests::fixtures::run_test; use common_tests::fixtures::server::AcpServerConnection; @@ -7,9 +8,10 @@ use common_tests::{ run_fs_write_text_file_true, run_initialize_doesnt_hit_provider, run_list_sessions, run_load_mode, run_load_model, run_load_session_error, run_load_session_mcp, run_mode_set, run_model_list, run_model_set, run_model_set_error_session_not_found, - run_permission_persistence, run_prompt_basic, run_prompt_codemode, run_prompt_error, - run_prompt_image, run_prompt_image_attachment, run_prompt_mcp, run_prompt_model_mismatch, - run_prompt_skill, run_shell_terminal_false, run_shell_terminal_true, + run_new_session_returns_initial_config, run_permission_persistence, run_prompt_basic, + run_prompt_codemode, run_prompt_error, run_prompt_image, run_prompt_image_attachment, + run_prompt_mcp, run_prompt_model_mismatch, run_prompt_skill, run_shell_terminal_false, + run_shell_terminal_true, }; tests_config_option_set_error!(AcpServerConnection); @@ -95,6 +97,11 @@ fn test_model_list() { run_test(async { run_model_list::().await }); } +#[test] +fn test_new_session_returns_initial_config() { + run_test(async { run_new_session_returns_initial_config::().await }); +} + #[test] fn test_model_set() { run_test(async { run_model_set::().await }); diff --git a/crates/goose-sdk/src/custom_requests.rs b/crates/goose-sdk/src/custom_requests.rs index 1f199aea..afa61149 100644 --- a/crates/goose-sdk/src/custom_requests.rs +++ b/crates/goose-sdk/src/custom_requests.rs @@ -1,6 +1,7 @@ use sacp::{JsonRpcRequest, JsonRpcResponse}; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; /// Schema descriptor for a single custom method, produced by the /// `#[custom_methods]` macro's generated `custom_method_schemas()` function. @@ -150,6 +151,109 @@ pub struct GetExtensionsResponse { pub warnings: Vec, } +/// Atomically update the provider for a live session. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/session/provider/update", response = UpdateProviderResponse)] +#[serde(rename_all = "camelCase")] +pub struct UpdateProviderRequest { + pub session_id: String, + pub provider: String, + pub model: Option, + pub context_limit: Option, + pub request_params: Option>, +} + +/// Provider update response. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct UpdateProviderResponse { + /// Refreshed session config options after the provider/model change. + pub config_options: Vec, +} + +/// Read a single non-secret config value. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/config/read", response = ReadConfigResponse)] +#[serde(rename_all = "camelCase")] +pub struct ReadConfigRequest { + pub key: String, +} + +/// Config read response. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct ReadConfigResponse { + #[serde(default)] + pub value: serde_json::Value, +} + +/// Upsert a single non-secret config value. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/config/upsert", response = EmptyResponse)] +#[serde(rename_all = "camelCase")] +pub struct UpsertConfigRequest { + pub key: String, + pub value: serde_json::Value, +} + +/// Remove a single non-secret config value. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/config/remove", response = EmptyResponse)] +#[serde(rename_all = "camelCase")] +pub struct RemoveConfigRequest { + pub key: String, +} + +/// Check whether a secret exists. Never returns the actual value. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/secret/check", response = CheckSecretResponse)] +#[serde(rename_all = "camelCase")] +pub struct CheckSecretRequest { + pub key: String, +} + +/// Secret check response. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct CheckSecretResponse { + pub exists: bool, +} + +/// Set a secret value (write-only). +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/secret/upsert", response = EmptyResponse)] +#[serde(rename_all = "camelCase")] +pub struct UpsertSecretRequest { + pub key: String, + pub value: serde_json::Value, +} + +/// Remove a secret. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/secret/remove", response = EmptyResponse)] +#[serde(rename_all = "camelCase")] +pub struct RemoveSecretRequest { + pub key: String, +} + +/// List providers available through goose, including the config-default sentinel. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/providers/list", response = ListProvidersResponse)] +pub struct ListProvidersRequest {} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct ProviderListEntry { + pub id: String, + pub label: String, +} + +/// Provider list response. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +pub struct ListProvidersResponse { + pub providers: Vec, +} + /// Empty success response for operations that return no data. #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] pub struct EmptyResponse {} diff --git a/crates/goose/src/acp/common.rs b/crates/goose/src/acp/common.rs index 759687b2..803dd718 100644 --- a/crates/goose/src/acp/common.rs +++ b/crates/goose/src/acp/common.rs @@ -3,27 +3,10 @@ use std::str::FromStr; use crate::permission::Permission; use sacp::schema::{ PermissionOption, PermissionOptionKind, RequestPermissionOutcome, RequestPermissionRequest, - RequestPermissionResponse, SelectedPermissionOutcome, ToolCallStatus, + RequestPermissionResponse, SelectedPermissionOutcome, }; use strum::{Display, EnumString}; -#[derive(Clone, Debug)] -pub struct PermissionMapping { - pub allow_option_id: Option, - pub reject_option_id: Option, - pub rejected_tool_status: ToolCallStatus, -} - -impl Default for PermissionMapping { - fn default() -> Self { - Self { - allow_option_id: None, - reject_option_id: None, - rejected_tool_status: ToolCallStatus::Failed, - } - } -} - #[derive(Clone, Copy, Debug, Eq, PartialEq, Display, EnumString)] #[strum(serialize_all = "snake_case")] pub enum PermissionDecision { @@ -81,60 +64,30 @@ impl From<&RequestPermissionOutcome> for PermissionDecision { } } +/// Map a permission decision to a response by matching the option kind from the +/// request. Each decision tries its preferred kind first, then falls back to the +/// closest alternative (e.g. AllowAlways falls back to AllowOnce). pub fn map_permission_response( - mapping: &PermissionMapping, request: &RequestPermissionRequest, decision: PermissionDecision, ) -> RequestPermissionResponse { let selected_id = match decision { - PermissionDecision::AllowAlways => select_option_id( - &request.options, - &mapping.allow_option_id, - PermissionOptionKind::AllowAlways, - ) - .or_else(|| { - select_option_id( - &request.options, - &mapping.allow_option_id, - PermissionOptionKind::AllowOnce, - ) - }), - PermissionDecision::AllowOnce => select_option_id( - &request.options, - &mapping.allow_option_id, - PermissionOptionKind::AllowOnce, - ) - .or_else(|| { - select_option_id( - &request.options, - &mapping.allow_option_id, - PermissionOptionKind::AllowAlways, - ) - }), - PermissionDecision::RejectAlways => select_option_id( - &request.options, - &mapping.reject_option_id, - PermissionOptionKind::RejectAlways, - ) - .or_else(|| { - select_option_id( - &request.options, - &mapping.reject_option_id, - PermissionOptionKind::RejectOnce, - ) - }), - PermissionDecision::RejectOnce => select_option_id( - &request.options, - &mapping.reject_option_id, - PermissionOptionKind::RejectOnce, - ) - .or_else(|| { - select_option_id( - &request.options, - &mapping.reject_option_id, - PermissionOptionKind::RejectAlways, - ) - }), + PermissionDecision::AllowAlways => { + find_option(&request.options, PermissionOptionKind::AllowAlways) + .or_else(|| find_option(&request.options, PermissionOptionKind::AllowOnce)) + } + PermissionDecision::AllowOnce => { + find_option(&request.options, PermissionOptionKind::AllowOnce) + .or_else(|| find_option(&request.options, PermissionOptionKind::AllowAlways)) + } + PermissionDecision::RejectAlways => { + find_option(&request.options, PermissionOptionKind::RejectAlways) + .or_else(|| find_option(&request.options, PermissionOptionKind::RejectOnce)) + } + PermissionDecision::RejectOnce => { + find_option(&request.options, PermissionOptionKind::RejectOnce) + .or_else(|| find_option(&request.options, PermissionOptionKind::RejectAlways)) + } PermissionDecision::Cancel => None, }; @@ -147,18 +100,7 @@ pub fn map_permission_response( } } -fn select_option_id( - options: &[PermissionOption], - preferred_id: &Option, - kind: PermissionOptionKind, -) -> Option { - if let Some(preferred_id) = preferred_id { - let preferred = sacp::schema::PermissionOptionId::new(preferred_id.clone()); - if options.iter().any(|opt| opt.option_id == preferred) { - return Some(preferred_id.clone()); - } - } - +fn find_option(options: &[PermissionOption], kind: PermissionOptionKind) -> Option { options .iter() .find(|opt| opt.kind == kind) @@ -186,88 +128,34 @@ mod tests { } #[test_case( - Some("allow"), - None, - PermissionDecision::AllowOnce, - "allow", - true; - "allow_uses_preferred_id" - )] - #[test_case( - None, - None, PermissionDecision::AllowAlways, - "allow_always", - false; - "allow_always_prefers_kind" + "allow_always"; + "allow_always_matches_kind" )] #[test_case( - Some("missing"), - None, PermissionDecision::AllowOnce, - "allow_once", - false; - "allow_falls_back_to_kind" + "allow_once"; + "allow_once_matches_kind" )] #[test_case( - None, - Some("reject"), PermissionDecision::RejectOnce, - "reject", - true; - "reject_uses_preferred_id" + "reject_once"; + "reject_once_matches_kind" )] #[test_case( - None, - Some("missing"), - PermissionDecision::RejectOnce, - "reject_once", - false; - "reject_falls_back_to_kind" + PermissionDecision::RejectAlways, + "reject_always"; + "reject_always_matches_kind" )] - fn test_permission_mapping( - allow_option_id: Option<&str>, - reject_option_id: Option<&str>, - decision: PermissionDecision, - expected_id: &str, - include_preferred: bool, - ) { - let mut options = vec![ + fn test_permission_response(decision: PermissionDecision, expected_id: &str) { + let options = vec![ option("allow_once", PermissionOptionKind::AllowOnce), option("allow_always", PermissionOptionKind::AllowAlways), option("reject_once", PermissionOptionKind::RejectOnce), - option("reject", PermissionOptionKind::RejectAlways), + option("reject_always", PermissionOptionKind::RejectAlways), ]; - - if include_preferred { - if let Some(preferred_allow) = allow_option_id { - if !options - .iter() - .any(|opt| opt.option_id.0.as_ref() == preferred_allow) - { - options.push(option(preferred_allow, PermissionOptionKind::AllowOnce)); - } - } - - if let Some(preferred_reject) = reject_option_id { - if !options - .iter() - .any(|opt| opt.option_id.0.as_ref() == preferred_reject) - { - options.push(option(preferred_reject, PermissionOptionKind::RejectOnce)); - } - } - } - let request = make_request(options); - - let mapping = PermissionMapping { - allow_option_id: allow_option_id.map(|s| s.to_string()), - reject_option_id: reject_option_id.map(|s| s.to_string()), - rejected_tool_status: ToolCallStatus::Failed, - }; - - let response = map_permission_response(&mapping, &request, decision); + let response = map_permission_response(&request, decision); match response.outcome { RequestPermissionOutcome::Selected(selected) => { assert_eq!(selected.option_id.0.as_ref(), expected_id); @@ -276,10 +164,22 @@ mod tests { } } + #[test] + fn test_allow_always_falls_back_to_allow_once() { + let request = make_request(vec![option("allow_once", PermissionOptionKind::AllowOnce)]); + let response = map_permission_response(&request, PermissionDecision::AllowAlways); + match response.outcome { + RequestPermissionOutcome::Selected(selected) => { + assert_eq!(selected.option_id.0.as_ref(), "allow_once"); + } + _ => panic!("expected selected outcome"), + } + } + #[test_case(PermissionDecision::Cancel; "cancelled")] fn test_permission_cancelled(decision: PermissionDecision) { let request = make_request(vec![option("allow_once", PermissionOptionKind::AllowOnce)]); - let response = map_permission_response(&PermissionMapping::default(), &request, decision); + let response = map_permission_response(&request, decision); assert!(matches!( response.outcome, RequestPermissionOutcome::Cancelled diff --git a/crates/goose/src/acp/mod.rs b/crates/goose/src/acp/mod.rs index 4692125e..6e995948 100644 --- a/crates/goose/src/acp/mod.rs +++ b/crates/goose/src/acp/mod.rs @@ -1,7 +1,7 @@ mod common; mod provider; -pub use common::{map_permission_response, PermissionDecision, PermissionMapping}; +pub use common::{map_permission_response, PermissionDecision}; pub use provider::{ extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, }; diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index 6f9d0f3b..575fe631 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -4,14 +4,14 @@ use async_stream::try_stream; use futures::future::BoxFuture; use rmcp::model::{Role, Tool}; use sacp::schema::{ - AuthMethod, CloseSessionRequest, ContentBlock, ContentChunk, EnvVariable, HttpHeader, - ImageContent, InitializeRequest, InitializeResponse, ListSessionsRequest, ListSessionsResponse, - McpCapabilities, McpServer, McpServerHttp, McpServerStdio, NewSessionRequest, - NewSessionResponse, PromptRequest, PromptResponse, ProtocolVersion, RequestPermissionOutcome, - RequestPermissionRequest, RequestPermissionResponse, SessionConfigKind, - SessionConfigOptionCategory, SessionConfigSelectOptions, SessionId, SessionNotification, - SessionUpdate, SetSessionConfigOptionRequest, SetSessionModeRequest, SetSessionModeResponse, - SetSessionModelRequest, StopReason, TextContent, ToolCallContent, + ClientCapabilities, CloseSessionRequest, ContentBlock, ContentChunk, EnvVariable, HttpHeader, + ImageContent, InitializeRequest, InitializeResponse, McpCapabilities, McpServer, McpServerHttp, + McpServerStdio, NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, + ProtocolVersion, RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, + SessionConfigKind, SessionConfigOption, SessionConfigOptionCategory, + SessionConfigSelectOptions, SessionId, SessionNotification, SessionUpdate, + SetSessionConfigOptionRequest, SetSessionModeRequest, SetSessionModeResponse, StopReason, + TextContent, ToolCallContent, }; use sacp::{Agent, Client, ConnectionTo}; use std::collections::{HashMap, HashSet}; @@ -21,10 +21,10 @@ use std::process::Stdio; use std::sync::{Arc, Mutex}; use std::thread::JoinHandle; use tokio::process::{Child, Command}; -use tokio::sync::{mpsc, oneshot, Mutex as TokioMutex, OnceCell}; -use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; +use tokio::sync::{mpsc, oneshot, Mutex as TokioMutex}; +use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; -use crate::acp::{map_permission_response, PermissionDecision, PermissionMapping}; +use crate::acp::{map_permission_response, PermissionDecision}; use crate::config::{ExtensionConfig, GooseMode}; use crate::conversation::message::{Message, MessageContent}; use crate::model::ModelConfig; @@ -34,7 +34,7 @@ use crate::providers::base::{MessageStream, PermissionRouting, Provider}; use crate::providers::errors::ProviderError; use crate::subprocess::configure_subprocess; -/// Sentinel: resolved to SessionModelState.current_model_id at connect time. +/// Sentinel: resolved to the actual model name during connect(). pub const ACP_CURRENT_MODEL: &str = "current"; pub struct AcpProviderConfig { @@ -46,7 +46,6 @@ pub struct AcpProviderConfig { pub mcp_servers: Vec, pub session_mode_id: Option, pub mode_mapping: HashMap, - pub permission_mapping: PermissionMapping, pub notification_callback: Option>, } @@ -54,19 +53,11 @@ enum ClientRequest { NewSession { response_tx: oneshot::Sender>, }, - ListSessions { - response_tx: oneshot::Sender>, - }, SetMode { session_id: SessionId, mode_id: String, response_tx: oneshot::Sender>, }, - SetModel { - session_id: SessionId, - model_id: String, - response_tx: oneshot::Sender>, - }, SetConfigOption { session_id: SessionId, config_id: String, @@ -78,16 +69,6 @@ enum ClientRequest { content: Vec, response_tx: mpsc::Sender, }, - CloseSession { - session_id: SessionId, - response_tx: oneshot::Sender>, - }, - // For ACP methods not yet in agent-client-protocol-schema (e.g. session/delete) - Untyped { - method: String, - params: serde_json::Value, - response_tx: oneshot::Sender>, - }, } // tokio I/O handles can't move between runtimes, so the child process must be @@ -119,24 +100,26 @@ enum AcpUpdate { Error(String), } +/// The single ACP session backing this provider instance. +#[derive(Clone)] +struct AcpSession { + id: SessionId, + response: NewSessionResponse, +} + pub struct AcpProvider { name: String, model: ModelConfig, goose_mode: Arc>, - tx: Option>, - loop_thread: Option>, mode_mapping: HashMap, - permission_mapping: PermissionMapping, - rejected_tool_calls: Arc>>, + + session: AcpSession, + pending_confirmations: Arc>>>, - goose_to_acp_id: Arc>>, - acp_to_goose_id: Arc>>, - /// Per-session model tracking for detecting model changes in stream(). - session_model: Arc>>, - auth_methods: Vec, - supports_close: bool, - init_session: OnceCell, + + tx: Option>, + loop_thread: Option>, } impl std::fmt::Debug for AcpProvider { @@ -148,8 +131,6 @@ impl std::fmt::Debug for AcpProvider { } } -// Dedicated runtime on an OS thread so session/close completes even during -// main runtime shutdown. See reqwest InnerClientHandle. fn spawn_client_loop(fut: impl Future + Send + 'static) -> JoinHandle<()> { std::thread::spawn(move || { let rt = tokio::runtime::Builder::new_current_thread() @@ -211,117 +192,63 @@ impl AcpProvider { let (tx, rx) = mpsc::channel(32); let (init_tx, init_rx) = oneshot::channel(); let mode_mapping = config.mode_mapping.clone(); - let permission_mapping = config.permission_mapping.clone(); - let rejected_tool_calls = Arc::new(TokioMutex::new(HashSet::new())); - let goose_mode = Arc::new(Mutex::new(goose_mode)); - let client_loop = AcpClientLoop::new(config, goose_mode.clone()); + let goose_mode_shared = Arc::new(Mutex::new(goose_mode)); + let client_loop = AcpClientLoop::new(config, goose_mode_shared.clone()); let loop_thread = spawn_client_loop(run(client_loop, rx, init_tx)); - let init_response = init_rx + let _init_response = init_rx .await .context("ACP client initialization cancelled")??; - let supports_close = init_response - .agent_capabilities - .session_capabilities - .close - .is_some(); - let mut provider = Self::new_with_runtime( - name, - model, - goose_mode, - tx, - loop_thread, - mode_mapping, - permission_mapping, - rejected_tool_calls, - init_response.auth_methods, - supports_close, - ); - if provider.model.model_name == ACP_CURRENT_MODEL { - let response = provider.get_init_session().await?; - let (current_model, _) = resolve_model_info(&provider.name, response)?; - tracing::info!(from = ACP_CURRENT_MODEL, to = %current_model, "resolved ACP model"); - provider.model.model_name = current_model; - } - Ok(provider) - } + // Create the ACP session eagerly during connect. + let (session_tx, session_rx) = oneshot::channel(); + tx.send(ClientRequest::NewSession { + response_tx: session_tx, + }) + .await + .context("ACP client is unavailable")?; + let response = session_rx + .await + .context("ACP session creation cancelled")??; - #[allow(clippy::too_many_arguments)] - fn new_with_runtime( - name: String, - model: ModelConfig, - goose_mode: Arc>, - tx: mpsc::Sender, - loop_thread: JoinHandle<()>, - mode_mapping: HashMap, - permission_mapping: PermissionMapping, - rejected_tool_calls: Arc>>, - auth_methods: Vec, - supports_close: bool, - ) -> Self { - Self { + // Resolve model from the session response. + let resolved_model = if model.model_name == ACP_CURRENT_MODEL { + if let Ok((resolved, _)) = resolve_model_info(&name, &response) { + tracing::info!(from = ACP_CURRENT_MODEL, to = %resolved, "resolved ACP model"); + ModelConfig { + model_name: resolved, + ..model + } + } else { + model + } + } else { + model + }; + + let session = AcpSession { + id: response.session_id.clone(), + response, + }; + + Ok(Self { name, - model, - goose_mode, + model: resolved_model, + goose_mode: goose_mode_shared, + mode_mapping, + session, + pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), tx: Some(tx), loop_thread: Some(loop_thread), - mode_mapping, - permission_mapping, - rejected_tool_calls, - pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), - goose_to_acp_id: Arc::new(TokioMutex::new(HashMap::new())), - acp_to_goose_id: Arc::new(TokioMutex::new(HashMap::new())), - session_model: Arc::new(TokioMutex::new(HashMap::new())), - auth_methods, - supports_close, - init_session: OnceCell::new(), - } + }) } - pub fn auth_methods(&self) -> &[AuthMethod] { - &self.auth_methods + fn acp_session_id(&self) -> SessionId { + self.session.id.clone() } - pub async fn new_session(&self) -> Result { - let (response_tx, response_rx) = oneshot::channel(); - self.tx - .as_ref() - .unwrap() - .send(ClientRequest::NewSession { response_tx }) - .await - .context("ACP client is unavailable")?; - response_rx - .await - .context(format!("ACP {} cancelled", AGENT_METHOD_NAMES.session_new))? - } - - pub async fn list_sessions(&self) -> Result { - let (response_tx, response_rx) = oneshot::channel(); - self.tx - .as_ref() - .unwrap() - .send(ClientRequest::ListSessions { response_tx }) - .await - .context("ACP client is unavailable")?; - let raw = response_rx.await.context("ACP request cancelled")??; - let acp_to_goose = self.acp_to_goose_id.lock().await; - Ok(map_sessions_to_goose_ids(raw, &acp_to_goose)) - } - - async fn resolve_acp_session_id(&self, goose_id: &str) -> Result { - let map = self.goose_to_acp_id.lock().await; - map.get(goose_id) - .map(|r| r.session_id.clone()) - .ok_or_else(|| { - sacp::Error::resource_not_found(Some(goose_id.to_string())) - .data(format!("Session not found: {goose_id}")) - .into() - }) - } - - pub(crate) async fn send_set_mode(&self, goose_id: &str, mode_id: String) -> Result<()> { - let session_id = self.resolve_acp_session_id(goose_id).await?; + pub(crate) async fn send_set_mode(&self, _goose_id: &str, mode_id: String) -> Result<()> { + let session_id = self.acp_session_id(); let (response_tx, response_rx) = oneshot::channel(); self.tx .as_ref() @@ -336,29 +263,13 @@ impl AcpProvider { response_rx.await.context("ACP request cancelled")? } - pub(crate) async fn send_set_model(&self, goose_id: &str, model_id: String) -> Result<()> { - let session_id = self.resolve_acp_session_id(goose_id).await?; - let (response_tx, response_rx) = oneshot::channel(); - self.tx - .as_ref() - .unwrap() - .send(ClientRequest::SetModel { - session_id, - model_id, - response_tx, - }) - .await - .context("ACP client is unavailable")?; - response_rx.await.context("ACP request cancelled")? - } - pub(crate) async fn send_set_config_option( &self, - goose_id: &str, + _goose_id: &str, config_id: String, value: String, ) -> Result<()> { - let session_id = self.resolve_acp_session_id(goose_id).await?; + let session_id = self.acp_session_id(); let (response_tx, response_rx) = oneshot::channel(); self.tx .as_ref() @@ -374,110 +285,6 @@ impl AcpProvider { response_rx.await.context("ACP request cancelled")? } - // Only used by tests; session/delete has no typed request in agent-client-protocol-schema yet. - #[doc(hidden)] - pub async fn delete_session(&self, goose_id: &str) -> Result<()> { - let session_id = self.resolve_acp_session_id(goose_id).await?; - self.send_untyped( - "session/delete", - serde_json::json!({ "sessionId": session_id.0 }), - ) - .await?; - - // Clean up cached mappings so ensure_session doesn't return a stale entry. - self.goose_to_acp_id.lock().await.remove(goose_id); - self.acp_to_goose_id - .lock() - .await - .remove(session_id.0.as_ref()); - self.session_model.lock().await.remove(goose_id); - Ok(()) - } - - pub async fn send_untyped( - &self, - method: &str, - params: serde_json::Value, - ) -> Result { - let (response_tx, response_rx) = oneshot::channel(); - self.tx - .as_ref() - .unwrap() - .send(ClientRequest::Untyped { - method: method.to_string(), - params, - response_tx, - }) - .await - .context("ACP client is unavailable")?; - response_rx.await.context("ACP request cancelled")? - } - - pub async fn has_session(&self, goose_id: &str) -> bool { - self.goose_to_acp_id.lock().await.contains_key(goose_id) - } - - // If false, callers fall back to legacy set_mode/set_model. - async fn session_has_config_option( - &self, - goose_id: &str, - category: SessionConfigOptionCategory, - ) -> bool { - let map = self.goose_to_acp_id.lock().await; - map.get(goose_id) - .and_then(|r| r.config_options.as_ref()) - .is_some_and(|opts| opts.iter().any(|o| o.category.as_ref() == Some(&category))) - } - - pub async fn handle_permission_confirmation( - &self, - request_id: &str, - confirmation: &PermissionConfirmation, - ) -> bool { - let mut pending = self.pending_confirmations.lock().await; - if let Some(tx) = pending.remove(request_id) { - let _ = tx.send(confirmation.clone()); - return true; - } - false - } - - pub async fn ensure_session( - &self, - session_id: Option<&str>, - ) -> Result { - if let Some(session_id) = session_id { - if let Some(response) = self.goose_to_acp_id.lock().await.get(session_id) { - return Ok(response.clone()); - } - } - - let response = self.new_session().await.map_err(|e| { - ProviderError::RequestFailed(format!("Failed to create ACP session: {e}")) - })?; - - if let Some(session_id) = session_id { - self.goose_to_acp_id - .lock() - .await - .insert(session_id.to_string(), response.clone()); - self.acp_to_goose_id - .lock() - .await - .insert(response.session_id.0.to_string(), session_id.to_string()); - - // Initialize model tracking so stream() can detect changes. - let (current_model, _) = resolve_model_info(&self.name, &response)?; - self.session_model - .lock() - .await - .entry(session_id.to_string()) - .or_insert(current_model); - } - - Ok(response) - } - async fn prompt( &self, session_id: SessionId, @@ -497,31 +304,12 @@ impl AcpProvider { Ok(response_rx) } - async fn get_init_session(&self) -> Result<&NewSessionResponse> { - self.init_session - .get_or_try_init(|| async { - let response = self.new_session().await?; - if self.supports_close { - self.close_session_by_acp_id(response.session_id.clone()) - .await?; - } - Ok(response) - }) - .await - } - - async fn close_session_by_acp_id(&self, session_id: SessionId) -> Result<()> { - let (response_tx, response_rx) = oneshot::channel(); - self.tx + fn session_has_config_option(&self, category: SessionConfigOptionCategory) -> bool { + self.session + .response + .config_options .as_ref() - .unwrap() - .send(ClientRequest::CloseSession { - session_id, - response_tx, - }) - .await - .context("ACP client is unavailable")?; - response_rx.await.context("ACP request cancelled")? + .is_some_and(|opts| opts.iter().any(|o| o.category.as_ref() == Some(&category))) } } @@ -536,37 +324,25 @@ impl Provider for AcpProvider { } async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> { - let map = self.goose_to_acp_id.lock().await; - if map.is_empty() { - // Pre-initialization: no ACP session yet, just store the mode. - // The shared Arc> is read at session creation time. - drop(map); - } else { - drop(map); - let mode_str = self.mode_mapping[&mode].clone(); - if self - .session_has_config_option(session_id, SessionConfigOptionCategory::Mode) + let mode_str = self + .mode_mapping + .get(&mode) + .cloned() + .unwrap_or_else(|| format!("{mode:?}")); + + if self.session_has_config_option(SessionConfigOptionCategory::Mode) { + self.send_set_config_option(session_id, "mode".into(), mode_str) .await - { - self.send_set_config_option(session_id, "mode".into(), mode_str) - .await - .map_err(|e| { - ProviderError::RequestFailed(format!("Failed to set mode: {e}")) - })?; - } else { - self.send_set_mode(session_id, mode_str) - .await - .map_err(|e| { - ProviderError::RequestFailed(format!("Failed to set mode: {e}")) - })?; - } + .map_err(|e| ProviderError::RequestFailed(format!("Failed to set mode: {e}")))?; + } else { + self.send_set_mode(session_id, mode_str) + .await + .map_err(|e| ProviderError::RequestFailed(format!("Failed to set mode: {e}")))?; } - let mut current = self - .goose_mode - .lock() - .map_err(|_| ProviderError::RequestFailed("Failed to update mode".into()))?; - *current = mode; + if let Ok(mut guard) = self.goose_mode.lock() { + *guard = mode; + } Ok(()) } @@ -574,76 +350,49 @@ impl Provider for AcpProvider { PermissionRouting::ActionRequired } + fn manages_own_context(&self) -> bool { + true + } + async fn handle_permission_confirmation( &self, request_id: &str, confirmation: &PermissionConfirmation, ) -> bool { - AcpProvider::handle_permission_confirmation(self, request_id, confirmation).await + let mut pending = self.pending_confirmations.lock().await; + if let Some(tx) = pending.remove(request_id) { + let _ = tx.send(confirmation.clone()); + return true; + } + false } async fn stream( &self, - model_config: &ModelConfig, - session_id: &str, + _model_config: &ModelConfig, + _session_id: &str, _system: &str, messages: &[Message], _tools: &[Tool], ) -> Result { - let response = self.ensure_session(Some(session_id)).await?; - - // Provider trait has no update_model — stream() is the only place to forward model changes. - { - let new_model = &model_config.model_name; - let tracked = self.session_model.lock().await.get(session_id).cloned(); - if tracked.as_deref() != Some(new_model) { - if self - .session_has_config_option(session_id, SessionConfigOptionCategory::Model) - .await - { - self.send_set_config_option(session_id, "model".into(), new_model.clone()) - .await - .map_err(|e| { - ProviderError::RequestFailed(format!("Failed to set model: {e}")) - })?; - } else { - self.send_set_model(session_id, new_model.clone()) - .await - .map_err(|e| { - ProviderError::RequestFailed(format!("Failed to set model: {e}")) - })?; - } - self.session_model - .lock() - .await - .insert(session_id.to_string(), new_model.clone()); - } - } + let session_id = self.acp_session_id(); let prompt_blocks = messages_to_prompt(messages); let mut rx = self - .prompt(response.session_id, prompt_blocks) + .prompt(session_id, prompt_blocks) .await .map_err(|e| ProviderError::RequestFailed(format!("Failed to send ACP prompt: {e}")))?; let pending_confirmations = self.pending_confirmations.clone(); - let rejected_tool_calls = self.rejected_tool_calls.clone(); - let permission_mapping = self.permission_mapping.clone(); let goose_mode = *self .goose_mode .lock() .map_err(|_| ProviderError::RequestFailed("goose_mode lock poisoned".into()))?; let reject_all_tools = goose_mode == GooseMode::Chat; - Ok(Box::pin(try_stream! { - // ACP agents execute tools internally. Goose never dispatches tool calls; - // it only sees text, thoughts, and permission requests from the agent. - // - // In Chat mode (reject_all_tools), we suppress all text after a tool - // starts because the agent may send tool results as AcpUpdate::Text, - // bypassing the permission response. let mut suppress_text = false; + let mut rejected_tool_calls: HashSet = HashSet::new(); while let Some(update) = rx.recv().await { match update { @@ -662,12 +411,11 @@ impl Provider for AcpProvider { AcpUpdate::ToolCallStart { id, .. } => { if reject_all_tools { suppress_text = true; - rejected_tool_calls.lock().await.insert(id); + rejected_tool_calls.insert(id); } } AcpUpdate::ToolCallComplete { id, .. } => { - let is_error = rejected_tool_calls.lock().await.remove(&id); - if is_error { + if rejected_tool_calls.remove(&id) { let message = Message::assistant().with_text("Tool call was denied."); yield (Some(message), None); } @@ -675,10 +423,9 @@ impl Provider for AcpProvider { AcpUpdate::PermissionRequest { request, response_tx } => { if let Some(decision) = permission_decision_from_mode(goose_mode) { if decision.should_record_rejection() { - rejected_tool_calls.lock().await.insert(request.tool_call.tool_call_id.0.to_string()); + rejected_tool_calls.insert(request.tool_call.tool_call_id.0.to_string()); } - let response = map_permission_response(&permission_mapping, &request, decision); - let _ = response_tx.send(response); + let _ = response_tx.send(map_permission_response(&request, decision)); continue; } @@ -703,10 +450,9 @@ impl Provider for AcpProvider { let decision = PermissionDecision::from(confirmation.permission); if decision.should_record_rejection() { - rejected_tool_calls.lock().await.insert(request.tool_call.tool_call_id.0.to_string()); + rejected_tool_calls.insert(request.tool_call.tool_call_id.0.to_string()); } - let response = map_permission_response(&permission_mapping, &request, decision); - let _ = response_tx.send(response); + let _ = response_tx.send(map_permission_response(&request, decision)); } AcpUpdate::Complete(_reason) => { break; @@ -720,17 +466,13 @@ impl Provider for AcpProvider { } async fn fetch_supported_models(&self) -> Result, ProviderError> { - let response = self.get_init_session().await.map_err(|e| { - ProviderError::RequestFailed(format!("Failed to create ACP session: {e}")) - })?; - let (_, available) = resolve_model_info(&self.name, response)?; + let (_, available) = resolve_model_info(&self.name, &self.session.response)?; Ok(available) } } impl Drop for AcpProvider { fn drop(&mut self) { - // Join OS thread so session/close completes before runtime exits (reqwest InnerClientHandle pattern). self.tx.take(); if let Some(h) = self.loop_thread.take() { if let Err(e) = h.join() { @@ -807,12 +549,11 @@ impl AcpClientLoop { { let prompt_response_tx = prompt_response_tx.clone(); let reverse_modes = reverse_modes.clone(); + let goose_mode = goose_mode.clone(); async move |notification: SessionNotification, _cx| { if let Some(ref cb) = notification_callback { cb(notification.clone()); } - // stream() reads goose_mode at call time, so it must - // reflect any prior set_mode before the next prompt. match ¬ification.update { SessionUpdate::CurrentModeUpdate(update) => { if let Some(mode) = resolve_mode( @@ -915,7 +656,7 @@ impl AcpClientLoop { sacp::on_receive_request!(), ) .connect_with(transport, async move |cx: ConnectionTo| { - handle_requests(config, cx, rx, prompt_response_tx, init_tx).await + handle_requests(config, goose_mode, cx, rx, prompt_response_tx, init_tx).await }) .await?; @@ -943,7 +684,6 @@ async fn spawn_acp_process(config: &AcpProviderConfig) -> Result { cmd.spawn().context("failed to spawn ACP process") } -// sacp panics on Err from connect_with handlers, so log send failures instead of ?. fn log_undelivered(result: Result<(), E>, method: &str) { if let Err(e) = result { tracing::debug!(method, error = ?e, "response not delivered"); @@ -952,6 +692,7 @@ fn log_undelivered(result: Result<(), E>, method: &str) { async fn handle_requests( config: AcpProviderConfig, + goose_mode: Arc>, cx: ConnectionTo, rx: &mut mpsc::Receiver, prompt_response_tx: Arc>>>, @@ -959,13 +700,16 @@ async fn handle_requests( ) -> Result<(), sacp::Error> { let mut init_tx = Some(init_tx); + let client_capabilities = ClientCapabilities::new(); let init_response: InitializeResponse = cx - .send_request(InitializeRequest::new(ProtocolVersion::LATEST)) + .send_request( + InitializeRequest::new(ProtocolVersion::LATEST) + .client_capabilities(client_capabilities), + ) .block_task() .await .map_err(|err| { let message = format!("ACP {} failed: {err}", AGENT_METHOD_NAMES.initialize); - // Attempt to send a specific error to the ctor waiting on init_rx; if let Some(tx) = init_tx.take() { let _ = tx.send(Err(anyhow::anyhow!(message.clone()))); } @@ -997,7 +741,7 @@ async fn handle_requests( let result = match session { Ok(session) => { session_ids.push(session.session_id.clone()); - apply_session_mode(&config, &cx, session).await + apply_session_mode(&config, &goose_mode, &cx, session).await } Err(err) => Err(anyhow::anyhow!( "ACP {} failed: {err}", @@ -1006,14 +750,6 @@ async fn handle_requests( }; log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_new); } - ClientRequest::ListSessions { response_tx } => { - let result: Result = cx - .send_request(ListSessionsRequest::new()) - .block_task() - .await - .map_err(anyhow::Error::from); - log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_list); - } ClientRequest::SetMode { session_id, mode_id, @@ -1030,22 +766,6 @@ async fn handle_requests( AGENT_METHOD_NAMES.session_set_mode, ); } - ClientRequest::SetModel { - session_id, - model_id, - response_tx, - } => { - let result: Result<()> = cx - .send_request(SetSessionModelRequest::new(session_id, model_id)) - .block_task() - .await - .map(|_| ()) - .map_err(anyhow::Error::from); - log_undelivered( - response_tx.send(result), - AGENT_METHOD_NAMES.session_set_model, - ); - } ClientRequest::SetConfigOption { session_id, config_id, @@ -1065,35 +785,6 @@ async fn handle_requests( AGENT_METHOD_NAMES.session_set_config_option, ); } - ClientRequest::CloseSession { - session_id, - response_tx, - } => { - let result: Result<()> = cx - .send_request(CloseSessionRequest::new(session_id.clone())) - .block_task() - .await - .map(|_| ()) - .map_err(anyhow::Error::from); - session_ids.retain(|s| s != &session_id); - log_undelivered(response_tx.send(result), AGENT_METHOD_NAMES.session_close); - } - ClientRequest::Untyped { - method, - params, - response_tx, - } => { - let result: Result = - match sacp::UntypedMessage::new(&method, params) { - Ok(msg) => cx - .send_request(msg) - .block_task() - .await - .map_err(anyhow::Error::from), - Err(e) => Err(anyhow::Error::from(e)), - }; - log_undelivered(response_tx.send(result), &method); - } ClientRequest::Prompt { session_id, content, @@ -1126,7 +817,6 @@ async fn handle_requests( } } - // After loop exits (channel closed by Drop): if supports_close { for session_id in session_ids { if let Err(e) = cx @@ -1144,10 +834,16 @@ async fn handle_requests( async fn apply_session_mode( config: &AcpProviderConfig, + goose_mode: &Arc>, cx: &ConnectionTo, session: NewSessionResponse, ) -> Result { - if let (Some(mode_id), Some(modes)) = (config.session_mode_id.clone(), session.modes.as_ref()) { + let current_mode = goose_mode.lock().ok().map(|mode| *mode); + let requested_mode_id = current_mode + .and_then(|mode| config.mode_mapping.get(&mode).cloned()) + .or_else(|| config.session_mode_id.clone()); + + if let (Some(mode_id), Some(modes)) = (requested_mode_id, session.modes.as_ref()) { if modes.current_mode_id.0.as_ref() != mode_id.as_str() { let available: Vec = modes .available_modes @@ -1325,32 +1021,45 @@ fn build_action_required_message(request: &RequestPermissionRequest) -> Option Option<(String, Vec)> { + let select = config_options.iter().find_map(|opt| { + if opt.category.as_ref() != Some(&SessionConfigOptionCategory::Model) { + return None; + } + match &opt.kind { + SessionConfigKind::Select(select) => Some(select), + _ => None, + } + })?; + + let current = select.current_value.0.to_string(); + let available = match &select.options { + SessionConfigSelectOptions::Ungrouped(options) => options + .iter() + .map(|option| option.value.0.to_string()) + .collect(), + SessionConfigSelectOptions::Grouped(groups) => groups + .iter() + .flat_map(|group| { + group + .options + .iter() + .map(|option| option.value.0.to_string()) + }) + .collect(), + _ => Vec::new(), + }; + Some((current, available)) +} + fn resolve_model_info( provider_name: &str, response: &NewSessionResponse, ) -> Result<(String, Vec), ProviderError> { if let Some(opts) = &response.config_options { - if let Some(sel) = opts.iter().find_map(|opt| { - if opt.category.as_ref() != Some(&SessionConfigOptionCategory::Model) { - return None; - } - match &opt.kind { - SessionConfigKind::Select(s) => Some(s), - _ => None, - } - }) { - let current = sel.current_value.0.to_string(); - let available = match &sel.options { - SessionConfigSelectOptions::Ungrouped(opts) => { - opts.iter().map(|o| o.value.0.to_string()).collect() - } - SessionConfigSelectOptions::Grouped(groups) => groups - .iter() - .flat_map(|g| g.options.iter().map(|o| o.value.0.to_string())) - .collect(), - _ => vec![], - }; + if let Some((current, available)) = extract_model_info_from_config_options(opts) { return Ok((current, available)); } } @@ -1379,8 +1088,6 @@ fn reverse_mode_mapping( reverse } -// When multiple GooseModes map to the same provider ID (e.g. codex "read-only"), -// prefer the current mode if it's among candidates. fn resolve_mode( reverse_modes: &HashMap>, mode_id: &str, @@ -1406,29 +1113,11 @@ fn permission_decision_from_mode(goose_mode: GooseMode) -> Option, -) -> ListSessionsResponse { - let sessions = response - .sessions - .into_iter() - .filter_map(|mut info| { - let goose_id = acp_to_goose.get(info.session_id.0.as_ref())?; - info.session_id = SessionId::new(goose_id.clone()); - Some(info) - }) - .collect(); - ListSessionsResponse::new(sessions) -} - #[cfg(test)] mod tests { use super::*; use crate::agents::extension::Envs; - use sacp::schema::{SessionConfigOption, SessionConfigSelectOption, SessionInfo}; + use sacp::schema::SessionConfigSelectOption; use test_case::test_case; #[test_case( @@ -1523,64 +1212,6 @@ mod tests { assert!(filtered.is_empty()); } - #[test_case( - ListSessionsResponse::new(vec![ - SessionInfo::new(SessionId::new("20260318_1"), "/Users/codefromthecrypt/oss/goose-2") - .title("Fix login bug".to_string()) - .updated_at("2026-03-18T07:02:42.549655Z".to_string()), - SessionInfo::new(SessionId::new("20260318_2"), "/tmp/test-acpx") - .title("Add caching layer".to_string()) - .updated_at("2026-03-18T07:05:01.123Z".to_string()), - ]), - HashMap::from([ - ("20260318_1".to_string(), "goose-session-1".to_string()), - ("20260318_2".to_string(), "goose-session-2".to_string()), - ]), - ListSessionsResponse::new(vec![ - SessionInfo::new(SessionId::new("goose-session-1"), "/Users/codefromthecrypt/oss/goose-2") - .title("Fix login bug".to_string()) - .updated_at("2026-03-18T07:02:42.549655Z".to_string()), - SessionInfo::new(SessionId::new("goose-session-2"), "/tmp/test-acpx") - .title("Add caching layer".to_string()) - .updated_at("2026-03-18T07:05:01.123Z".to_string()), - ]) - ; "all sessions mapped with all fields preserved" - )] - #[test_case( - ListSessionsResponse::new(vec![ - SessionInfo::new(SessionId::new("20260318_1"), "/Users/codefromthecrypt/oss/goose-2") - .title("Fix login bug".to_string()), - SessionInfo::new(SessionId::new("other-agent-session"), "/tmp/other") - .title("Not our session".to_string()), - ]), - HashMap::from([ - ("20260318_1".to_string(), "goose-session-1".to_string()), - ]), - ListSessionsResponse::new(vec![ - SessionInfo::new(SessionId::new("goose-session-1"), "/Users/codefromthecrypt/oss/goose-2") - .title("Fix login bug".to_string()), - ]) - ; "unmapped sessions filtered out" - )] - #[test_case( - ListSessionsResponse::new(vec![ - SessionInfo::new(SessionId::new("20260318_1"), "/Users/codefromthecrypt/oss/goose-2") - .title("ACP Session".to_string()) - .updated_at("2026-03-18T01:29:02.141700Z".to_string()), - ]), - HashMap::new(), - ListSessionsResponse::new(vec![]) - ; "empty map returns empty list" - )] - fn test_map_sessions_to_goose_ids( - response: ListSessionsResponse, - acp_to_goose: HashMap, - expected: ListSessionsResponse, - ) { - let result = map_sessions_to_goose_ids(response, &acp_to_goose); - assert_eq!(result, expected); - } - #[test_case(GooseMode::Auto => Some(PermissionDecision::AllowOnce) ; "auto allows")] #[test_case(GooseMode::Chat => Some(PermissionDecision::RejectOnce) ; "chat rejects")] #[test_case(GooseMode::Approve => None ; "approve defers")] @@ -1699,7 +1330,6 @@ mod tests { resolve_model_info("test", &response) } - // Codex mapping: read-only maps to both Approve and Chat. fn codex_reverse_modes() -> HashMap> { HashMap::from([ ("full-access".to_string(), vec![GooseMode::Auto]), @@ -1735,9 +1365,7 @@ mod tests { let reverse_modes = codex_reverse_modes(); let current = Arc::new(Mutex::new(current)); let result = resolve_mode(&reverse_modes, mode_id, ¤t); - // For the fallback case, just check we got *some* candidate (order is nondeterministic). if mode_id == "read-only" && expected == Some(GooseMode::Approve) { - // Current (Auto) not in candidates — any candidate is valid. assert!(result == Some(GooseMode::Approve) || result == Some(GooseMode::Chat)); } else { assert_eq!(result, expected); diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 92de31bc..57446f07 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -1578,6 +1578,7 @@ impl Agent { no_tools_called = false; } } + #[allow(unused_variables)] Err(ref provider_err @ ProviderError::ContextLengthExceeded(_)) => { #[cfg(feature = "telemetry")] crate::posthog::emit_error(provider_err.telemetry_type(), &provider_err.to_string()); diff --git a/crates/goose/src/agents/platform_extensions/orchestrator.rs b/crates/goose/src/agents/platform_extensions/orchestrator.rs index b01b0cb6..45800343 100644 --- a/crates/goose/src/agents/platform_extensions/orchestrator.rs +++ b/crates/goose/src/agents/platform_extensions/orchestrator.rs @@ -2,11 +2,13 @@ use crate::agents::extension::PlatformExtensionContext; use crate::agents::mcp_client::{Error, McpClientTrait}; use crate::agents::tool_execution::ToolCallContext; use crate::agents::{AgentEvent, SessionConfig}; -use crate::config::GooseMode; +use crate::config::{Config, ExtensionConfig, GooseMode}; use crate::context_mgmt::format_message_for_compacting; use crate::conversation::message::Message; use crate::execution::manager::AgentManager; +use crate::providers; use crate::providers::base::Provider; +use crate::session::extension_data::EnabledExtensionsState; use crate::session::session_manager::SessionType; use anyhow::Result; use async_trait::async_trait; @@ -139,6 +141,11 @@ impl OrchestratorClient { .ok_or_else(|| "Provider not available".to_string()) } + fn parent_extensions(&self) -> Vec { + let extension_data = self.context.session.as_ref().map(|s| &s.extension_data); + EnabledExtensionsState::extensions_or_default(extension_data, Config::global()) + } + async fn handle_list_sessions( &self, arguments: Option, @@ -397,8 +404,15 @@ impl OrchestratorClient { .await .map_err(|e| format!("Failed to create agent: {}", e))?; - // Inherit the orchestrator's provider and model - let provider = self.get_provider().await?; + let parent_provider = self.get_provider().await?; + let extensions = self.parent_extensions(); + let provider = providers::create( + parent_provider.get_name(), + parent_provider.get_model_config(), + extensions, + ) + .await + .map_err(|e| format!("Failed to create provider for new agent: {}", e))?; agent .update_provider(provider, &session.id) .await @@ -432,11 +446,20 @@ impl OrchestratorClient { .map_err(|e| format!("Failed to get agent for session '{}': {}", session_id, e))?; if agent.provider().await.is_err() { - if let Ok(provider) = self.get_provider().await { - agent - .update_provider(provider, &session_id) - .await - .map_err(|e| format!("Failed to set provider: {}", e))?; + if let Ok(parent_provider) = self.get_provider().await { + let extensions = self.parent_extensions(); + if let Ok(provider) = providers::create( + parent_provider.get_name(), + parent_provider.get_model_config(), + extensions, + ) + .await + { + agent + .update_provider(provider, &session_id) + .await + .map_err(|e| format!("Failed to set provider: {}", e))?; + } } } diff --git a/crates/goose/src/context_mgmt/mod.rs b/crates/goose/src/context_mgmt/mod.rs index fd7c7857..1cde777b 100644 --- a/crates/goose/src/context_mgmt/mod.rs +++ b/crates/goose/src/context_mgmt/mod.rs @@ -188,6 +188,10 @@ pub async fn check_if_compaction_needed( threshold_override: Option, session: &crate::session::Session, ) -> Result { + if provider.manages_own_context() { + return Ok(false); + } + let messages = conversation.messages(); let config = Config::global(); let threshold = threshold_override.unwrap_or_else(|| { diff --git a/crates/goose/src/providers/amp_acp.rs b/crates/goose/src/providers/amp_acp.rs new file mode 100644 index 00000000..987b370b --- /dev/null +++ b/crates/goose/src/providers/amp_acp.rs @@ -0,0 +1,74 @@ +use anyhow::Result; +use futures::future::BoxFuture; +use std::collections::HashMap; +use std::path::PathBuf; + +use crate::acp::{ + extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, +}; +use crate::config::search_path::SearchPaths; +use crate::config::{Config, GooseMode}; +use crate::model::ModelConfig; +use crate::providers::base::{ProviderDef, ProviderMetadata}; + +const AMP_ACP_PROVIDER_NAME: &str = "amp-acp"; +const AMP_ACP_DOC_URL: &str = "https://ampcode.com"; +const AMP_ACP_BINARY: &str = "amp-acp"; + +pub struct AmpAcpProvider; + +impl ProviderDef for AmpAcpProvider { + type Provider = AcpProvider; + + fn metadata() -> ProviderMetadata { + ProviderMetadata::new( + AMP_ACP_PROVIDER_NAME, + "Amp", + "Use goose with your Amp subscription via the amp-acp adapter.", + ACP_CURRENT_MODEL, + vec![], + AMP_ACP_DOC_URL, + vec![], + ) + .with_setup_steps(vec![ + "Install the Amp CLI: `curl -fsSL https://ampcode.com/install.sh | bash`", + "Install the ACP adapter: `npm install -g amp-acp`", + "Ensure your Amp CLI is authenticated (run `amp` to verify)", + "Set in your goose config file (`~/.config/goose/config.yaml` on macOS/Linux):\n GOOSE_PROVIDER: amp-acp\n GOOSE_MODEL: current", + "Restart goose for changes to take effect", + ]) + } + + fn from_env( + model: ModelConfig, + extensions: Vec, + ) -> BoxFuture<'static, Result> { + Box::pin(async move { + let config = Config::global(); + let resolved_command = SearchPaths::builder().with_npm().resolve(AMP_ACP_BINARY)?; + let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); + + let mode_mapping = HashMap::from([ + (GooseMode::Auto, "auto".to_string()), + (GooseMode::Approve, "approve".to_string()), + (GooseMode::SmartApprove, "smart-approve".to_string()), + (GooseMode::Chat, "chat".to_string()), + ]); + + let provider_config = AcpProviderConfig { + command: resolved_command, + args: vec![], + env: vec![], + env_remove: vec![], + work_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")), + mcp_servers: extension_configs_to_mcp_servers(&extensions), + session_mode_id: Some(mode_mapping[&goose_mode].clone()), + mode_mapping, + notification_callback: None, + }; + + let metadata = Self::metadata(); + AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + }) + } +} diff --git a/crates/goose/src/providers/claude_acp.rs b/crates/goose/src/providers/claude_acp.rs index 3d8c6a95..9fb25be2 100644 --- a/crates/goose/src/providers/claude_acp.rs +++ b/crates/goose/src/providers/claude_acp.rs @@ -4,8 +4,7 @@ use std::collections::HashMap; use std::path::PathBuf; use crate::acp::{ - extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, PermissionMapping, - ACP_CURRENT_MODEL, + extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, }; use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; @@ -51,13 +50,6 @@ impl ProviderDef for ClaudeAcpProvider { .resolve(CLAUDE_ACP_BINARY)?; let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); - // claude-agent-acp permission option_ids - let permission_mapping = PermissionMapping { - allow_option_id: Some("allow".to_string()), - reject_option_id: Some("reject".to_string()), - rejected_tool_status: sacp::schema::ToolCallStatus::Failed, - }; - let mode_mapping = HashMap::from([ // Closest to "autonomous": bypassPermissions skips confirmations. (GooseMode::Auto, "bypassPermissions".to_string()), @@ -79,7 +71,6 @@ impl ProviderDef for ClaudeAcpProvider { mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), mode_mapping, - permission_mapping, notification_callback: None, }; diff --git a/crates/goose/src/providers/codex_acp.rs b/crates/goose/src/providers/codex_acp.rs index bdecb7db..f28bd029 100644 --- a/crates/goose/src/providers/codex_acp.rs +++ b/crates/goose/src/providers/codex_acp.rs @@ -4,8 +4,7 @@ use std::collections::HashMap; use std::path::PathBuf; use crate::acp::{ - extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, PermissionMapping, - ACP_CURRENT_MODEL, + extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, }; use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; @@ -74,13 +73,6 @@ impl ProviderDef for CodexAcpProvider { ]); } - // codex-acp permission option_ids - let permission_mapping = PermissionMapping { - allow_option_id: Some("approved".to_string()), - reject_option_id: Some("abort".to_string()), - rejected_tool_status: sacp::schema::ToolCallStatus::Failed, - }; - // Chat and Approve both map to "read-only". let mode_mapping = HashMap::from([ (GooseMode::Auto, "full-access".to_string()), @@ -99,7 +91,6 @@ impl ProviderDef for CodexAcpProvider { // Disabled until https://github.com/zed-industries/codex-acp/issues/179 is fixed. session_mode_id: None, mode_mapping, - permission_mapping, notification_callback: None, }; diff --git a/crates/goose/src/providers/copilot_acp.rs b/crates/goose/src/providers/copilot_acp.rs index 662a2d65..9810547f 100644 --- a/crates/goose/src/providers/copilot_acp.rs +++ b/crates/goose/src/providers/copilot_acp.rs @@ -4,8 +4,7 @@ use std::collections::HashMap; use std::path::PathBuf; use crate::acp::{ - extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, PermissionMapping, - ACP_CURRENT_MODEL, + extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, }; use crate::config::search_path::SearchPaths; use crate::config::{Config, GooseMode}; @@ -54,10 +53,6 @@ impl ProviderDef for CopilotAcpProvider { .resolve(COPILOT_ACP_BINARY)?; let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); - // Copilot uses standard ACP permission option IDs (allow_once, - // allow_always, reject_once) so kind-based fallback handles them. - let permission_mapping = PermissionMapping::default(); - let mut args = vec!["--acp".to_string()]; if model.model_name != ACP_CURRENT_MODEL { args.push("--model".to_string()); @@ -82,7 +77,6 @@ impl ProviderDef for CopilotAcpProvider { mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), mode_mapping, - permission_mapping, notification_callback: None, }; diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index e7ff25ec..34febc8f 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -7,6 +7,7 @@ use super::local_inference::LocalInferenceProvider; #[cfg(feature = "aws-providers")] use super::sagemaker_tgi::SageMakerTgiProvider; use super::{ + amp_acp::AmpAcpProvider, anthropic::AnthropicProvider, avian::AvianProvider, azure::AzureProvider, @@ -29,6 +30,7 @@ use super::{ ollama::OllamaProvider, openai::OpenAiProvider, openrouter::OpenRouterProvider, + pi_acp::PiAcpProvider, provider_registry::ProviderRegistry, snowflake::SnowflakeProvider, tetrate::TetrateProvider, @@ -49,6 +51,7 @@ static REGISTRY: OnceCell> = OnceCell::const_new(); async fn init_registry() -> RwLock { let mut registry = ProviderRegistry::new().with_providers(|registry| { + registry.register::(false); registry.register::(true); registry.register::(false); registry.register::(false); @@ -74,6 +77,7 @@ async fn init_registry() -> RwLock { registry.register::(true); registry.register::(true); registry.register::(true); + registry.register::(false); #[cfg(feature = "aws-providers")] registry.register::(false); registry.register::(false); diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 6934c4d6..a49b5a48 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -1,3 +1,4 @@ +pub mod amp_acp; pub mod anthropic; pub mod api_client; pub mod avian; @@ -36,6 +37,7 @@ pub mod ollama; pub mod openai; pub mod openai_compatible; pub mod openrouter; +pub mod pi_acp; pub mod provider_registry; pub mod provider_test; mod retry; diff --git a/crates/goose/src/providers/pi_acp.rs b/crates/goose/src/providers/pi_acp.rs new file mode 100644 index 00000000..5c3eaa72 --- /dev/null +++ b/crates/goose/src/providers/pi_acp.rs @@ -0,0 +1,73 @@ +use anyhow::Result; +use futures::future::BoxFuture; +use std::collections::HashMap; +use std::path::PathBuf; + +use crate::acp::{ + extension_configs_to_mcp_servers, AcpProvider, AcpProviderConfig, ACP_CURRENT_MODEL, +}; +use crate::config::search_path::SearchPaths; +use crate::config::{Config, GooseMode}; +use crate::model::ModelConfig; +use crate::providers::base::{ProviderDef, ProviderMetadata}; + +const PI_ACP_PROVIDER_NAME: &str = "pi-acp"; +const PI_ACP_DOC_URL: &str = "https://github.com/anthropics/pi"; +const PI_ACP_BINARY: &str = "pi-acp"; + +pub struct PiAcpProvider; + +impl ProviderDef for PiAcpProvider { + type Provider = AcpProvider; + + fn metadata() -> ProviderMetadata { + ProviderMetadata::new( + PI_ACP_PROVIDER_NAME, + "Pi", + "Use goose with Pi via the pi-acp adapter.", + ACP_CURRENT_MODEL, + vec![], + PI_ACP_DOC_URL, + vec![], + ) + .with_setup_steps(vec![ + "Install the Pi CLI and the pi-acp adapter", + "Ensure your Pi CLI is authenticated (run `pi` to verify)", + "Set in your goose config file (`~/.config/goose/config.yaml` on macOS/Linux):\n GOOSE_PROVIDER: pi-acp\n GOOSE_MODEL: current", + "Restart goose for changes to take effect", + ]) + } + + fn from_env( + model: ModelConfig, + extensions: Vec, + ) -> BoxFuture<'static, Result> { + Box::pin(async move { + let config = Config::global(); + let resolved_command = SearchPaths::builder().with_npm().resolve(PI_ACP_BINARY)?; + let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); + + let mode_mapping = HashMap::from([ + (GooseMode::Auto, "auto".to_string()), + (GooseMode::Approve, "approve".to_string()), + (GooseMode::SmartApprove, "smart-approve".to_string()), + (GooseMode::Chat, "chat".to_string()), + ]); + + let provider_config = AcpProviderConfig { + command: resolved_command, + args: vec![], + env: vec![], + env_remove: vec![], + work_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")), + mcp_servers: extension_configs_to_mcp_servers(&extensions), + session_mode_id: Some(mode_mapping[&goose_mode].clone()), + mode_mapping, + notification_callback: None, + }; + + let metadata = Self::metadata(); + AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + }) + } +} diff --git a/crates/goose/src/providers/provider_registry.rs b/crates/goose/src/providers/provider_registry.rs index c6420344..860421c4 100644 --- a/crates/goose/src/providers/provider_registry.rs +++ b/crates/goose/src/providers/provider_registry.rs @@ -23,6 +23,10 @@ pub struct ProviderEntry { } impl ProviderEntry { + pub fn metadata(&self) -> &ProviderMetadata { + &self.metadata + } + pub async fn create_with_default_model( &self, extensions: Vec, diff --git a/crates/goose/src/scheduler.rs b/crates/goose/src/scheduler.rs index 970e32c7..9c6793c6 100644 --- a/crates/goose/src/scheduler.rs +++ b/crates/goose/src/scheduler.rs @@ -839,6 +839,7 @@ async fn execute_job( } drop(jobs_guard); + #[cfg(feature = "telemetry")] let start_time = std::time::Instant::now(); #[cfg(feature = "telemetry")] tokio::spawn(async move { diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index 5f95bc97..3da0a0a7 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -232,6 +232,11 @@ impl<'a> SessionUpdateBuilder<'a> { self } + pub fn clear_model_config(mut self) -> Self { + self.model_config = Some(None); + self + } + pub fn goose_mode(mut self, mode: GooseMode) -> Self { self.goose_mode = Some(mode); self diff --git a/documentation/docs/guides/acp-providers.md b/documentation/docs/guides/acp-providers.md index 35c078d2..40ec9325 100644 --- a/documentation/docs/guides/acp-providers.md +++ b/documentation/docs/guides/acp-providers.md @@ -22,6 +22,16 @@ ACP providers let you use goose with your existing Claude Code, ChatGPT Plus/Pro ## Available ACP Providers +### Amp ACP + +Wraps [amp-acp](https://www.npmjs.com/package/amp-acp), an ACP adapter for [Amp](https://ampcode.com). Uses your existing Amp subscription. + +**Requirements:** +- Node.js and npm +- Amp CLI installed (`curl -fsSL https://ampcode.com/install.sh | bash`) +- ACP adapter installed (`npm install -g amp-acp`) +- Authenticated with your Amp account (`amp` CLI working) + ### Claude ACP Wraps [claude-agent-acp](https://github.com/zed-industries/claude-agent-acp), an ACP adapter for Anthropic's Claude Code. Uses the same Claude subscription as the deprecated `claude-code` CLI provider. @@ -49,8 +59,44 @@ Uses Google's [Gemini CLI](https://github.com/google-gemini/gemini-cli) directly - Gemini CLI installed (`npm install -g @google/gemini-cli`) - Authenticated with your Google account (run `gemini` once to authenticate via browser) +### Pi ACP + +Wraps `pi-acp`, an ACP adapter for Pi. Uses your existing Pi installation. + +**Requirements:** +- Pi CLI installed +- ACP adapter installed (`pi-acp` binary available) +- Authenticated with your Pi account (`pi` CLI working) + ## Setup Instructions +### Amp ACP + +1. **Install the Amp CLI** + + ```bash + curl -fsSL https://ampcode.com/install.sh | bash + ``` + +2. **Install the ACP adapter** + + ```bash + npm install -g amp-acp + ``` + +3. **Authenticate with Amp** + + Run `amp` and follow the authentication prompts. + +4. **Configure goose** + + Set the provider environment variable: + ```bash + export GOOSE_PROVIDER=amp-acp + ``` + + Or configure through the goose CLI using `goose configure`. + ### Claude ACP 1. **Install the ACP adapter** @@ -159,6 +205,25 @@ Uses Google's [Gemini CLI](https://github.com/google-gemini/gemini-cli) directly │ default ``` +### Pi ACP + +1. **Install the Pi CLI and ACP adapter** + + Install the `pi` CLI and the `pi-acp` ACP adapter following the project's installation instructions. + +2. **Authenticate with Pi** + + Run `pi` and follow the authentication prompts. + +3. **Configure goose** + + Set the provider environment variable: + ```bash + export GOOSE_PROVIDER=pi-acp + ``` + + Or configure through the goose CLI using `goose configure`. + ## Usage Examples ### Basic Usage @@ -191,6 +256,14 @@ GOOSE_PROVIDER=gemini-acp goose run \ ## Configuration Options +### Amp ACP Configuration + +| Environment Variable | Description | Default | +|----------------------|-------------------|-----------| +| `GOOSE_PROVIDER` | Set to `amp-acp` | None | +| `GOOSE_MODEL` | Model to use | `current` | +| `GOOSE_MODE` | Permission mode | `auto` | + ### Claude ACP Configuration | Environment Variable | Description | Default | @@ -259,13 +332,21 @@ See [codex-acp](https://github.com/zed-industries/codex-acp) for approval policy See the [Gemini CLI documentation](https://github.com/google-gemini/gemini-cli) for approval mode details. +### Pi ACP Configuration + +| Environment Variable | Description | Default | +|----------------------|------------------|-----------| +| `GOOSE_PROVIDER` | Set to `pi-acp` | None | +| `GOOSE_MODEL` | Model to use | `current` | +| `GOOSE_MODE` | Permission mode | `auto` | + ## Error Handling -ACP providers depend on external npm packages, so ensure: +ACP providers depend on external binaries, so ensure: -- The ACP agent binary is installed and in your PATH (`claude-agent-acp`, `codex-acp`, or `gemini`) +- The ACP agent binary is installed and in your PATH (`amp-acp`, `claude-agent-acp`, `codex-acp`, `gemini`, `pi-acp`, or `copilot`) - The underlying CLI tool is authenticated and working - Subscription limits are not exceeded -- Node.js and npm are installed +- Node.js and npm are installed (for npm-distributed adapters) -If goose can't find the binary, session startup will fail with an error. Run `which claude-agent-acp`, `which codex-acp`, or `which gemini` to verify installation. +If goose can't find the binary, session startup will fail with an error. Run `which ` to verify installation. diff --git a/ui/acp/src/generated/client.gen.ts b/ui/acp/src/generated/client.gen.ts index dd46316c..f89d8536 100644 --- a/ui/acp/src/generated/client.gen.ts +++ b/ui/acp/src/generated/client.gen.ts @@ -9,6 +9,8 @@ export interface ExtMethodProvider { import type { AddExtensionRequest, + CheckSecretRequest, + CheckSecretResponse, DeleteSessionRequest, ExportSessionRequest, ExportSessionResponse, @@ -20,18 +22,32 @@ import type { GetToolsResponse, ImportSessionRequest, ImportSessionResponse, + ListProvidersRequest, + ListProvidersResponse, + ReadConfigRequest, + ReadConfigResponse, ReadResourceRequest, ReadResourceResponse, + RemoveConfigRequest, RemoveExtensionRequest, + RemoveSecretRequest, + UpdateProviderRequest, + UpdateProviderResponse, UpdateWorkingDirRequest, + UpsertConfigRequest, + UpsertSecretRequest, } from './types.gen.js'; import { + zCheckSecretResponse, zExportSessionResponse, zGetExtensionsResponse, zGetSessionResponse, zGetToolsResponse, zImportSessionResponse, + zListProvidersResponse, + zReadConfigResponse, zReadResourceResponse, + zUpdateProviderResponse, } from './zod.gen.js'; export class GooseExtClient { @@ -90,4 +106,51 @@ export class GooseExtClient { const raw = await this.conn.extMethod("_goose/config/extensions", params); return zGetExtensionsResponse.parse(raw) as GetExtensionsResponse; } + + async GooseSessionProviderUpdate( + params: UpdateProviderRequest, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/session/provider/update", + params, + ); + return zUpdateProviderResponse.parse(raw) as UpdateProviderResponse; + } + + async GooseProvidersList( + params: ListProvidersRequest, + ): Promise { + const raw = await this.conn.extMethod("_goose/providers/list", params); + return zListProvidersResponse.parse(raw) as ListProvidersResponse; + } + + async GooseConfigRead( + params: ReadConfigRequest, + ): Promise { + const raw = await this.conn.extMethod("_goose/config/read", params); + return zReadConfigResponse.parse(raw) as ReadConfigResponse; + } + + async GooseConfigUpsert(params: UpsertConfigRequest): Promise { + await this.conn.extMethod("_goose/config/upsert", params); + } + + async GooseConfigRemove(params: RemoveConfigRequest): Promise { + await this.conn.extMethod("_goose/config/remove", params); + } + + async GooseSecretCheck( + params: CheckSecretRequest, + ): Promise { + const raw = await this.conn.extMethod("_goose/secret/check", params); + return zCheckSecretResponse.parse(raw) as CheckSecretResponse; + } + + async GooseSecretUpsert(params: UpsertSecretRequest): Promise { + await this.conn.extMethod("_goose/secret/upsert", params); + } + + async GooseSecretRemove(params: RemoveSecretRequest): Promise { + await this.conn.extMethod("_goose/secret/remove", params); + } } diff --git a/ui/acp/src/generated/index.ts b/ui/acp/src/generated/index.ts index f82eb9b9..e6c1f0d6 100644 --- a/ui/acp/src/generated/index.ts +++ b/ui/acp/src/generated/index.ts @@ -1,6 +1,6 @@ // This file is auto-generated by @hey-api/openapi-ts -export type { AddExtensionRequest, DeleteSessionRequest, EmptyResponse, ExportSessionRequest, ExportSessionResponse, ExtRequest, ExtResponse, GetExtensionsRequest, GetExtensionsResponse, GetSessionRequest, GetSessionResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, ImportSessionResponse, ReadResourceRequest, ReadResourceResponse, RemoveExtensionRequest, UpdateWorkingDirRequest } from './types.gen.js'; +export type { AddExtensionRequest, CheckSecretRequest, CheckSecretResponse, DeleteSessionRequest, EmptyResponse, ExportSessionRequest, ExportSessionResponse, ExtRequest, ExtResponse, GetExtensionsRequest, GetExtensionsResponse, GetSessionRequest, GetSessionResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, ImportSessionResponse, ListProvidersRequest, ListProvidersResponse, ProviderListEntry, ReadConfigRequest, ReadConfigResponse, ReadResourceRequest, ReadResourceResponse, RemoveConfigRequest, RemoveExtensionRequest, RemoveSecretRequest, UpdateProviderRequest, UpdateProviderResponse, UpdateWorkingDirRequest, UpsertConfigRequest, UpsertSecretRequest } from './types.gen.js'; export const GOOSE_EXT_METHODS = [ { @@ -53,6 +53,46 @@ export const GOOSE_EXT_METHODS = [ requestType: "GetExtensionsRequest", responseType: "GetExtensionsResponse", }, + { + method: "_goose/session/provider/update", + requestType: "UpdateProviderRequest", + responseType: "UpdateProviderResponse", + }, + { + method: "_goose/providers/list", + requestType: "ListProvidersRequest", + responseType: "ListProvidersResponse", + }, + { + method: "_goose/config/read", + requestType: "ReadConfigRequest", + responseType: "ReadConfigResponse", + }, + { + method: "_goose/config/upsert", + requestType: "UpsertConfigRequest", + responseType: "EmptyResponse", + }, + { + method: "_goose/config/remove", + requestType: "RemoveConfigRequest", + responseType: "EmptyResponse", + }, + { + method: "_goose/secret/check", + requestType: "CheckSecretRequest", + responseType: "CheckSecretResponse", + }, + { + method: "_goose/secret/upsert", + requestType: "UpsertSecretRequest", + responseType: "EmptyResponse", + }, + { + method: "_goose/secret/remove", + requestType: "RemoveSecretRequest", + responseType: "EmptyResponse", + }, ] as const; export type GooseExtMethod = (typeof GOOSE_EXT_METHODS)[number]; diff --git a/ui/acp/src/generated/types.gen.ts b/ui/acp/src/generated/types.gen.ts index 4ce7b461..4d00fa20 100644 --- a/ui/acp/src/generated/types.gen.ts +++ b/ui/acp/src/generated/types.gen.ts @@ -145,17 +145,117 @@ export type GetExtensionsResponse = { warnings: Array; }; +/** + * Atomically update the provider for a live session. + */ +export type UpdateProviderRequest = { + sessionId: string; + provider: string; + model?: string | null; + contextLimit?: number | null; + requestParams?: { + [key: string]: unknown; + } | null; +}; + +/** + * Provider update response. + */ +export type UpdateProviderResponse = { + /** + * Refreshed session config options after the provider/model change. + */ + configOptions: Array; +}; + +/** + * List providers available through goose, including the config-default sentinel. + */ +export type ListProvidersRequest = { + [key: string]: unknown; +}; + +/** + * Provider list response. + */ +export type ListProvidersResponse = { + providers: Array; +}; + +export type ProviderListEntry = { + id: string; + label: string; +}; + +/** + * Read a single non-secret config value. + */ +export type ReadConfigRequest = { + key: string; +}; + +/** + * Config read response. + */ +export type ReadConfigResponse = { + value?: unknown; +}; + +/** + * Upsert a single non-secret config value. + */ +export type UpsertConfigRequest = { + key: string; + value: unknown; +}; + +/** + * Remove a single non-secret config value. + */ +export type RemoveConfigRequest = { + key: string; +}; + +/** + * Check whether a secret exists. Never returns the actual value. + */ +export type CheckSecretRequest = { + key: string; +}; + +/** + * Secret check response. + */ +export type CheckSecretResponse = { + exists: boolean; +}; + +/** + * Set a secret value (write-only). + */ +export type UpsertSecretRequest = { + key: string; + value: unknown; +}; + +/** + * Remove a secret. + */ +export type RemoveSecretRequest = { + key: string; +}; + export type ExtRequest = { id: string; method: string; - params?: AddExtensionRequest | RemoveExtensionRequest | GetToolsRequest | ReadResourceRequest | UpdateWorkingDirRequest | GetSessionRequest | DeleteSessionRequest | ExportSessionRequest | ImportSessionRequest | GetExtensionsRequest | { + params?: AddExtensionRequest | RemoveExtensionRequest | GetToolsRequest | ReadResourceRequest | UpdateWorkingDirRequest | GetSessionRequest | DeleteSessionRequest | ExportSessionRequest | ImportSessionRequest | GetExtensionsRequest | UpdateProviderRequest | ListProvidersRequest | ReadConfigRequest | UpsertConfigRequest | RemoveConfigRequest | CheckSecretRequest | UpsertSecretRequest | RemoveSecretRequest | { [key: string]: unknown; } | null; }; export type ExtResponse = { id: string; - result?: EmptyResponse | GetToolsResponse | ReadResourceResponse | GetSessionResponse | ExportSessionResponse | ImportSessionResponse | GetExtensionsResponse | unknown; + result?: EmptyResponse | GetToolsResponse | ReadResourceResponse | GetSessionResponse | ExportSessionResponse | ImportSessionResponse | GetExtensionsResponse | UpdateProviderResponse | ListProvidersResponse | ReadConfigResponse | CheckSecretResponse | unknown; } | { error: { code: number; diff --git a/ui/acp/src/generated/zod.gen.ts b/ui/acp/src/generated/zod.gen.ts index 834d26e1..e1f3768c 100644 --- a/ui/acp/src/generated/zod.gen.ts +++ b/ui/acp/src/generated/zod.gen.ts @@ -124,6 +124,108 @@ export const zGetExtensionsResponse = z.object({ warnings: z.array(z.string()) }); +/** + * Atomically update the provider for a live session. + */ +export const zUpdateProviderRequest = z.object({ + sessionId: z.string(), + provider: z.string(), + model: z.union([ + z.string(), + z.null() + ]).optional(), + contextLimit: z.union([ + z.number().int().gte(0), + z.null() + ]).optional(), + requestParams: z.union([ + z.record(z.unknown()), + z.null() + ]).optional() +}); + +/** + * Provider update response. + */ +export const zUpdateProviderResponse = z.object({ + configOptions: z.array(z.unknown()) +}); + +/** + * List providers available through goose, including the config-default sentinel. + */ +export const zListProvidersRequest = z.record(z.unknown()); + +export const zProviderListEntry = z.object({ + id: z.string(), + label: z.string() +}); + +/** + * Provider list response. + */ +export const zListProvidersResponse = z.object({ + providers: z.array(zProviderListEntry) +}); + +/** + * Read a single non-secret config value. + */ +export const zReadConfigRequest = z.object({ + key: z.string() +}); + +/** + * Config read response. + */ +export const zReadConfigResponse = z.object({ + value: z.unknown().optional().default(null) +}); + +/** + * Upsert a single non-secret config value. + */ +export const zUpsertConfigRequest = z.object({ + key: z.string(), + value: z.unknown() +}); + +/** + * Remove a single non-secret config value. + */ +export const zRemoveConfigRequest = z.object({ + key: z.string() +}); + +/** + * Check whether a secret exists. Never returns the actual value. + */ +export const zCheckSecretRequest = z.object({ + key: z.string() +}); + +/** + * Secret check response. + */ +export const zCheckSecretResponse = z.object({ + exists: z.boolean() +}); + +/** + * Set a secret value (write-only). + */ +export const zUpsertSecretRequest = z.object({ + key: z.string(), + value: z.unknown() +}); + +/** + * Remove a secret. + */ +export const zRemoveSecretRequest = z.object({ + key: z.string() +}); + export const zExtRequest = z.object({ id: z.string(), method: z.string(), @@ -138,7 +240,15 @@ export const zExtRequest = z.object({ zDeleteSessionRequest, zExportSessionRequest, zImportSessionRequest, - zGetExtensionsRequest + zGetExtensionsRequest, + zUpdateProviderRequest, + zListProvidersRequest, + zReadConfigRequest, + zUpsertConfigRequest, + zRemoveConfigRequest, + zCheckSecretRequest, + zUpsertSecretRequest, + zRemoveSecretRequest ]), z.union([ z.record(z.unknown()), @@ -158,7 +268,11 @@ export const zExtResponse = z.union([ zGetSessionResponse, zExportSessionResponse, zImportSessionResponse, - zGetExtensionsResponse + zGetExtensionsResponse, + zUpdateProviderResponse, + zListProvidersResponse, + zReadConfigResponse, + zCheckSecretResponse ]), z.unknown() ]).optional()