diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 0089b8a7..64af66d7 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -13,13 +13,6 @@ use crate::config::paths::Paths; use crate::config::permission::PermissionManager; use crate::config::{Config, GooseMode}; use crate::conversation::message::{ActionRequiredData, Message, MessageContent}; -#[cfg(feature = "local-inference")] -use crate::dictation::providers::transcribe_local; -use crate::dictation::providers::{ - all_providers, is_configured, transcribe_with_provider, DictationProvider, -}; -#[cfg(feature = "local-inference")] -use crate::dictation::whisper; use crate::mcp_utils::ToolResult; use crate::permission::permission_confirmation::PrincipalType; use crate::permission::{Permission, PermissionConfirmation}; @@ -36,7 +29,6 @@ use fs_err as fs; use futures::future::{BoxFuture, Either}; use futures::stream::{self, StreamExt}; use futures::FutureExt; -use goose_acp_macros::custom_methods; use rmcp::model::{ AnnotateAble, CallToolResult, RawContent, RawTextContent, ResourceContents, Role, }; @@ -74,6 +66,18 @@ use tokio_util::sync::CancellationToken; use tracing::{debug, error, info, warn}; use url::Url; +mod config; +mod custom_dispatch; +mod dictation; +mod dispatch; +mod extensions; +mod providers; +mod resources; +mod secrets; +mod sessions; +mod sources; +mod tools; + pub type AcpProviderFactory = Arc< dyn Fn( String, @@ -114,12 +118,6 @@ impl ResultExt for Result { const DEFAULT_PROVIDER_ID: &str = "goose"; const DEFAULT_PROVIDER_LABEL: &str = "Goose (Default)"; -const OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "OPENAI_TRANSCRIPTION_MODEL"; -const GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "GROQ_TRANSCRIPTION_MODEL"; -const ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "ELEVENLABS_TRANSCRIPTION_MODEL"; -const OPENAI_TRANSCRIPTION_MODEL: &str = "whisper-1"; -const GROQ_TRANSCRIPTION_MODEL: &str = "whisper-large-v3-turbo"; -const ELEVENLABS_TRANSCRIPTION_MODEL: &str = "scribe_v1"; const PROVIDER_CONFIG_STATUS_CHECK_CONCURRENCY: usize = 16; async fn ensure_refresh_identity_current( @@ -583,145 +581,6 @@ fn builtin_to_extension_config(name: &str) -> ExtensionConfig { } } -fn inventory_entry_to_dto(entry: ProviderInventoryEntry) -> ProviderInventoryEntryDto { - let stale = ProviderInventoryService::is_stale(&entry); - ProviderInventoryEntryDto { - provider_id: entry.provider_id, - provider_name: entry.provider_name, - description: entry.description, - default_model: entry.default_model, - configured: entry.configured, - provider_type: format!("{:?}", entry.provider_type), - config_keys: entry - .config_keys - .into_iter() - .map(provider_config_key_to_dto) - .collect(), - setup_steps: entry.setup_steps, - supports_refresh: entry.supports_refresh, - refreshing: entry.refreshing, - models: entry - .models - .into_iter() - .map(|m| ProviderInventoryModelDto { - id: m.id, - name: m.name, - family: m.family, - context_limit: m.context_limit, - reasoning: m.reasoning, - recommended: m.recommended, - }) - .collect(), - last_updated_at: entry.last_updated_at.map(|t| t.to_rfc3339()), - last_refresh_attempt_at: entry.last_refresh_attempt_at.map(|t| t.to_rfc3339()), - last_refresh_error: entry.last_refresh_error, - stale, - model_selection_hint: entry.model_selection_hint, - } -} - -fn provider_config_key_to_dto(key: crate::providers::base::ConfigKey) -> ProviderConfigKey { - ProviderConfigKey { - name: key.name, - required: key.required, - secret: key.secret, - default: key.default, - oauth_flow: key.oauth_flow, - device_code_flow: key.device_code_flow, - primary: key.primary, - } -} - -const SECRET_MASK_PREFIX_LEN: usize = 4; -const SECRET_MASK_SUFFIX_LEN: usize = 3; -const SECRET_MASK_FALLBACK: &str = "***"; - -fn mask_secret_value(value: &str) -> String { - let prefix: String = value.chars().take(SECRET_MASK_PREFIX_LEN).collect(); - let suffix_chars: Vec = value.chars().rev().take(SECRET_MASK_SUFFIX_LEN).collect(); - let suffix: String = suffix_chars.into_iter().rev().collect(); - - if prefix.is_empty() - || suffix.is_empty() - || value.chars().count() <= SECRET_MASK_PREFIX_LEN + SECRET_MASK_SUFFIX_LEN - { - return SECRET_MASK_FALLBACK.to_string(); - } - - format!("{prefix}...{suffix}") -} - -fn config_value_to_string(value: &serde_json::Value) -> Option { - match value { - serde_json::Value::Null => None, - serde_json::Value::String(value) if value.is_empty() => None, - serde_json::Value::String(value) => Some(value.clone()), - other => serde_json::to_string(other).ok(), - } -} - -fn provider_config_field_value( - config: &Config, - key: &crate::providers::base::ConfigKey, - secrets: Option<&HashMap>, -) -> ProviderConfigFieldValueDto { - let value = if key.secret { - std::env::var(key.name.to_uppercase()).ok().or_else(|| { - secrets - .and_then(|values| values.get(&key.name)) - .and_then(config_value_to_string) - }) - } else { - config - .get_param::(&key.name) - .ok() - .and_then(|value| config_value_to_string(&value)) - }; - - ProviderConfigFieldValueDto { - key: key.name.clone(), - value: value.as_deref().map(|value| { - if key.secret { - mask_secret_value(value) - } else { - value.to_string() - } - }), - is_set: value.is_some(), - is_secret: key.secret, - required: key.required, - } -} - -fn refresh_skip_reason_to_dto(reason: RefreshSkipReason) -> RefreshProviderInventorySkipReasonDto { - match reason { - RefreshSkipReason::UnknownProvider => { - RefreshProviderInventorySkipReasonDto::UnknownProvider - } - RefreshSkipReason::NotConfigured => RefreshProviderInventorySkipReasonDto::NotConfigured, - RefreshSkipReason::DoesNotSupportRefresh => { - RefreshProviderInventorySkipReasonDto::DoesNotSupportRefresh - } - RefreshSkipReason::AlreadyRefreshing => { - RefreshProviderInventorySkipReasonDto::AlreadyRefreshing - } - } -} - -fn refresh_plan_to_response(refresh_plan: RefreshPlan) -> RefreshProviderInventoryResponse { - RefreshProviderInventoryResponse { - started: refresh_plan.started, - skipped: refresh_plan - .skipped - .into_iter() - .map(|entry| RefreshProviderInventorySkipDto { - provider_id: entry.provider_id, - reason: refresh_skip_reason_to_dto(entry.reason), - }) - .collect(), - } -} - fn build_model_state(current_model: &str, inventory: &ProviderInventoryEntry) -> SessionModelState { let mut available_models = inventory .models @@ -2925,1577 +2784,10 @@ impl GooseAcpAgent { } } -#[custom_methods] -impl GooseAcpAgent { - #[custom_method(AddExtensionRequest)] - async fn on_add_extension( - &self, - req: AddExtensionRequest, - ) -> Result { - let internal_id = self.internal_session_id(&req.session_id).await?; - let config: ExtensionConfig = serde_json::from_value(req.config) - .map_err(|e| sacp::Error::invalid_params().data(format!("bad config: {e}")))?; - let agent = self.get_session_agent(&req.session_id, None).await?; - agent - .add_extension(config, &internal_id) - .await - .internal_err()?; - Ok(EmptyResponse {}) - } - - #[custom_method(RemoveExtensionRequest)] - async fn on_remove_extension( - &self, - req: RemoveExtensionRequest, - ) -> Result { - let internal_id = self.internal_session_id(&req.session_id).await?; - let agent = self.get_session_agent(&req.session_id, None).await?; - agent - .remove_extension(&req.name, &internal_id) - .await - .internal_err()?; - Ok(EmptyResponse {}) - } - - #[custom_method(GetToolsRequest)] - async fn on_get_tools(&self, req: GetToolsRequest) -> Result { - let internal_id = self.internal_session_id(&req.session_id).await?; - let agent = self.get_session_agent(&req.session_id, None).await?; - let tools = agent.list_tools(&internal_id, None).await; - let tools_json = tools - .into_iter() - .map(|t| serde_json::to_value(&t)) - .collect::, _>>() - .internal_err()?; - Ok(GetToolsResponse { tools: tools_json }) - } - - #[custom_method(ReadResourceRequest)] - async fn on_read_resource( - &self, - req: ReadResourceRequest, - ) -> Result { - let internal_id = self.internal_session_id(&req.session_id).await?; - let agent = self.get_session_agent(&req.session_id, None).await?; - let cancel_token = CancellationToken::new(); - let result = agent - .extension_manager - .read_resource(&internal_id, &req.uri, &req.extension_name, cancel_token) - .await - .internal_err()?; - let result_json = serde_json::to_value(&result).internal_err()?; - Ok(ReadResourceResponse { - result: result_json, - }) - } - - #[custom_method(UpdateWorkingDirRequest)] - async fn on_update_working_dir( - &self, - req: UpdateWorkingDirRequest, - ) -> Result { - let working_dir = req.working_dir.trim().to_string(); - if working_dir.is_empty() { - return Err(sacp::Error::invalid_params().data("working directory cannot be empty")); - } - let path = std::path::PathBuf::from(&working_dir); - if !path.exists() || !path.is_dir() { - return Err(sacp::Error::invalid_params().data("invalid directory path")); - } - let internal_id = self.internal_session_id(&req.session_id).await?; - self.session_manager - .update(&internal_id) - .working_dir(path.clone()) - .apply() - .await - .internal_err()?; - - self.thread_manager - .update_working_dir(&req.session_id, &working_dir) - .await - .internal_err()?; - - if let Some(session) = self.sessions.lock().await.get_mut(&req.session_id) { - match &session.agent { - AgentHandle::Ready(agent) => { - agent.extension_manager.update_working_dir(&path).await; - } - AgentHandle::Loading(_) => { - session.pending_working_dir = Some(path); - } - } - } - - Ok(EmptyResponse {}) - } - - #[custom_method(DeleteSessionRequest)] - async fn on_delete_session( - &self, - req: DeleteSessionRequest, - ) -> Result { - // Delete the thread and all its internal sessions + messages. - self.thread_manager - .delete_thread(&req.session_id) - .await - .internal_err()?; - self.sessions.lock().await.remove(&req.session_id); - Ok(EmptyResponse {}) - } - - #[custom_method(GetExtensionsRequest)] - async fn on_get_extensions(&self) -> Result { - let extensions = crate::config::extensions::get_all_extensions(); - let warnings = crate::config::extensions::get_warnings(); - let extensions_json = extensions - .into_iter() - .map(|e| { - let config_key = e.config.key(); - let mut value = serde_json::to_value(&e)?; - if let Some(obj) = value.as_object_mut() { - obj.insert( - "config_key".to_string(), - serde_json::Value::String(config_key), - ); - } - Ok::<_, serde_json::Error>(value) - }) - .collect::, _>>() - .internal_err()?; - Ok(GetExtensionsResponse { - extensions: extensions_json, - warnings, - }) - } - - #[custom_method(AddConfigExtensionRequest)] - async fn on_add_config_extension( - &self, - req: AddConfigExtensionRequest, - ) -> Result { - let mut obj = match req.extension_config { - serde_json::Value::Object(obj) => obj, - _ => { - return Err( - sacp::Error::invalid_params().data("extensionConfig must be a JSON object") - ); - } - }; - obj.insert( - "name".to_string(), - serde_json::Value::String(req.name.clone()), - ); - - let config: crate::agents::ExtensionConfig = - serde_json::from_value(serde_json::Value::Object(obj)) - .map_err(|e| sacp::Error::invalid_params().data(format!("bad config: {e}")))?; - - crate::config::extensions::set_extension(crate::config::extensions::ExtensionEntry { - enabled: req.enabled, - config, - }); - Ok(EmptyResponse {}) - } - - #[custom_method(RemoveConfigExtensionRequest)] - async fn on_remove_config_extension( - &self, - req: RemoveConfigExtensionRequest, - ) -> Result { - let keys = crate::config::extensions::get_all_extension_names(); - if !keys.iter().any(|k| k == &req.config_key) { - return Err(sacp::Error::invalid_params() - .data(format!("Extension '{}' not found", req.config_key))); - } - crate::config::extensions::remove_extension(&req.config_key); - Ok(EmptyResponse {}) - } - - #[custom_method(ToggleConfigExtensionRequest)] - async fn on_toggle_config_extension( - &self, - req: ToggleConfigExtensionRequest, - ) -> Result { - let keys = crate::config::extensions::get_all_extension_names(); - if !keys.iter().any(|k| k == &req.config_key) { - return Err(sacp::Error::invalid_params() - .data(format!("Extension '{}' not found", req.config_key))); - } - crate::config::extensions::set_extension_enabled(&req.config_key, req.enabled); - Ok(EmptyResponse {}) - } - - #[custom_method(GetSessionExtensionsRequest)] - async fn on_get_session_extensions( - &self, - req: GetSessionExtensionsRequest, - ) -> Result { - let internal_id = self.internal_session_id(&req.session_id).await?; - let session = self - .session_manager - .get_session(&internal_id, false) - .await - .internal_err()?; - - let extensions = EnabledExtensionsState::extensions_or_default( - Some(&session.extension_data), - crate::config::Config::global(), - ); - - let extensions_json = extensions - .into_iter() - .map(|e| serde_json::to_value(&e)) - .collect::, _>>() - .internal_err()?; - - Ok(GetSessionExtensionsResponse { - extensions: extensions_json, - }) - } - - #[custom_method(ListProvidersRequest)] - async fn on_list_providers( - &self, - req: ListProvidersRequest, - ) -> Result { - let entries = self - .provider_inventory - .entries(&req.provider_ids) - .await - .internal_err()?; - Ok(ListProvidersResponse { - entries: entries.into_iter().map(inventory_entry_to_dto).collect(), - }) - } - - async fn provider_config_status(provider_id: String) -> ProviderConfigStatusDto { - let is_configured = match crate::providers::get_from_registry(&provider_id).await { - Ok(entry) => { - match tokio::task::spawn_blocking(move || entry.inventory_configured()).await { - Ok(is_configured) => is_configured, - Err(error) => { - warn!( - provider = %provider_id, - error = %error, - "provider config status check failed" - ); - false - } - } - } - Err(_) => false, - }; - - ProviderConfigStatusDto { - provider_id, - is_configured, - } - } - - async fn provider_config_statuses(provider_ids: &[String]) -> Vec { - let mut ids = if provider_ids.is_empty() { - crate::providers::providers() - .await - .into_iter() - .map(|(metadata, _)| metadata.name) - .collect::>() - } else { - provider_ids.to_vec() - }; - ids.sort(); - ids.dedup(); - - let mut statuses = stream::iter(ids) - .map(Self::provider_config_status) - .buffer_unordered(PROVIDER_CONFIG_STATUS_CHECK_CONCURRENCY) - .collect::>() - .await; - statuses.sort_by(|a, b| a.provider_id.cmp(&b.provider_id)); - statuses - } - - fn spawn_provider_inventory_refresh_jobs(&self, refresh_plan: &RefreshJobPlan) { - for refresh_job in refresh_plan.started.iter().cloned() { - let provider_inventory = self.provider_inventory.clone(); - let provider_factory = Arc::clone(&self.provider_factory); - let provider_id = refresh_job.provider_id.clone(); - let identity = refresh_job.identity.clone(); - tokio::spawn(async move { - let mut refresh_guard = provider_inventory.refresh_guard(&identity); - let provider_result = AssertUnwindSafe(async { - let metadata = crate::providers::get_from_registry(&provider_id).await?; - let model_config = - crate::model::ModelConfig::new(&metadata.metadata().default_model)? - .with_canonical_limits(&provider_id); - provider_factory(provider_id.clone(), model_config, Vec::new()).await - }) - .catch_unwind() - .await; - - let fetch_result: Result> = match provider_result { - Ok(Ok(provider)) => { - match ensure_refresh_identity_current(&provider_id, &identity).await { - Ok(()) => match AssertUnwindSafe(provider.fetch_recommended_models()) - .catch_unwind() - .await - { - Ok(Ok(models)) => Ok(models), - Ok(Err(error)) => Err(anyhow::anyhow!(error.to_string())), - Err(_) => { - Err(anyhow::anyhow!("provider inventory refresh task panicked")) - } - }, - Err(error) => Err(error), - } - } - Ok(Err(error)) => Err(error), - Err(_) => Err(anyhow::anyhow!("provider inventory refresh task panicked")), - }; - - match fetch_result { - Ok(models) => match provider_inventory - .store_refreshed_models_for_identity(&identity, &models) - .await - { - Ok(()) => refresh_guard.complete(), - Err(error) => warn!( - provider = %provider_id, - error = %error, - "failed to store refreshed provider inventory" - ), - }, - Err(error) => { - let error_message = error.to_string(); - match provider_inventory - .store_refresh_error_for_identity(&identity, error_message.clone()) - .await - { - Ok(()) => refresh_guard.complete(), - Err(store_error) => warn!( - provider = %provider_id, - error = %store_error, - refresh_error = %error_message, - "failed to store provider inventory refresh error" - ), - } - warn!(provider = %provider_id, error = %error_message, "provider inventory refresh failed"); - } - } - }); - } - } - - async fn start_provider_inventory_refresh( - &self, - provider_ids: &[String], - ) -> Result { - let refresh_job_plan = self - .provider_inventory - .plan_refresh_jobs(provider_ids) - .await - .internal_err()?; - self.spawn_provider_inventory_refresh_jobs(&refresh_job_plan); - Ok(refresh_plan_to_response( - refresh_job_plan.into_public_plan(), - )) - } - - #[custom_method(RefreshProviderInventoryRequest)] - async fn on_refresh_provider_inventory( - &self, - req: RefreshProviderInventoryRequest, - ) -> Result { - Config::global().invalidate_secrets_cache(); - self.start_provider_inventory_refresh(&req.provider_ids) - .await - } - - #[custom_method(ProviderConfigReadRequest)] - async fn on_read_provider_config( - &self, - req: ProviderConfigReadRequest, - ) -> Result { - let entry = crate::providers::get_from_registry(&req.provider_id) - .await - .invalid_params_err_ctx("Unknown provider")?; - let config = Config::global(); - let config_keys = &entry.metadata().config_keys; - let secrets = if config_keys.iter().any(|key| key.secret) { - Some(config.all_secrets().internal_err()?) - } else { - None - }; - - Ok(ProviderConfigReadResponse { - fields: config_keys - .iter() - .map(|key| provider_config_field_value(config, key, secrets.as_ref())) - .collect(), - }) - } - - #[custom_method(ProviderConfigStatusRequest)] - async fn on_provider_config_status( - &self, - req: ProviderConfigStatusRequest, - ) -> Result { - Ok(ProviderConfigStatusResponse { - statuses: Self::provider_config_statuses(&req.provider_ids).await, - }) - } - - #[custom_method(ProviderConfigSaveRequest)] - async fn on_save_provider_config( - &self, - req: ProviderConfigSaveRequest, - ) -> Result { - let entry = crate::providers::get_from_registry(&req.provider_id) - .await - .invalid_params_err_ctx("Unknown provider")?; - let metadata = entry.metadata().clone(); - let config = Config::global(); - let mut config_updates = Vec::new(); - let mut secret_updates = Vec::new(); - - for field in &req.fields { - let Some(config_key) = metadata - .config_keys - .iter() - .find(|config_key| config_key.name == field.key) - else { - return Err(sacp::Error::invalid_params() - .data(format!("Unsupported provider config field: {}", field.key))); - }; - - let value = field.value.trim(); - if value.is_empty() { - return Err(sacp::Error::invalid_params().data(format!( - "Provider config field cannot be empty: {}", - field.key - ))); - } - - if config_key.secret { - secret_updates.push(( - config_key.name.clone(), - serde_json::Value::String(value.to_string()), - )); - } else { - config_updates.push((config_key.name.clone(), value.to_string())); - } - } - - for (key, value) in config_updates { - config - .set_param(&key, &value) - .internal_err_ctx("Failed to save provider config field")?; - } - config - .set_secret_values(&secret_updates) - .internal_err_ctx("Failed to save provider secret fields")?; - - let provider_ids = [req.provider_id.clone()]; - let status = Self::provider_config_status(req.provider_id.clone()).await; - let refresh = self.start_provider_inventory_refresh(&provider_ids).await?; - Ok(ProviderConfigChangeResponse { status, refresh }) - } - - #[custom_method(ProviderConfigDeleteRequest)] - async fn on_delete_provider_config( - &self, - req: ProviderConfigDeleteRequest, - ) -> Result { - let entry = crate::providers::get_from_registry(&req.provider_id) - .await - .invalid_params_err_ctx("Unknown provider")?; - let metadata = entry.metadata().clone(); - let config = Config::global(); - let mut secret_keys = Vec::new(); - - for config_key in &metadata.config_keys { - if config_key.secret { - secret_keys.push(config_key.name.clone()); - } else { - config - .delete(&config_key.name) - .internal_err_ctx("Failed to delete provider config field")?; - } - } - - config - .delete_secret_values(&secret_keys) - .internal_err_ctx("Failed to delete provider secret fields")?; - crate::providers::cleanup_provider(&req.provider_id) - .await - .internal_err_ctx("Failed to clean up provider state")?; - - let provider_ids = [req.provider_id.clone()]; - let status = Self::provider_config_status(req.provider_id.clone()).await; - let refresh = self.start_provider_inventory_refresh(&provider_ids).await?; - Ok(ProviderConfigChangeResponse { status, refresh }) - } - - #[custom_method(ReadConfigRequest)] - async fn on_read_config( - &self, - req: ReadConfigRequest, - ) -> Result { - let config = self.config()?; - let response = match config.get_param::(&req.key) { - Ok(value) => ReadConfigResponse { value }, - Err(crate::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.config()?; - config.set_param(&req.key, &req.value).internal_err()?; - Ok(EmptyResponse {}) - } - - #[custom_method(RemoveConfigRequest)] - async fn on_remove_config( - &self, - req: RemoveConfigRequest, - ) -> Result { - let config = self.config()?; - config.delete(&req.key).internal_err()?; - Ok(EmptyResponse {}) - } - - #[custom_method(CheckSecretRequest)] - async fn on_check_secret( - &self, - req: CheckSecretRequest, - ) -> Result { - let config = self.config()?; - 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.config()?; - config.set_secret(&req.key, &req.value).internal_err()?; - Config::global().invalidate_secrets_cache(); - Ok(EmptyResponse {}) - } - - #[custom_method(RemoveSecretRequest)] - async fn on_remove_secret( - &self, - req: RemoveSecretRequest, - ) -> Result { - let config = self.config()?; - config.delete_secret(&req.key).internal_err()?; - Config::global().invalidate_secrets_cache(); - Ok(EmptyResponse {}) - } - - #[custom_method(ExportSessionRequest)] - async fn on_export_session( - &self, - req: ExportSessionRequest, - ) -> Result { - let thread = self - .thread_manager - .get_thread(&req.session_id) - .await - .internal_err()?; - let internal_id = thread - .current_session_id - .ok_or_else(|| sacp::Error::internal_error().data("Thread has no internal session"))?; - let data = self - .session_manager - .export_session(&internal_id) - .await - .internal_err()?; - Ok(ExportSessionResponse { data }) - } - - #[custom_method(ImportSessionRequest)] - async fn on_import_session( - &self, - req: ImportSessionRequest, - ) -> Result { - let session = self - .session_manager - .import_session(&req.data, Some(SessionType::Acp)) - .await - .internal_err()?; - - // Create a thread for the imported session. - let thread = self - .thread_manager - .create_thread( - Some(session.name.clone()), - None, - Some(session.working_dir.display().to_string()), - ) - .await - .internal_err()?; - - // Link the internal session to the thread. - self.session_manager - .update(&session.id) - .thread_id(Some(thread.id.clone())) - .apply() - .await - .internal_err()?; - - // Copy conversation messages into thread_messages so they appear in the thread. - if let Some(ref conversation) = session.conversation { - for msg in conversation.messages() { - self.thread_manager - .append_message(&thread.id, Some(&session.id), msg) - .await - .internal_err()?; - } - } - - // Re-fetch thread to get accurate message_count. - let thread = self - .thread_manager - .get_thread(&thread.id) - .await - .internal_err()?; - - Ok(ImportSessionResponse { - session_id: thread.id, - title: Some(thread.name), - updated_at: Some(thread.updated_at.to_rfc3339()), - message_count: thread.message_count as u64, - }) - } - - #[custom_method(UpdateSessionProjectRequest)] - async fn on_update_session_project( - &self, - req: UpdateSessionProjectRequest, - ) -> Result { - let project_id = req.project_id; - self.update_thread_metadata(&req.session_id, move |meta| { - meta.project_id = project_id; - }) - .await?; - Ok(EmptyResponse {}) - } - - #[custom_method(RenameSessionRequest)] - async fn on_rename_session( - &self, - req: RenameSessionRequest, - ) -> Result { - self.thread_manager - .update_thread(&req.session_id, Some(req.title), Some(true), None) - .await - .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; - Ok(EmptyResponse {}) - } - - #[custom_method(ArchiveSessionRequest)] - async fn on_archive_session( - &self, - req: ArchiveSessionRequest, - ) -> Result { - self.thread_manager - .archive_thread(&req.session_id) - .await - .internal_err()?; - self.sessions.lock().await.remove(&req.session_id); - Ok(EmptyResponse {}) - } - - #[custom_method(UnarchiveSessionRequest)] - async fn on_unarchive_session( - &self, - req: UnarchiveSessionRequest, - ) -> Result { - self.thread_manager - .unarchive_thread(&req.session_id) - .await - .internal_err()?; - Ok(EmptyResponse {}) - } - - #[custom_method(CreateSourceRequest)] - async fn on_create_source( - &self, - req: CreateSourceRequest, - ) -> Result { - let source = crate::sources::create_source( - req.source_type, - &req.name, - &req.description, - &req.content, - req.global, - req.project_dir.as_deref(), - )?; - Ok(CreateSourceResponse { source }) - } - - #[custom_method(ListSourcesRequest)] - async fn on_list_sources( - &self, - req: ListSourcesRequest, - ) -> Result { - let sources = crate::sources::list_sources(req.source_type, req.project_dir.as_deref())?; - Ok(ListSourcesResponse { sources }) - } - - #[custom_method(UpdateSourceRequest)] - async fn on_update_source( - &self, - req: UpdateSourceRequest, - ) -> Result { - let source = crate::sources::update_source( - req.source_type, - &req.path, - &req.name, - &req.description, - &req.content, - )?; - Ok(UpdateSourceResponse { source }) - } - - #[custom_method(DeleteSourceRequest)] - async fn on_delete_source( - &self, - req: DeleteSourceRequest, - ) -> Result { - crate::sources::delete_source(req.source_type, &req.path)?; - Ok(EmptyResponse {}) - } - - #[custom_method(ExportSourceRequest)] - async fn on_export_source( - &self, - req: ExportSourceRequest, - ) -> Result { - let (json, filename) = crate::sources::export_source(req.source_type, &req.path)?; - Ok(ExportSourceResponse { json, filename }) - } - - #[custom_method(ImportSourcesRequest)] - async fn on_import_sources( - &self, - req: ImportSourcesRequest, - ) -> Result { - let sources = - crate::sources::import_sources(&req.data, req.global, req.project_dir.as_deref())?; - Ok(ImportSourcesResponse { sources }) - } - - #[custom_method(DictationTranscribeRequest)] - async fn on_dictation_transcribe( - &self, - req: DictationTranscribeRequest, - ) -> Result { - use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; - let config = crate::config::Config::global(); - - #[cfg(not(feature = "local-inference"))] - if req.provider == "local" { - return Err(sacp::Error::invalid_params() - .data("Local inference is not available in this build")); - } - - let provider: DictationProvider = serde_json::from_value(serde_json::Value::String( - req.provider.clone(), - )) - .map_err(|_| { - sacp::Error::invalid_params().data(format!("Unknown provider: {}", req.provider)) - })?; - - let audio_bytes = BASE64 - .decode(&req.audio) - .map_err(|_| sacp::Error::invalid_params().data("Invalid base64 audio data"))?; - - if audio_bytes.len() > 50 * 1024 * 1024 { - return Err(sacp::Error::invalid_params().data("Audio too large (max 50MB)")); - } - - let extension = match req.mime_type.as_str() { - "audio/webm" | "audio/webm;codecs=opus" => "webm", - "audio/mp4" => "mp4", - "audio/mpeg" | "audio/mpga" => "mp3", - "audio/m4a" => "m4a", - "audio/wav" | "audio/x-wav" => "wav", - other => { - return Err( - sacp::Error::invalid_params().data(format!("Unsupported format: {other}")) - ); - } - }; - - let text = match provider { - #[cfg(feature = "local-inference")] - DictationProvider::Local => transcribe_local(audio_bytes).await, - remote => { - let (model_param, default_model) = dictation_transcribe_params(remote); - let model = dictation_selected_model(config, remote) - .unwrap_or_else(|| default_model.to_string()); - transcribe_with_provider( - remote, - model_param.to_string(), - model, - audio_bytes, - extension, - &req.mime_type, - ) - .await - } - } - .internal_err()?; - - Ok(DictationTranscribeResponse { text }) - } - - #[custom_method(DictationConfigRequest)] - async fn on_dictation_config( - &self, - _req: DictationConfigRequest, - ) -> Result { - let config = crate::config::Config::global(); - let mut providers = std::collections::HashMap::new(); - - for def in all_providers() { - let provider = def.provider; - let host = if let Some(host_key) = def.host_key { - config - .get(host_key, false) - .ok() - .and_then(|v| v.as_str().map(|s| s.to_string())) - } else { - None - }; - - let provider_key = serde_json::to_value(provider) - .ok() - .and_then(|v| v.as_str().map(|s| s.to_string())) - .unwrap_or_else(|| format!("{:?}", provider).to_lowercase()); - providers.insert( - provider_key, - DictationProviderStatusEntry { - configured: is_configured(provider), - host, - description: def.description.to_string(), - uses_provider_config: def.uses_provider_config, - settings_path: def.settings_path.map(|s| s.to_string()), - config_key: if !def.uses_provider_config { - Some(def.config_key.to_string()) - } else { - None - }, - model_config_key: dictation_model_config_key(provider), - default_model: dictation_default_model(provider), - selected_model: dictation_selected_model(config, provider), - available_models: dictation_available_models(provider), - }, - ); - } - - Ok(DictationConfigResponse { providers }) - } - - #[custom_method(DictationModelsListRequest)] - async fn on_dictation_models_list( - &self, - _req: DictationModelsListRequest, - ) -> Result { - #[cfg(feature = "local-inference")] - { - use crate::download_manager::{get_download_manager, DownloadStatus}; - - let manager = get_download_manager(); - let models = whisper::available_models() - .iter() - .map(|model| DictationLocalModelStatus { - id: model.id.to_string(), - label: model.id.to_string(), - description: model.description.to_string(), - size_mb: model.size_mb, - downloaded: model.is_downloaded(), - download_in_progress: manager - .get_progress(model.id) - .map(|progress| progress.status == DownloadStatus::Downloading) - .unwrap_or(false), - }) - .collect(); - - Ok(DictationModelsListResponse { models }) - } - - #[cfg(not(feature = "local-inference"))] - Ok(DictationModelsListResponse::default()) - } - - #[custom_method(DictationModelDownloadRequest)] - async fn on_dictation_model_download( - &self, - _req: DictationModelDownloadRequest, - ) -> Result { - #[cfg(feature = "local-inference")] - { - use crate::download_manager::get_download_manager; - - let model = whisper::get_model(&_req.model_id) - .ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?; - let manager = get_download_manager(); - let model_id_for_config = model.id.to_string(); - - manager - .download_model( - model.id.to_string(), - model.url.to_string(), - model.local_path(), - Some(Box::new(move || { - let config = crate::config::Config::global(); - // Only auto-select this model if the user has no model - // currently selected. This prevents silently switching - // the active model mid-session when a user downloads an - // additional model while one is already in use. - let already_selected = config - .get(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, false) - .ok() - .and_then(|value| value.as_str().map(str::to_owned)) - .filter(|model_id| { - // Treat a deleted model file as no active selection - // so a fresh download can auto-select cleanly. - whisper::get_model(model_id) - .is_some_and(|model| model.is_downloaded()) - }); - if already_selected.is_none() { - if let Err(e) = config.set_param( - whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, - model_id_for_config.clone(), - ) { - error!("Failed to save LOCAL_WHISPER_MODEL after download: {}", e); - } - } - })), - ) - .await - .internal_err()?; - - Ok(EmptyResponse {}) - } - - #[cfg(not(feature = "local-inference"))] - Err(sacp::Error::invalid_params().data("Local inference not enabled")) - } - - #[custom_method(DictationModelDownloadProgressRequest)] - async fn on_dictation_model_download_progress( - &self, - _req: DictationModelDownloadProgressRequest, - ) -> Result { - #[cfg(feature = "local-inference")] - { - use crate::download_manager::get_download_manager; - - let manager = get_download_manager(); - let progress = - manager - .get_progress(&_req.model_id) - .map(|progress| DictationDownloadProgress { - bytes_downloaded: progress.bytes_downloaded, - total_bytes: progress.total_bytes, - progress_percent: progress.progress_percent, - status: serde_json::to_value(&progress.status) - .ok() - .and_then(|value| value.as_str().map(ToOwned::to_owned)) - .unwrap_or_else(|| "unknown".to_string()), - error: progress.error, - }); - - Ok(DictationModelDownloadProgressResponse { progress }) - } - - #[cfg(not(feature = "local-inference"))] - Ok(DictationModelDownloadProgressResponse { progress: None }) - } - - #[custom_method(DictationModelCancelRequest)] - async fn on_dictation_model_cancel( - &self, - _req: DictationModelCancelRequest, - ) -> Result { - #[cfg(feature = "local-inference")] - { - use crate::download_manager::get_download_manager; - - let manager = get_download_manager(); - manager.cancel_download(&_req.model_id).internal_err()?; - - Ok(EmptyResponse {}) - } - - #[cfg(not(feature = "local-inference"))] - Err(sacp::Error::invalid_params().data("Local inference not enabled")) - } - - #[custom_method(DictationModelDeleteRequest)] - async fn on_dictation_model_delete( - &self, - _req: DictationModelDeleteRequest, - ) -> Result { - #[cfg(feature = "local-inference")] - { - let model = whisper::get_model(&_req.model_id) - .ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?; - let path = model.local_path(); - - if !path.exists() { - return Err(sacp::Error::invalid_params().data("Model not downloaded")); - } - - std::fs::remove_file(path).internal_err()?; - - Ok(EmptyResponse {}) - } - - #[cfg(not(feature = "local-inference"))] - Err(sacp::Error::invalid_params().data("Local inference not enabled")) - } - - #[custom_method(DictationModelSelectRequest)] - async fn on_dictation_model_select( - &self, - req: DictationModelSelectRequest, - ) -> Result { - #[cfg(not(feature = "local-inference"))] - if req.provider == "local" { - return Err(sacp::Error::invalid_params().data("Local inference not enabled")); - } - - let provider: DictationProvider = serde_json::from_value(serde_json::Value::String( - req.provider.clone(), - )) - .map_err(|_| { - sacp::Error::invalid_params().data(format!("Unknown provider: {}", req.provider)) - })?; - - let key = match provider { - DictationProvider::OpenAI => OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY, - DictationProvider::Groq => GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY, - DictationProvider::ElevenLabs => ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY, - #[cfg(feature = "local-inference")] - DictationProvider::Local => { - let model = whisper::get_model(&req.model_id) - .ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?; - if !model.is_downloaded() { - return Err( - sacp::Error::invalid_params().data("Local Whisper model is not downloaded") - ); - } - whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY - } - }; - - crate::config::Config::global() - .set_param(key, req.model_id) - .internal_err()?; - - Ok(EmptyResponse {}) - } -} - -fn dictation_model_config_key(provider: DictationProvider) -> Option { - match provider { - DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()), - DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()), - DictationProvider::ElevenLabs => { - Some(ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()) - } - #[cfg(feature = "local-inference")] - DictationProvider::Local => Some(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY.to_string()), - } -} - -/// Returns the (param_name, default_model) pair used by `transcribe_with_provider` -/// for remote dictation providers. Local inference is handled separately. -fn dictation_transcribe_params(provider: DictationProvider) -> (&'static str, &'static str) { - match provider { - DictationProvider::OpenAI => ("model", OPENAI_TRANSCRIPTION_MODEL), - DictationProvider::Groq => ("model", GROQ_TRANSCRIPTION_MODEL), - DictationProvider::ElevenLabs => ("model_id", ELEVENLABS_TRANSCRIPTION_MODEL), - #[cfg(feature = "local-inference")] - DictationProvider::Local => ("", ""), - } -} - -fn dictation_default_model(provider: DictationProvider) -> Option { - match provider { - DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL.to_string()), - DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL.to_string()), - DictationProvider::ElevenLabs => Some(ELEVENLABS_TRANSCRIPTION_MODEL.to_string()), - #[cfg(feature = "local-inference")] - DictationProvider::Local => Some(whisper::recommend_model().to_string()), - } -} - -fn dictation_selected_model(config: &Config, provider: DictationProvider) -> Option { - #[cfg(feature = "local-inference")] - if provider == DictationProvider::Local { - return config - .get(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, false) - .ok() - .and_then(|value| value.as_str().map(str::to_owned)) - .filter(|model_id| whisper::get_model(model_id).is_some()) - .or_else(|| dictation_default_model(provider)); - } - - dictation_model_config_key(provider) - .and_then(|key| { - config - .get(&key, false) - .ok() - .and_then(|value| value.as_str().map(str::to_owned)) - }) - .or_else(|| dictation_default_model(provider)) -} - -fn dictation_available_models(provider: DictationProvider) -> Vec { - match provider { - DictationProvider::OpenAI => vec![DictationModelOption { - id: OPENAI_TRANSCRIPTION_MODEL.to_string(), - label: "Whisper-1".to_string(), - description: "OpenAI's hosted Whisper transcription model.".to_string(), - }], - DictationProvider::Groq => vec![DictationModelOption { - id: GROQ_TRANSCRIPTION_MODEL.to_string(), - label: "Whisper Large V3 Turbo".to_string(), - description: "Groq's fast hosted Whisper transcription model.".to_string(), - }], - DictationProvider::ElevenLabs => vec![DictationModelOption { - id: ELEVENLABS_TRANSCRIPTION_MODEL.to_string(), - label: "Scribe v1".to_string(), - description: "ElevenLabs' hosted speech-to-text model.".to_string(), - }], - #[cfg(feature = "local-inference")] - DictationProvider::Local => whisper::available_models() - .iter() - .map(|model| DictationModelOption { - id: model.id.to_string(), - label: model.id.to_string(), - description: model.description.to_string(), - }) - .collect(), - } -} - pub struct GooseAcpHandler { pub agent: Arc, } -impl HandleDispatchFrom for GooseAcpHandler { - fn describe_chain(&self) -> impl std::fmt::Debug { - "goose-acp" - } - - fn handle_dispatch_from( - &mut self, - message: Dispatch, - cx: ConnectionTo, - ) -> impl std::future::Future, sacp::Error>> + Send { - let agent = self.agent.clone(); - - // The MatchDispatchFrom chain produces an ~85KB async state machine. - // Box::pin moves it to the heap so it doesn't overflow the tokio worker stack. - Box::pin(async move { - // InitializeRequest runs inline: it sets connection-scoped state - // (client fs/terminal capabilities) that later handlers read with - // defaults, so a pipelined NewSessionRequest must not race ahead of it. - MatchDispatchFrom::new(message, &cx) - .if_request( - |req: InitializeRequest, responder: Responder| async { - responder.respond_with_result(agent.on_initialize(req).await) - }, - ) - .await - .if_request( - |_req: AuthenticateRequest, responder: Responder| async { - responder.respond(AuthenticateResponse::new()) - }, - ) - .await - .if_request( - |req: NewSessionRequest, responder: Responder| async { - let agent = agent.clone(); - let cx_clone = cx.clone(); - cx.spawn(async move { - responder.respond_with_result(agent.on_new_session(&cx_clone, req).await)?; - Ok(()) - })?; - Ok(()) - }, - ) - .await - .if_request( - |req: LoadSessionRequest, responder: Responder| async { - let agent = agent.clone(); - let cx_clone = cx.clone(); - cx.spawn(async move { - match agent.on_load_session(&cx_clone, req).await { - Ok(response) => { - responder.respond(response)?; - } - Err(e) => { - responder.respond_with_error(e)?; - } - } - Ok(()) - })?; - Ok(()) - }, - ) - .await - .if_request( - |req: PromptRequest, responder: Responder| async { - let agent = agent.clone(); - let cx_clone = cx.clone(); - cx.spawn(async move { - match agent.on_prompt(&cx_clone, req).await { - Ok(response) => { - responder.respond(response)?; - } - Err(e) => { - responder.respond_with_error(e)?; - } - } - Ok(()) - })?; - Ok(()) - }, - ) - .await - .if_notification(|notif: CancelNotification| async { - let agent = agent.clone(); - agent.on_cancel(notif).await?; - Ok(()) - }) - .await - // set_config_option (SACP 11) and legacy set_mode/set_model; custom _goose/* in otherwise. - .if_request({ - let agent = agent.clone(); - let cx = cx.clone(); - |req: SetSessionConfigOptionRequest, responder: Responder| async move { - let cx_spawn = cx.clone(); - cx.spawn(async move { - let cx = cx_spawn; - let value_id = req.value.as_value_id() - .ok_or_else(|| sacp::Error::invalid_params().data("Expected a value ID"))? - .clone(); - let session_id = req.session_id.clone(); - let sid = sid_short(session_id.0.as_ref()); - let config_id = req.config_id.0.to_string(); - let t_handler = std::time::Instant::now(); - match config_id.as_ref() { - "provider" => { - Config::global().invalidate_secrets_cache(); - 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(_) => {} - Err(e) => { responder.respond_with_error(e)?; return Ok(()); } - } - } - "model" => { - match agent.on_set_model(&session_id.0, &value_id.0).await { - Ok(_) => {} - Err(e) => { responder.respond_with_error(e)?; return Ok(()); } - } - } - other => { - responder.respond_with_error( - sacp::Error::invalid_params().data(format!("Unsupported config option: {}", other)) - )?; - return Ok(()); - } - } - // Respond immediately using the current provider inventory snapshot. - let (notification, config_options) = agent.build_config_update(&session_id).await?; - cx.send_notification(notification)?; - responder.respond(SetSessionConfigOptionResponse::new(config_options))?; - - let maybe_refresh = if config_id == "provider" { - let provider_id = value_id.0.to_string(); - agent - .provider_inventory - .plan_refresh_jobs(std::slice::from_ref(&provider_id)) - .await - .ok() - .and_then(|plan| { - plan.started - .into_iter() - .find(|job| job.provider_id == provider_id) - }) - } else { - None - }; - if let Some(refresh_job) = maybe_refresh { - let agent_bg = agent.clone(); - let cx_bg = cx.clone(); - let session_id_bg = session_id.clone(); - tokio::spawn(async move { - let refresh_identity = refresh_job.identity; - let refresh_provider_id = refresh_job.provider_id; - let mut refresh_guard = - agent_bg.provider_inventory.refresh_guard(&refresh_identity); - let provider_result: Result> = - AssertUnwindSafe(async { - let session_agent = - agent_bg.get_session_agent(&session_id_bg.0, None).await?; - let provider = session_agent - .provider() - .await - .map_err(|e| anyhow::anyhow!(e.to_string()))?; - let provider_name = provider.get_name().to_string(); - if provider_name != refresh_provider_id { - return Err(anyhow::anyhow!( - "provider changed before inventory refresh completed" - )); - } - Ok(provider) - }) - .catch_unwind() - .await - .map_err(|_| { - anyhow::anyhow!("provider inventory refresh task panicked") - }) - .and_then(|result| result); - - let fetch_result = match provider_result { - Ok(provider) => { - match ensure_refresh_identity_current( - &refresh_provider_id, - &refresh_identity, - ) - .await - { - Ok(()) => match AssertUnwindSafe( - provider.fetch_recommended_models(), - ) - .catch_unwind() - .await - { - Ok(Ok(models)) => Ok(models), - Ok(Err(error)) => { - Err(anyhow::anyhow!(error.to_string())) - } - Err(_) => Err(anyhow::anyhow!( - "provider inventory refresh task panicked" - )), - }, - Err(error) => Err(error), - } - } - Err(error) => Err(error), - }; - - match fetch_result { - Ok(models) => match agent_bg - .provider_inventory - .store_refreshed_models_for_identity( - &refresh_identity, - &models, - ) - .await - { - Ok(()) => { - refresh_guard.complete(); - match agent_bg.build_config_update(&session_id_bg).await - { - Ok((fresh_notification, _)) => { - let _ = cx_bg - .send_notification(fresh_notification); - } - Err(error) => warn!( - provider = %refresh_provider_id, - error = %error, - "failed to build config update after provider inventory refresh" - ), - } - } - Err(error) => warn!( - provider = %refresh_provider_id, - error = %error, - "failed to store refreshed provider inventory after config change" - ), - }, - Err(error) => { - let error_message = error.to_string(); - match agent_bg - .provider_inventory - .store_refresh_error_for_identity( - &refresh_identity, - error_message.clone(), - ) - .await - { - Ok(()) => refresh_guard.complete(), - Err(store_error) => warn!( - provider = %refresh_provider_id, - error = %store_error, - refresh_error = %error_message, - "failed to store provider inventory refresh error after config change" - ), - } - warn!( - provider = %refresh_provider_id, - error = %error_message, - "provider inventory refresh failed after config change" - ); - } - } - }); - } - - debug!(target: "perf", sid = %sid, ms = t_handler.elapsed().as_millis() as u64, config_id = %config_id, "perf: set_config_option done"); - Ok(()) - })?; - Ok(()) - } - }) - .await - .if_request({ - let agent = agent.clone(); - let cx = cx.clone(); - |req: SetSessionModeRequest, responder: Responder| async move { - let cx_spawn = cx.clone(); - cx.spawn(async move { - let cx = cx_spawn; - let session_id = req.session_id.clone(); - let mode_id = req.mode_id.clone(); - match agent.on_set_mode(&session_id.0, &mode_id.0).await { - Ok(resp) => { - // Notify before responding so clients see the mode update before block_task unblocks. - cx.send_notification(SessionNotification::new( - session_id, - SessionUpdate::CurrentModeUpdate( - CurrentModeUpdate::new(mode_id), - ), - ))?; - responder.respond(resp)?; - } - Err(e) => { - responder.respond_with_error(e)?; - } - } - Ok(()) - })?; - Ok(()) - } - }) - .await - .if_request({ - let agent = agent.clone(); - let cx = cx.clone(); - |req: SetSessionModelRequest, responder: Responder| async move { - let cx_spawn = cx.clone(); - cx.spawn(async move { - let cx = cx_spawn; - let session_id = req.session_id.clone(); - match agent.on_set_model(&session_id.0, &req.model_id.0).await { - Ok(resp) => { - let (notification, _) = agent.build_config_update(&session_id).await?; - cx.send_notification(notification)?; - responder.respond(resp)?; - } - Err(e) => responder.respond_with_error(e)?, - } - Ok(()) - })?; - Ok(()) - } - }) - .await - .if_request({ - let agent = agent.clone(); - let cx = cx.clone(); - |_req: ListSessionsRequest, responder: Responder| async move { - cx.spawn(async move { - responder.respond(agent.on_list_sessions().await?)?; - Ok(()) - })?; - Ok(()) - } - }) - .await - .if_request({ - let agent = agent.clone(); - let cx = cx.clone(); - |req: CloseSessionRequest, responder: Responder| async move { - cx.spawn(async move { - responder.respond(agent.on_close_session(&req.session_id.0).await?)?; - Ok(()) - })?; - Ok(()) - } - }) - .await - .if_request({ - let agent = agent.clone(); - let cx = cx.clone(); - |req: ForkSessionRequest, responder: Responder| async move { - let cx_spawn = cx.clone(); - cx.spawn(async move { - responder.respond_with_result(agent.on_fork_session(&cx_spawn, req).await)?; - Ok(()) - })?; - Ok(()) - } - }) - .await - .otherwise({ - let agent = agent.clone(); - let cx = cx.clone(); - |message: Dispatch| async move { - match message { - Dispatch::Request(req, responder) => { - cx.spawn(async move { - match agent.handle_custom_request(&req.method, req.params).await { - Ok(json) => responder.respond(json)?, - Err(e) => responder.respond_with_error(e)?, - } - Ok(()) - })?; - Ok(()) - } - Dispatch::Response(result, router) => { - debug!(method = %router.method(), id = %router.id(), ok = result.is_ok(), "routing response"); - router.respond_with_result(result)?; - Ok(()) - } - Dispatch::Notification(notif) => { - debug!(method = %notif.method, "unhandled notification"); - Ok(()) - } - } - } - }) - .await - .map(|()| Handled::Yes) - }) - } -} - pub fn serve( agent: Arc, read: R, diff --git a/crates/goose/src/acp/server/config.rs b/crates/goose/src/acp/server/config.rs new file mode 100644 index 00000000..7610433f --- /dev/null +++ b/crates/goose/src/acp/server/config.rs @@ -0,0 +1,36 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_read_config( + &self, + req: ReadConfigRequest, + ) -> Result { + let config = self.config()?; + let response = match config.get_param::(&req.key) { + Ok(value) => ReadConfigResponse { value }, + Err(crate::config::ConfigError::NotFound(_)) => ReadConfigResponse { + value: serde_json::Value::Null, + }, + Err(e) => return Err(sacp::Error::internal_error().data(e.to_string())), + }; + Ok(response) + } + + pub(super) async fn on_upsert_config( + &self, + req: UpsertConfigRequest, + ) -> Result { + let config = self.config()?; + config.set_param(&req.key, &req.value).internal_err()?; + Ok(EmptyResponse {}) + } + + pub(super) async fn on_remove_config( + &self, + req: RemoveConfigRequest, + ) -> Result { + let config = self.config()?; + config.delete(&req.key).internal_err()?; + Ok(EmptyResponse {}) + } +} diff --git a/crates/goose/src/acp/server/custom_dispatch.rs b/crates/goose/src/acp/server/custom_dispatch.rs new file mode 100644 index 00000000..19a49843 --- /dev/null +++ b/crates/goose/src/acp/server/custom_dispatch.rs @@ -0,0 +1,354 @@ +use super::*; +use goose_acp_macros::custom_methods; + +#[custom_methods] +impl GooseAcpAgent { + pub async fn dispatch_custom_request( + &self, + method: &str, + params: serde_json::Value, + ) -> Result { + self.handle_custom_request(method, params).await + } + + #[custom_method(AddExtensionRequest)] + async fn dispatch_add_extension( + &self, + req: AddExtensionRequest, + ) -> Result { + self.on_add_extension(req).await + } + + #[custom_method(RemoveExtensionRequest)] + async fn dispatch_remove_extension( + &self, + req: RemoveExtensionRequest, + ) -> Result { + self.on_remove_extension(req).await + } + + #[custom_method(GetToolsRequest)] + async fn dispatch_get_tools( + &self, + req: GetToolsRequest, + ) -> Result { + self.on_get_tools(req).await + } + + #[custom_method(ReadResourceRequest)] + async fn dispatch_read_resource( + &self, + req: ReadResourceRequest, + ) -> Result { + self.on_read_resource(req).await + } + + #[custom_method(UpdateWorkingDirRequest)] + async fn dispatch_update_working_dir( + &self, + req: UpdateWorkingDirRequest, + ) -> Result { + self.on_update_working_dir(req).await + } + + #[custom_method(DeleteSessionRequest)] + async fn dispatch_delete_session( + &self, + req: DeleteSessionRequest, + ) -> Result { + self.on_delete_session(req).await + } + + #[custom_method(GetExtensionsRequest)] + async fn dispatch_get_extensions(&self) -> Result { + self.on_get_extensions().await + } + + #[custom_method(AddConfigExtensionRequest)] + async fn dispatch_add_config_extension( + &self, + req: AddConfigExtensionRequest, + ) -> Result { + self.on_add_config_extension(req).await + } + + #[custom_method(RemoveConfigExtensionRequest)] + async fn dispatch_remove_config_extension( + &self, + req: RemoveConfigExtensionRequest, + ) -> Result { + self.on_remove_config_extension(req).await + } + + #[custom_method(ToggleConfigExtensionRequest)] + async fn dispatch_toggle_config_extension( + &self, + req: ToggleConfigExtensionRequest, + ) -> Result { + self.on_toggle_config_extension(req).await + } + + #[custom_method(GetSessionExtensionsRequest)] + async fn dispatch_get_session_extensions( + &self, + req: GetSessionExtensionsRequest, + ) -> Result { + self.on_get_session_extensions(req).await + } + + #[custom_method(ListProvidersRequest)] + async fn dispatch_list_providers( + &self, + req: ListProvidersRequest, + ) -> Result { + self.on_list_providers(req).await + } + + #[custom_method(RefreshProviderInventoryRequest)] + async fn dispatch_refresh_provider_inventory( + &self, + req: RefreshProviderInventoryRequest, + ) -> Result { + self.on_refresh_provider_inventory(req).await + } + + #[custom_method(ProviderConfigReadRequest)] + async fn dispatch_read_provider_config( + &self, + req: ProviderConfigReadRequest, + ) -> Result { + self.on_read_provider_config(req).await + } + + #[custom_method(ProviderConfigStatusRequest)] + async fn dispatch_provider_config_status( + &self, + req: ProviderConfigStatusRequest, + ) -> Result { + self.on_provider_config_status(req).await + } + + #[custom_method(ProviderConfigSaveRequest)] + async fn dispatch_save_provider_config( + &self, + req: ProviderConfigSaveRequest, + ) -> Result { + self.on_save_provider_config(req).await + } + + #[custom_method(ProviderConfigDeleteRequest)] + async fn dispatch_delete_provider_config( + &self, + req: ProviderConfigDeleteRequest, + ) -> Result { + self.on_delete_provider_config(req).await + } + + #[custom_method(ReadConfigRequest)] + async fn dispatch_read_config( + &self, + req: ReadConfigRequest, + ) -> Result { + self.on_read_config(req).await + } + + #[custom_method(UpsertConfigRequest)] + async fn dispatch_upsert_config( + &self, + req: UpsertConfigRequest, + ) -> Result { + self.on_upsert_config(req).await + } + + #[custom_method(RemoveConfigRequest)] + async fn dispatch_remove_config( + &self, + req: RemoveConfigRequest, + ) -> Result { + self.on_remove_config(req).await + } + + #[custom_method(CheckSecretRequest)] + async fn dispatch_check_secret( + &self, + req: CheckSecretRequest, + ) -> Result { + self.on_check_secret(req).await + } + + #[custom_method(UpsertSecretRequest)] + async fn dispatch_upsert_secret( + &self, + req: UpsertSecretRequest, + ) -> Result { + self.on_upsert_secret(req).await + } + + #[custom_method(RemoveSecretRequest)] + async fn dispatch_remove_secret( + &self, + req: RemoveSecretRequest, + ) -> Result { + self.on_remove_secret(req).await + } + + #[custom_method(ExportSessionRequest)] + async fn dispatch_export_session( + &self, + req: ExportSessionRequest, + ) -> Result { + self.on_export_session(req).await + } + + #[custom_method(ImportSessionRequest)] + async fn dispatch_import_session( + &self, + req: ImportSessionRequest, + ) -> Result { + self.on_import_session(req).await + } + + #[custom_method(UpdateSessionProjectRequest)] + async fn dispatch_update_session_project( + &self, + req: UpdateSessionProjectRequest, + ) -> Result { + self.on_update_session_project(req).await + } + + #[custom_method(RenameSessionRequest)] + async fn dispatch_rename_session( + &self, + req: RenameSessionRequest, + ) -> Result { + self.on_rename_session(req).await + } + + #[custom_method(ArchiveSessionRequest)] + async fn dispatch_archive_session( + &self, + req: ArchiveSessionRequest, + ) -> Result { + self.on_archive_session(req).await + } + + #[custom_method(UnarchiveSessionRequest)] + async fn dispatch_unarchive_session( + &self, + req: UnarchiveSessionRequest, + ) -> Result { + self.on_unarchive_session(req).await + } + + #[custom_method(CreateSourceRequest)] + async fn dispatch_create_source( + &self, + req: CreateSourceRequest, + ) -> Result { + self.on_create_source(req).await + } + + #[custom_method(ListSourcesRequest)] + async fn dispatch_list_sources( + &self, + req: ListSourcesRequest, + ) -> Result { + self.on_list_sources(req).await + } + + #[custom_method(UpdateSourceRequest)] + async fn dispatch_update_source( + &self, + req: UpdateSourceRequest, + ) -> Result { + self.on_update_source(req).await + } + + #[custom_method(DeleteSourceRequest)] + async fn dispatch_delete_source( + &self, + req: DeleteSourceRequest, + ) -> Result { + self.on_delete_source(req).await + } + + #[custom_method(ExportSourceRequest)] + async fn dispatch_export_source( + &self, + req: ExportSourceRequest, + ) -> Result { + self.on_export_source(req).await + } + + #[custom_method(ImportSourcesRequest)] + async fn dispatch_import_sources( + &self, + req: ImportSourcesRequest, + ) -> Result { + self.on_import_sources(req).await + } + + #[custom_method(DictationTranscribeRequest)] + async fn dispatch_dictation_transcribe( + &self, + req: DictationTranscribeRequest, + ) -> Result { + self.on_dictation_transcribe(req).await + } + + #[custom_method(DictationConfigRequest)] + async fn dispatch_dictation_config( + &self, + _req: DictationConfigRequest, + ) -> Result { + self.on_dictation_config(_req).await + } + + #[custom_method(DictationModelsListRequest)] + async fn dispatch_dictation_models_list( + &self, + _req: DictationModelsListRequest, + ) -> Result { + self.on_dictation_models_list(_req).await + } + + #[custom_method(DictationModelDownloadRequest)] + async fn dispatch_dictation_model_download( + &self, + _req: DictationModelDownloadRequest, + ) -> Result { + self.on_dictation_model_download(_req).await + } + + #[custom_method(DictationModelDownloadProgressRequest)] + async fn dispatch_dictation_model_download_progress( + &self, + _req: DictationModelDownloadProgressRequest, + ) -> Result { + self.on_dictation_model_download_progress(_req).await + } + + #[custom_method(DictationModelCancelRequest)] + async fn dispatch_dictation_model_cancel( + &self, + _req: DictationModelCancelRequest, + ) -> Result { + self.on_dictation_model_cancel(_req).await + } + + #[custom_method(DictationModelDeleteRequest)] + async fn dispatch_dictation_model_delete( + &self, + _req: DictationModelDeleteRequest, + ) -> Result { + self.on_dictation_model_delete(_req).await + } + + #[custom_method(DictationModelSelectRequest)] + async fn dispatch_dictation_model_select( + &self, + req: DictationModelSelectRequest, + ) -> Result { + self.on_dictation_model_select(req).await + } +} diff --git a/crates/goose/src/acp/server/dictation.rs b/crates/goose/src/acp/server/dictation.rs new file mode 100644 index 00000000..135d9d4b --- /dev/null +++ b/crates/goose/src/acp/server/dictation.rs @@ -0,0 +1,406 @@ +use super::*; +#[cfg(feature = "local-inference")] +use crate::dictation::providers::transcribe_local; +use crate::dictation::providers::{ + all_providers, is_configured, transcribe_with_provider, DictationProvider, +}; +#[cfg(feature = "local-inference")] +use crate::dictation::whisper; + +const OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "OPENAI_TRANSCRIPTION_MODEL"; +const GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "GROQ_TRANSCRIPTION_MODEL"; +const ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY: &str = "ELEVENLABS_TRANSCRIPTION_MODEL"; +const OPENAI_TRANSCRIPTION_MODEL: &str = "whisper-1"; +const GROQ_TRANSCRIPTION_MODEL: &str = "whisper-large-v3-turbo"; +const ELEVENLABS_TRANSCRIPTION_MODEL: &str = "scribe_v1"; + +impl GooseAcpAgent { + pub(super) async fn on_dictation_transcribe( + &self, + req: DictationTranscribeRequest, + ) -> Result { + use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; + let config = crate::config::Config::global(); + + #[cfg(not(feature = "local-inference"))] + if req.provider == "local" { + return Err(sacp::Error::invalid_params() + .data("Local inference is not available in this build")); + } + + let provider: DictationProvider = serde_json::from_value(serde_json::Value::String( + req.provider.clone(), + )) + .map_err(|_| { + sacp::Error::invalid_params().data(format!("Unknown provider: {}", req.provider)) + })?; + + let audio_bytes = BASE64 + .decode(&req.audio) + .map_err(|_| sacp::Error::invalid_params().data("Invalid base64 audio data"))?; + + if audio_bytes.len() > 50 * 1024 * 1024 { + return Err(sacp::Error::invalid_params().data("Audio too large (max 50MB)")); + } + + let extension = match req.mime_type.as_str() { + "audio/webm" | "audio/webm;codecs=opus" => "webm", + "audio/mp4" => "mp4", + "audio/mpeg" | "audio/mpga" => "mp3", + "audio/m4a" => "m4a", + "audio/wav" | "audio/x-wav" => "wav", + other => { + return Err( + sacp::Error::invalid_params().data(format!("Unsupported format: {other}")) + ); + } + }; + + let text = match provider { + #[cfg(feature = "local-inference")] + DictationProvider::Local => transcribe_local(audio_bytes).await, + remote => { + let (model_param, default_model) = dictation_transcribe_params(remote); + let model = dictation_selected_model(config, remote) + .unwrap_or_else(|| default_model.to_string()); + transcribe_with_provider( + remote, + model_param.to_string(), + model, + audio_bytes, + extension, + &req.mime_type, + ) + .await + } + } + .internal_err()?; + + Ok(DictationTranscribeResponse { text }) + } + + pub(super) async fn on_dictation_config( + &self, + _req: DictationConfigRequest, + ) -> Result { + let config = crate::config::Config::global(); + let mut providers = std::collections::HashMap::new(); + + for def in all_providers() { + let provider = def.provider; + let host = if let Some(host_key) = def.host_key { + config + .get(host_key, false) + .ok() + .and_then(|v| v.as_str().map(|s| s.to_string())) + } else { + None + }; + + let provider_key = serde_json::to_value(provider) + .ok() + .and_then(|v| v.as_str().map(|s| s.to_string())) + .unwrap_or_else(|| format!("{:?}", provider).to_lowercase()); + providers.insert( + provider_key, + DictationProviderStatusEntry { + configured: is_configured(provider), + host, + description: def.description.to_string(), + uses_provider_config: def.uses_provider_config, + settings_path: def.settings_path.map(|s| s.to_string()), + config_key: if !def.uses_provider_config { + Some(def.config_key.to_string()) + } else { + None + }, + model_config_key: dictation_model_config_key(provider), + default_model: dictation_default_model(provider), + selected_model: dictation_selected_model(config, provider), + available_models: dictation_available_models(provider), + }, + ); + } + + Ok(DictationConfigResponse { providers }) + } + + pub(super) async fn on_dictation_models_list( + &self, + _req: DictationModelsListRequest, + ) -> Result { + #[cfg(feature = "local-inference")] + { + use crate::download_manager::{get_download_manager, DownloadStatus}; + + let manager = get_download_manager(); + let models = whisper::available_models() + .iter() + .map(|model| DictationLocalModelStatus { + id: model.id.to_string(), + label: model.id.to_string(), + description: model.description.to_string(), + size_mb: model.size_mb, + downloaded: model.is_downloaded(), + download_in_progress: manager + .get_progress(model.id) + .map(|progress| progress.status == DownloadStatus::Downloading) + .unwrap_or(false), + }) + .collect(); + + Ok(DictationModelsListResponse { models }) + } + + #[cfg(not(feature = "local-inference"))] + Ok(DictationModelsListResponse::default()) + } + + pub(super) async fn on_dictation_model_download( + &self, + _req: DictationModelDownloadRequest, + ) -> Result { + #[cfg(feature = "local-inference")] + { + use crate::download_manager::get_download_manager; + + let model = whisper::get_model(&_req.model_id) + .ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?; + let manager = get_download_manager(); + let model_id_for_config = model.id.to_string(); + + manager + .download_model( + model.id.to_string(), + model.url.to_string(), + model.local_path(), + Some(Box::new(move || { + let config = crate::config::Config::global(); + // Only auto-select this model if the user has no model + // currently selected. This prevents silently switching + // the active model mid-session when a user downloads an + // additional model while one is already in use. + let already_selected = config + .get(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, false) + .ok() + .and_then(|value| value.as_str().map(str::to_owned)) + .filter(|model_id| { + // Treat a deleted model file as no active selection + // so a fresh download can auto-select cleanly. + whisper::get_model(model_id) + .is_some_and(|model| model.is_downloaded()) + }); + if already_selected.is_none() { + if let Err(e) = config.set_param( + whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, + model_id_for_config.clone(), + ) { + error!("Failed to save LOCAL_WHISPER_MODEL after download: {}", e); + } + } + })), + ) + .await + .internal_err()?; + + Ok(EmptyResponse {}) + } + + #[cfg(not(feature = "local-inference"))] + Err(sacp::Error::invalid_params().data("Local inference not enabled")) + } + + pub(super) async fn on_dictation_model_download_progress( + &self, + _req: DictationModelDownloadProgressRequest, + ) -> Result { + #[cfg(feature = "local-inference")] + { + use crate::download_manager::get_download_manager; + + let manager = get_download_manager(); + let progress = + manager + .get_progress(&_req.model_id) + .map(|progress| DictationDownloadProgress { + bytes_downloaded: progress.bytes_downloaded, + total_bytes: progress.total_bytes, + progress_percent: progress.progress_percent, + status: serde_json::to_value(&progress.status) + .ok() + .and_then(|value| value.as_str().map(ToOwned::to_owned)) + .unwrap_or_else(|| "unknown".to_string()), + error: progress.error, + }); + + Ok(DictationModelDownloadProgressResponse { progress }) + } + + #[cfg(not(feature = "local-inference"))] + Ok(DictationModelDownloadProgressResponse { progress: None }) + } + + pub(super) async fn on_dictation_model_cancel( + &self, + _req: DictationModelCancelRequest, + ) -> Result { + #[cfg(feature = "local-inference")] + { + use crate::download_manager::get_download_manager; + + let manager = get_download_manager(); + manager.cancel_download(&_req.model_id).internal_err()?; + + Ok(EmptyResponse {}) + } + + #[cfg(not(feature = "local-inference"))] + Err(sacp::Error::invalid_params().data("Local inference not enabled")) + } + + pub(super) async fn on_dictation_model_delete( + &self, + _req: DictationModelDeleteRequest, + ) -> Result { + #[cfg(feature = "local-inference")] + { + let model = whisper::get_model(&_req.model_id) + .ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?; + let path = model.local_path(); + + if !path.exists() { + return Err(sacp::Error::invalid_params().data("Model not downloaded")); + } + + std::fs::remove_file(path).internal_err()?; + + Ok(EmptyResponse {}) + } + + #[cfg(not(feature = "local-inference"))] + Err(sacp::Error::invalid_params().data("Local inference not enabled")) + } + + pub(super) async fn on_dictation_model_select( + &self, + req: DictationModelSelectRequest, + ) -> Result { + #[cfg(not(feature = "local-inference"))] + if req.provider == "local" { + return Err(sacp::Error::invalid_params().data("Local inference not enabled")); + } + + let provider: DictationProvider = serde_json::from_value(serde_json::Value::String( + req.provider.clone(), + )) + .map_err(|_| { + sacp::Error::invalid_params().data(format!("Unknown provider: {}", req.provider)) + })?; + + let key = match provider { + DictationProvider::OpenAI => OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY, + DictationProvider::Groq => GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY, + DictationProvider::ElevenLabs => ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY, + #[cfg(feature = "local-inference")] + DictationProvider::Local => { + let model = whisper::get_model(&req.model_id) + .ok_or_else(|| sacp::Error::invalid_params().data("Unknown model id"))?; + if !model.is_downloaded() { + return Err( + sacp::Error::invalid_params().data("Local Whisper model is not downloaded") + ); + } + whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY + } + }; + + crate::config::Config::global() + .set_param(key, req.model_id) + .internal_err()?; + + Ok(EmptyResponse {}) + } +} +fn dictation_model_config_key(provider: DictationProvider) -> Option { + match provider { + DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()), + DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()), + DictationProvider::ElevenLabs => { + Some(ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()) + } + #[cfg(feature = "local-inference")] + DictationProvider::Local => Some(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY.to_string()), + } +} + +/// Returns the (param_name, default_model) pair used by `transcribe_with_provider` +/// for remote dictation providers. Local inference is handled separately. +fn dictation_transcribe_params(provider: DictationProvider) -> (&'static str, &'static str) { + match provider { + DictationProvider::OpenAI => ("model", OPENAI_TRANSCRIPTION_MODEL), + DictationProvider::Groq => ("model", GROQ_TRANSCRIPTION_MODEL), + DictationProvider::ElevenLabs => ("model_id", ELEVENLABS_TRANSCRIPTION_MODEL), + #[cfg(feature = "local-inference")] + DictationProvider::Local => ("", ""), + } +} + +fn dictation_default_model(provider: DictationProvider) -> Option { + match provider { + DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL.to_string()), + DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL.to_string()), + DictationProvider::ElevenLabs => Some(ELEVENLABS_TRANSCRIPTION_MODEL.to_string()), + #[cfg(feature = "local-inference")] + DictationProvider::Local => Some(whisper::recommend_model().to_string()), + } +} + +fn dictation_selected_model(config: &Config, provider: DictationProvider) -> Option { + #[cfg(feature = "local-inference")] + if provider == DictationProvider::Local { + return config + .get(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY, false) + .ok() + .and_then(|value| value.as_str().map(str::to_owned)) + .filter(|model_id| whisper::get_model(model_id).is_some()) + .or_else(|| dictation_default_model(provider)); + } + + dictation_model_config_key(provider) + .and_then(|key| { + config + .get(&key, false) + .ok() + .and_then(|value| value.as_str().map(str::to_owned)) + }) + .or_else(|| dictation_default_model(provider)) +} + +fn dictation_available_models(provider: DictationProvider) -> Vec { + match provider { + DictationProvider::OpenAI => vec![DictationModelOption { + id: OPENAI_TRANSCRIPTION_MODEL.to_string(), + label: "Whisper-1".to_string(), + description: "OpenAI's hosted Whisper transcription model.".to_string(), + }], + DictationProvider::Groq => vec![DictationModelOption { + id: GROQ_TRANSCRIPTION_MODEL.to_string(), + label: "Whisper Large V3 Turbo".to_string(), + description: "Groq's fast hosted Whisper transcription model.".to_string(), + }], + DictationProvider::ElevenLabs => vec![DictationModelOption { + id: ELEVENLABS_TRANSCRIPTION_MODEL.to_string(), + label: "Scribe v1".to_string(), + description: "ElevenLabs' hosted speech-to-text model.".to_string(), + }], + #[cfg(feature = "local-inference")] + DictationProvider::Local => whisper::available_models() + .iter() + .map(|model| DictationModelOption { + id: model.id.to_string(), + label: model.id.to_string(), + description: model.description.to_string(), + }) + .collect(), + } +} diff --git a/crates/goose/src/acp/server/dispatch.rs b/crates/goose/src/acp/server/dispatch.rs new file mode 100644 index 00000000..e00bec39 --- /dev/null +++ b/crates/goose/src/acp/server/dispatch.rs @@ -0,0 +1,397 @@ +use super::*; + +impl HandleDispatchFrom for GooseAcpHandler { + fn describe_chain(&self) -> impl std::fmt::Debug { + "goose-acp" + } + + fn handle_dispatch_from( + &mut self, + message: Dispatch, + cx: ConnectionTo, + ) -> impl std::future::Future, sacp::Error>> + Send { + let agent = self.agent.clone(); + + // The MatchDispatchFrom chain produces an ~85KB async state machine. + // Box::pin moves it to the heap so it doesn't overflow the tokio worker stack. + Box::pin(async move { + // InitializeRequest runs inline: it sets connection-scoped state + // (client fs/terminal capabilities) that later handlers read with + // defaults, so a pipelined NewSessionRequest must not race ahead of it. + MatchDispatchFrom::new(message, &cx) + .if_request( + |req: InitializeRequest, responder: Responder| async { + responder.respond_with_result(agent.on_initialize(req).await) + }, + ) + .await + .if_request( + |_req: AuthenticateRequest, responder: Responder| async { + responder.respond(AuthenticateResponse::new()) + }, + ) + .await + .if_request( + |req: NewSessionRequest, responder: Responder| async { + let agent = agent.clone(); + let cx_clone = cx.clone(); + cx.spawn(async move { + responder.respond_with_result(agent.on_new_session(&cx_clone, req).await)?; + Ok(()) + })?; + Ok(()) + }, + ) + .await + .if_request( + |req: LoadSessionRequest, responder: Responder| async { + let agent = agent.clone(); + let cx_clone = cx.clone(); + cx.spawn(async move { + match agent.on_load_session(&cx_clone, req).await { + Ok(response) => { + responder.respond(response)?; + } + Err(e) => { + responder.respond_with_error(e)?; + } + } + Ok(()) + })?; + Ok(()) + }, + ) + .await + .if_request( + |req: PromptRequest, responder: Responder| async { + let agent = agent.clone(); + let cx_clone = cx.clone(); + cx.spawn(async move { + match agent.on_prompt(&cx_clone, req).await { + Ok(response) => { + responder.respond(response)?; + } + Err(e) => { + responder.respond_with_error(e)?; + } + } + Ok(()) + })?; + Ok(()) + }, + ) + .await + .if_notification(|notif: CancelNotification| async { + let agent = agent.clone(); + agent.on_cancel(notif).await?; + Ok(()) + }) + .await + // set_config_option (SACP 11) and legacy set_mode/set_model; custom _goose/* in otherwise. + .if_request({ + let agent = agent.clone(); + let cx = cx.clone(); + |req: SetSessionConfigOptionRequest, responder: Responder| async move { + let cx_spawn = cx.clone(); + cx.spawn(async move { + let cx = cx_spawn; + let value_id = req.value.as_value_id() + .ok_or_else(|| sacp::Error::invalid_params().data("Expected a value ID"))? + .clone(); + let session_id = req.session_id.clone(); + let sid = sid_short(session_id.0.as_ref()); + let config_id = req.config_id.0.to_string(); + let t_handler = std::time::Instant::now(); + match config_id.as_ref() { + "provider" => { + Config::global().invalidate_secrets_cache(); + 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(_) => {} + Err(e) => { responder.respond_with_error(e)?; return Ok(()); } + } + } + "model" => { + match agent.on_set_model(&session_id.0, &value_id.0).await { + Ok(_) => {} + Err(e) => { responder.respond_with_error(e)?; return Ok(()); } + } + } + other => { + responder.respond_with_error( + sacp::Error::invalid_params().data(format!("Unsupported config option: {}", other)) + )?; + return Ok(()); + } + } + // Respond immediately using the current provider inventory snapshot. + let (notification, config_options) = agent.build_config_update(&session_id).await?; + cx.send_notification(notification)?; + responder.respond(SetSessionConfigOptionResponse::new(config_options))?; + + let maybe_refresh = if config_id == "provider" { + let provider_id = value_id.0.to_string(); + agent + .provider_inventory + .plan_refresh_jobs(std::slice::from_ref(&provider_id)) + .await + .ok() + .and_then(|plan| { + plan.started + .into_iter() + .find(|job| job.provider_id == provider_id) + }) + } else { + None + }; + if let Some(refresh_job) = maybe_refresh { + let agent_bg = agent.clone(); + let cx_bg = cx.clone(); + let session_id_bg = session_id.clone(); + tokio::spawn(async move { + let refresh_identity = refresh_job.identity; + let refresh_provider_id = refresh_job.provider_id; + let mut refresh_guard = + agent_bg.provider_inventory.refresh_guard(&refresh_identity); + let provider_result: Result> = + AssertUnwindSafe(async { + let session_agent = + agent_bg.get_session_agent(&session_id_bg.0, None).await?; + let provider = session_agent + .provider() + .await + .map_err(|e| anyhow::anyhow!(e.to_string()))?; + let provider_name = provider.get_name().to_string(); + if provider_name != refresh_provider_id { + return Err(anyhow::anyhow!( + "provider changed before inventory refresh completed" + )); + } + Ok(provider) + }) + .catch_unwind() + .await + .map_err(|_| { + anyhow::anyhow!("provider inventory refresh task panicked") + }) + .and_then(|result| result); + + let fetch_result = match provider_result { + Ok(provider) => { + match ensure_refresh_identity_current( + &refresh_provider_id, + &refresh_identity, + ) + .await + { + Ok(()) => match AssertUnwindSafe( + provider.fetch_recommended_models(), + ) + .catch_unwind() + .await + { + Ok(Ok(models)) => Ok(models), + Ok(Err(error)) => { + Err(anyhow::anyhow!(error.to_string())) + } + Err(_) => Err(anyhow::anyhow!( + "provider inventory refresh task panicked" + )), + }, + Err(error) => Err(error), + } + } + Err(error) => Err(error), + }; + + match fetch_result { + Ok(models) => match agent_bg + .provider_inventory + .store_refreshed_models_for_identity( + &refresh_identity, + &models, + ) + .await + { + Ok(()) => { + refresh_guard.complete(); + match agent_bg.build_config_update(&session_id_bg).await + { + Ok((fresh_notification, _)) => { + let _ = cx_bg + .send_notification(fresh_notification); + } + Err(error) => warn!( + provider = %refresh_provider_id, + error = %error, + "failed to build config update after provider inventory refresh" + ), + } + } + Err(error) => warn!( + provider = %refresh_provider_id, + error = %error, + "failed to store refreshed provider inventory after config change" + ), + }, + Err(error) => { + let error_message = error.to_string(); + match agent_bg + .provider_inventory + .store_refresh_error_for_identity( + &refresh_identity, + error_message.clone(), + ) + .await + { + Ok(()) => refresh_guard.complete(), + Err(store_error) => warn!( + provider = %refresh_provider_id, + error = %store_error, + refresh_error = %error_message, + "failed to store provider inventory refresh error after config change" + ), + } + warn!( + provider = %refresh_provider_id, + error = %error_message, + "provider inventory refresh failed after config change" + ); + } + } + }); + } + + debug!(target: "perf", sid = %sid, ms = t_handler.elapsed().as_millis() as u64, config_id = %config_id, "perf: set_config_option done"); + Ok(()) + })?; + Ok(()) + } + }) + .await + .if_request({ + let agent = agent.clone(); + let cx = cx.clone(); + |req: SetSessionModeRequest, responder: Responder| async move { + let cx_spawn = cx.clone(); + cx.spawn(async move { + let cx = cx_spawn; + let session_id = req.session_id.clone(); + let mode_id = req.mode_id.clone(); + match agent.on_set_mode(&session_id.0, &mode_id.0).await { + Ok(resp) => { + // Notify before responding so clients see the mode update before block_task unblocks. + cx.send_notification(SessionNotification::new( + session_id, + SessionUpdate::CurrentModeUpdate( + CurrentModeUpdate::new(mode_id), + ), + ))?; + responder.respond(resp)?; + } + Err(e) => { + responder.respond_with_error(e)?; + } + } + Ok(()) + })?; + Ok(()) + } + }) + .await + .if_request({ + let agent = agent.clone(); + let cx = cx.clone(); + |req: SetSessionModelRequest, responder: Responder| async move { + let cx_spawn = cx.clone(); + cx.spawn(async move { + let cx = cx_spawn; + let session_id = req.session_id.clone(); + match agent.on_set_model(&session_id.0, &req.model_id.0).await { + Ok(resp) => { + let (notification, _) = agent.build_config_update(&session_id).await?; + cx.send_notification(notification)?; + responder.respond(resp)?; + } + Err(e) => responder.respond_with_error(e)?, + } + Ok(()) + })?; + Ok(()) + } + }) + .await + .if_request({ + let agent = agent.clone(); + let cx = cx.clone(); + |_req: ListSessionsRequest, responder: Responder| async move { + cx.spawn(async move { + responder.respond(agent.on_list_sessions().await?)?; + Ok(()) + })?; + Ok(()) + } + }) + .await + .if_request({ + let agent = agent.clone(); + let cx = cx.clone(); + |req: CloseSessionRequest, responder: Responder| async move { + cx.spawn(async move { + responder.respond(agent.on_close_session(&req.session_id.0).await?)?; + Ok(()) + })?; + Ok(()) + } + }) + .await + .if_request({ + let agent = agent.clone(); + let cx = cx.clone(); + |req: ForkSessionRequest, responder: Responder| async move { + let cx_spawn = cx.clone(); + cx.spawn(async move { + responder.respond_with_result(agent.on_fork_session(&cx_spawn, req).await)?; + Ok(()) + })?; + Ok(()) + } + }) + .await + .otherwise({ + let agent = agent.clone(); + let cx = cx.clone(); + |message: Dispatch| async move { + match message { + Dispatch::Request(req, responder) => { + cx.spawn(async move { + match agent.dispatch_custom_request(&req.method, req.params).await { + Ok(json) => responder.respond(json)?, + Err(e) => responder.respond_with_error(e)?, + } + Ok(()) + })?; + Ok(()) + } + Dispatch::Response(result, router) => { + debug!(method = %router.method(), id = %router.id(), ok = result.is_ok(), "routing response"); + router.respond_with_result(result)?; + Ok(()) + } + Dispatch::Notification(notif) => { + debug!(method = %notif.method, "unhandled notification"); + Ok(()) + } + } + } + }) + .await + .map(|()| Handled::Yes) + }) + } +} diff --git a/crates/goose/src/acp/server/extensions.rs b/crates/goose/src/acp/server/extensions.rs new file mode 100644 index 00000000..f78777c4 --- /dev/null +++ b/crates/goose/src/acp/server/extensions.rs @@ -0,0 +1,136 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_add_extension( + &self, + req: AddExtensionRequest, + ) -> Result { + let internal_id = self.internal_session_id(&req.session_id).await?; + let config: ExtensionConfig = serde_json::from_value(req.config) + .map_err(|e| sacp::Error::invalid_params().data(format!("bad config: {e}")))?; + let agent = self.get_session_agent(&req.session_id, None).await?; + agent + .add_extension(config, &internal_id) + .await + .internal_err()?; + Ok(EmptyResponse {}) + } + + pub(super) async fn on_remove_extension( + &self, + req: RemoveExtensionRequest, + ) -> Result { + let internal_id = self.internal_session_id(&req.session_id).await?; + let agent = self.get_session_agent(&req.session_id, None).await?; + agent + .remove_extension(&req.name, &internal_id) + .await + .internal_err()?; + Ok(EmptyResponse {}) + } + + pub(super) async fn on_get_extensions(&self) -> Result { + let extensions = crate::config::extensions::get_all_extensions(); + let warnings = crate::config::extensions::get_warnings(); + let extensions_json = extensions + .into_iter() + .map(|e| { + let config_key = e.config.key(); + let mut value = serde_json::to_value(&e)?; + if let Some(obj) = value.as_object_mut() { + obj.insert( + "config_key".to_string(), + serde_json::Value::String(config_key), + ); + } + Ok::<_, serde_json::Error>(value) + }) + .collect::, _>>() + .internal_err()?; + Ok(GetExtensionsResponse { + extensions: extensions_json, + warnings, + }) + } + + pub(super) async fn on_add_config_extension( + &self, + req: AddConfigExtensionRequest, + ) -> Result { + let mut obj = match req.extension_config { + serde_json::Value::Object(obj) => obj, + _ => { + return Err( + sacp::Error::invalid_params().data("extensionConfig must be a JSON object") + ); + } + }; + obj.insert( + "name".to_string(), + serde_json::Value::String(req.name.clone()), + ); + + let config: crate::agents::ExtensionConfig = + serde_json::from_value(serde_json::Value::Object(obj)) + .map_err(|e| sacp::Error::invalid_params().data(format!("bad config: {e}")))?; + + crate::config::extensions::set_extension(crate::config::extensions::ExtensionEntry { + enabled: req.enabled, + config, + }); + Ok(EmptyResponse {}) + } + + pub(super) async fn on_remove_config_extension( + &self, + req: RemoveConfigExtensionRequest, + ) -> Result { + let keys = crate::config::extensions::get_all_extension_names(); + if !keys.iter().any(|k| k == &req.config_key) { + return Err(sacp::Error::invalid_params() + .data(format!("Extension '{}' not found", req.config_key))); + } + crate::config::extensions::remove_extension(&req.config_key); + Ok(EmptyResponse {}) + } + + pub(super) async fn on_toggle_config_extension( + &self, + req: ToggleConfigExtensionRequest, + ) -> Result { + let keys = crate::config::extensions::get_all_extension_names(); + if !keys.iter().any(|k| k == &req.config_key) { + return Err(sacp::Error::invalid_params() + .data(format!("Extension '{}' not found", req.config_key))); + } + crate::config::extensions::set_extension_enabled(&req.config_key, req.enabled); + Ok(EmptyResponse {}) + } + + pub(super) async fn on_get_session_extensions( + &self, + req: GetSessionExtensionsRequest, + ) -> Result { + let internal_id = self.internal_session_id(&req.session_id).await?; + let session = self + .session_manager + .get_session(&internal_id, false) + .await + .internal_err()?; + + let extensions = EnabledExtensionsState::extensions_or_default( + Some(&session.extension_data), + crate::config::Config::global(), + ); + + let extensions_json = extensions + .into_iter() + .map(|e| serde_json::to_value(&e)) + .collect::, _>>() + .internal_err()?; + + Ok(GetSessionExtensionsResponse { + extensions: extensions_json, + }) + } +} diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs new file mode 100644 index 00000000..7297e1c7 --- /dev/null +++ b/crates/goose/src/acp/server/providers.rs @@ -0,0 +1,420 @@ +use super::*; + +fn inventory_entry_to_dto(entry: ProviderInventoryEntry) -> ProviderInventoryEntryDto { + let stale = ProviderInventoryService::is_stale(&entry); + ProviderInventoryEntryDto { + provider_id: entry.provider_id, + provider_name: entry.provider_name, + description: entry.description, + default_model: entry.default_model, + configured: entry.configured, + provider_type: format!("{:?}", entry.provider_type), + config_keys: entry + .config_keys + .into_iter() + .map(provider_config_key_to_dto) + .collect(), + setup_steps: entry.setup_steps, + supports_refresh: entry.supports_refresh, + refreshing: entry.refreshing, + models: entry + .models + .into_iter() + .map(|m| ProviderInventoryModelDto { + id: m.id, + name: m.name, + family: m.family, + context_limit: m.context_limit, + reasoning: m.reasoning, + recommended: m.recommended, + }) + .collect(), + last_updated_at: entry.last_updated_at.map(|t| t.to_rfc3339()), + last_refresh_attempt_at: entry.last_refresh_attempt_at.map(|t| t.to_rfc3339()), + last_refresh_error: entry.last_refresh_error, + stale, + model_selection_hint: entry.model_selection_hint, + } +} + +fn provider_config_key_to_dto(key: crate::providers::base::ConfigKey) -> ProviderConfigKey { + ProviderConfigKey { + name: key.name, + required: key.required, + secret: key.secret, + default: key.default, + oauth_flow: key.oauth_flow, + device_code_flow: key.device_code_flow, + primary: key.primary, + } +} + +const SECRET_MASK_PREFIX_LEN: usize = 4; +const SECRET_MASK_SUFFIX_LEN: usize = 3; +const SECRET_MASK_FALLBACK: &str = "***"; + +fn mask_secret_value(value: &str) -> String { + let prefix: String = value.chars().take(SECRET_MASK_PREFIX_LEN).collect(); + let suffix_chars: Vec = value.chars().rev().take(SECRET_MASK_SUFFIX_LEN).collect(); + let suffix: String = suffix_chars.into_iter().rev().collect(); + + if prefix.is_empty() + || suffix.is_empty() + || value.chars().count() <= SECRET_MASK_PREFIX_LEN + SECRET_MASK_SUFFIX_LEN + { + return SECRET_MASK_FALLBACK.to_string(); + } + + format!("{prefix}...{suffix}") +} + +fn config_value_to_string(value: &serde_json::Value) -> Option { + match value { + serde_json::Value::Null => None, + serde_json::Value::String(value) if value.is_empty() => None, + serde_json::Value::String(value) => Some(value.clone()), + other => serde_json::to_string(other).ok(), + } +} + +fn provider_config_field_value( + config: &Config, + key: &crate::providers::base::ConfigKey, + secrets: Option<&HashMap>, +) -> ProviderConfigFieldValueDto { + let value = if key.secret { + std::env::var(key.name.to_uppercase()).ok().or_else(|| { + secrets + .and_then(|values| values.get(&key.name)) + .and_then(config_value_to_string) + }) + } else { + config + .get_param::(&key.name) + .ok() + .and_then(|value| config_value_to_string(&value)) + }; + + ProviderConfigFieldValueDto { + key: key.name.clone(), + value: value.as_deref().map(|value| { + if key.secret { + mask_secret_value(value) + } else { + value.to_string() + } + }), + is_set: value.is_some(), + is_secret: key.secret, + required: key.required, + } +} + +fn refresh_skip_reason_to_dto(reason: RefreshSkipReason) -> RefreshProviderInventorySkipReasonDto { + match reason { + RefreshSkipReason::UnknownProvider => { + RefreshProviderInventorySkipReasonDto::UnknownProvider + } + RefreshSkipReason::NotConfigured => RefreshProviderInventorySkipReasonDto::NotConfigured, + RefreshSkipReason::DoesNotSupportRefresh => { + RefreshProviderInventorySkipReasonDto::DoesNotSupportRefresh + } + RefreshSkipReason::AlreadyRefreshing => { + RefreshProviderInventorySkipReasonDto::AlreadyRefreshing + } + } +} + +fn refresh_plan_to_response(refresh_plan: RefreshPlan) -> RefreshProviderInventoryResponse { + RefreshProviderInventoryResponse { + started: refresh_plan.started, + skipped: refresh_plan + .skipped + .into_iter() + .map(|entry| RefreshProviderInventorySkipDto { + provider_id: entry.provider_id, + reason: refresh_skip_reason_to_dto(entry.reason), + }) + .collect(), + } +} + +impl GooseAcpAgent { + pub(super) async fn on_list_providers( + &self, + req: ListProvidersRequest, + ) -> Result { + let entries = self + .provider_inventory + .entries(&req.provider_ids) + .await + .internal_err()?; + Ok(ListProvidersResponse { + entries: entries.into_iter().map(inventory_entry_to_dto).collect(), + }) + } + + pub(super) async fn provider_config_status(provider_id: String) -> ProviderConfigStatusDto { + let is_configured = match crate::providers::get_from_registry(&provider_id).await { + Ok(entry) => { + match tokio::task::spawn_blocking(move || entry.inventory_configured()).await { + Ok(is_configured) => is_configured, + Err(error) => { + warn!( + provider = %provider_id, + error = %error, + "provider config status check failed" + ); + false + } + } + } + Err(_) => false, + }; + + ProviderConfigStatusDto { + provider_id, + is_configured, + } + } + + pub(super) async fn provider_config_statuses( + provider_ids: &[String], + ) -> Vec { + let mut ids = if provider_ids.is_empty() { + crate::providers::providers() + .await + .into_iter() + .map(|(metadata, _)| metadata.name) + .collect::>() + } else { + provider_ids.to_vec() + }; + ids.sort(); + ids.dedup(); + + let mut statuses = stream::iter(ids) + .map(Self::provider_config_status) + .buffer_unordered(PROVIDER_CONFIG_STATUS_CHECK_CONCURRENCY) + .collect::>() + .await; + statuses.sort_by(|a, b| a.provider_id.cmp(&b.provider_id)); + statuses + } + + pub(super) fn spawn_provider_inventory_refresh_jobs(&self, refresh_plan: &RefreshJobPlan) { + for refresh_job in refresh_plan.started.iter().cloned() { + let provider_inventory = self.provider_inventory.clone(); + let provider_factory = Arc::clone(&self.provider_factory); + let provider_id = refresh_job.provider_id.clone(); + let identity = refresh_job.identity.clone(); + tokio::spawn(async move { + let mut refresh_guard = provider_inventory.refresh_guard(&identity); + let provider_result = AssertUnwindSafe(async { + let metadata = crate::providers::get_from_registry(&provider_id).await?; + let model_config = + crate::model::ModelConfig::new(&metadata.metadata().default_model)? + .with_canonical_limits(&provider_id); + provider_factory(provider_id.clone(), model_config, Vec::new()).await + }) + .catch_unwind() + .await; + + let fetch_result: Result> = match provider_result { + Ok(Ok(provider)) => { + match ensure_refresh_identity_current(&provider_id, &identity).await { + Ok(()) => match AssertUnwindSafe(provider.fetch_recommended_models()) + .catch_unwind() + .await + { + Ok(Ok(models)) => Ok(models), + Ok(Err(error)) => Err(anyhow::anyhow!(error.to_string())), + Err(_) => { + Err(anyhow::anyhow!("provider inventory refresh task panicked")) + } + }, + Err(error) => Err(error), + } + } + Ok(Err(error)) => Err(error), + Err(_) => Err(anyhow::anyhow!("provider inventory refresh task panicked")), + }; + + match fetch_result { + Ok(models) => match provider_inventory + .store_refreshed_models_for_identity(&identity, &models) + .await + { + Ok(()) => refresh_guard.complete(), + Err(error) => warn!( + provider = %provider_id, + error = %error, + "failed to store refreshed provider inventory" + ), + }, + Err(error) => { + let error_message = error.to_string(); + match provider_inventory + .store_refresh_error_for_identity(&identity, error_message.clone()) + .await + { + Ok(()) => refresh_guard.complete(), + Err(store_error) => warn!( + provider = %provider_id, + error = %store_error, + refresh_error = %error_message, + "failed to store provider inventory refresh error" + ), + } + warn!(provider = %provider_id, error = %error_message, "provider inventory refresh failed"); + } + } + }); + } + } + + pub(super) async fn start_provider_inventory_refresh( + &self, + provider_ids: &[String], + ) -> Result { + let refresh_job_plan = self + .provider_inventory + .plan_refresh_jobs(provider_ids) + .await + .internal_err()?; + self.spawn_provider_inventory_refresh_jobs(&refresh_job_plan); + Ok(refresh_plan_to_response( + refresh_job_plan.into_public_plan(), + )) + } + + pub(super) async fn on_refresh_provider_inventory( + &self, + req: RefreshProviderInventoryRequest, + ) -> Result { + Config::global().invalidate_secrets_cache(); + self.start_provider_inventory_refresh(&req.provider_ids) + .await + } + + pub(super) async fn on_read_provider_config( + &self, + req: ProviderConfigReadRequest, + ) -> Result { + let entry = crate::providers::get_from_registry(&req.provider_id) + .await + .invalid_params_err_ctx("Unknown provider")?; + let config = Config::global(); + let config_keys = &entry.metadata().config_keys; + let secrets = if config_keys.iter().any(|key| key.secret) { + Some(config.all_secrets().internal_err()?) + } else { + None + }; + + Ok(ProviderConfigReadResponse { + fields: config_keys + .iter() + .map(|key| provider_config_field_value(config, key, secrets.as_ref())) + .collect(), + }) + } + + pub(super) async fn on_provider_config_status( + &self, + req: ProviderConfigStatusRequest, + ) -> Result { + Ok(ProviderConfigStatusResponse { + statuses: Self::provider_config_statuses(&req.provider_ids).await, + }) + } + + pub(super) async fn on_save_provider_config( + &self, + req: ProviderConfigSaveRequest, + ) -> Result { + let entry = crate::providers::get_from_registry(&req.provider_id) + .await + .invalid_params_err_ctx("Unknown provider")?; + let metadata = entry.metadata().clone(); + let config = Config::global(); + let mut config_updates = Vec::new(); + let mut secret_updates = Vec::new(); + + for field in &req.fields { + let Some(config_key) = metadata + .config_keys + .iter() + .find(|config_key| config_key.name == field.key) + else { + return Err(sacp::Error::invalid_params() + .data(format!("Unsupported provider config field: {}", field.key))); + }; + + let value = field.value.trim(); + if value.is_empty() { + return Err(sacp::Error::invalid_params().data(format!( + "Provider config field cannot be empty: {}", + field.key + ))); + } + + if config_key.secret { + secret_updates.push(( + config_key.name.clone(), + serde_json::Value::String(value.to_string()), + )); + } else { + config_updates.push((config_key.name.clone(), value.to_string())); + } + } + + for (key, value) in config_updates { + config + .set_param(&key, &value) + .internal_err_ctx("Failed to save provider config field")?; + } + config + .set_secret_values(&secret_updates) + .internal_err_ctx("Failed to save provider secret fields")?; + + let provider_ids = [req.provider_id.clone()]; + let status = Self::provider_config_status(req.provider_id.clone()).await; + let refresh = self.start_provider_inventory_refresh(&provider_ids).await?; + Ok(ProviderConfigChangeResponse { status, refresh }) + } + + pub(super) async fn on_delete_provider_config( + &self, + req: ProviderConfigDeleteRequest, + ) -> Result { + let entry = crate::providers::get_from_registry(&req.provider_id) + .await + .invalid_params_err_ctx("Unknown provider")?; + let metadata = entry.metadata().clone(); + let config = Config::global(); + let mut secret_keys = Vec::new(); + + for config_key in &metadata.config_keys { + if config_key.secret { + secret_keys.push(config_key.name.clone()); + } else { + config + .delete(&config_key.name) + .internal_err_ctx("Failed to delete provider config field")?; + } + } + + config + .delete_secret_values(&secret_keys) + .internal_err_ctx("Failed to delete provider secret fields")?; + crate::providers::cleanup_provider(&req.provider_id) + .await + .internal_err_ctx("Failed to clean up provider state")?; + + let provider_ids = [req.provider_id.clone()]; + let status = Self::provider_config_status(req.provider_id.clone()).await; + let refresh = self.start_provider_inventory_refresh(&provider_ids).await?; + Ok(ProviderConfigChangeResponse { status, refresh }) + } +} diff --git a/crates/goose/src/acp/server/resources.rs b/crates/goose/src/acp/server/resources.rs new file mode 100644 index 00000000..f50a0d96 --- /dev/null +++ b/crates/goose/src/acp/server/resources.rs @@ -0,0 +1,21 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_read_resource( + &self, + req: ReadResourceRequest, + ) -> Result { + let internal_id = self.internal_session_id(&req.session_id).await?; + let agent = self.get_session_agent(&req.session_id, None).await?; + let cancel_token = CancellationToken::new(); + let result = agent + .extension_manager + .read_resource(&internal_id, &req.uri, &req.extension_name, cancel_token) + .await + .internal_err()?; + let result_json = serde_json::to_value(&result).internal_err()?; + Ok(ReadResourceResponse { + result: result_json, + }) + } +} diff --git a/crates/goose/src/acp/server/secrets.rs b/crates/goose/src/acp/server/secrets.rs new file mode 100644 index 00000000..16ab2250 --- /dev/null +++ b/crates/goose/src/acp/server/secrets.rs @@ -0,0 +1,32 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_check_secret( + &self, + req: CheckSecretRequest, + ) -> Result { + let config = self.config()?; + let exists = config.get_secret::(&req.key).is_ok(); + Ok(CheckSecretResponse { exists }) + } + + pub(super) async fn on_upsert_secret( + &self, + req: UpsertSecretRequest, + ) -> Result { + let config = self.config()?; + config.set_secret(&req.key, &req.value).internal_err()?; + Config::global().invalidate_secrets_cache(); + Ok(EmptyResponse {}) + } + + pub(super) async fn on_remove_secret( + &self, + req: RemoveSecretRequest, + ) -> Result { + let config = self.config()?; + config.delete_secret(&req.key).internal_err()?; + Config::global().invalidate_secrets_cache(); + Ok(EmptyResponse {}) + } +} diff --git a/crates/goose/src/acp/server/sessions.rs b/crates/goose/src/acp/server/sessions.rs new file mode 100644 index 00000000..8b5812f1 --- /dev/null +++ b/crates/goose/src/acp/server/sessions.rs @@ -0,0 +1,175 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_update_working_dir( + &self, + req: UpdateWorkingDirRequest, + ) -> Result { + let working_dir = req.working_dir.trim().to_string(); + if working_dir.is_empty() { + return Err(sacp::Error::invalid_params().data("working directory cannot be empty")); + } + let path = std::path::PathBuf::from(&working_dir); + if !path.exists() || !path.is_dir() { + return Err(sacp::Error::invalid_params().data("invalid directory path")); + } + let internal_id = self.internal_session_id(&req.session_id).await?; + self.session_manager + .update(&internal_id) + .working_dir(path.clone()) + .apply() + .await + .internal_err()?; + + self.thread_manager + .update_working_dir(&req.session_id, &working_dir) + .await + .internal_err()?; + + if let Some(session) = self.sessions.lock().await.get_mut(&req.session_id) { + match &session.agent { + AgentHandle::Ready(agent) => { + agent.extension_manager.update_working_dir(&path).await; + } + AgentHandle::Loading(_) => { + session.pending_working_dir = Some(path); + } + } + } + + Ok(EmptyResponse {}) + } + + pub(super) async fn on_delete_session( + &self, + req: DeleteSessionRequest, + ) -> Result { + // Delete the thread and all its internal sessions + messages. + self.thread_manager + .delete_thread(&req.session_id) + .await + .internal_err()?; + self.sessions.lock().await.remove(&req.session_id); + Ok(EmptyResponse {}) + } + + pub(super) async fn on_export_session( + &self, + req: ExportSessionRequest, + ) -> Result { + let thread = self + .thread_manager + .get_thread(&req.session_id) + .await + .internal_err()?; + let internal_id = thread + .current_session_id + .ok_or_else(|| sacp::Error::internal_error().data("Thread has no internal session"))?; + let data = self + .session_manager + .export_session(&internal_id) + .await + .internal_err()?; + Ok(ExportSessionResponse { data }) + } + + pub(super) async fn on_import_session( + &self, + req: ImportSessionRequest, + ) -> Result { + let session = self + .session_manager + .import_session(&req.data, Some(SessionType::Acp)) + .await + .internal_err()?; + + // Create a thread for the imported session. + let thread = self + .thread_manager + .create_thread( + Some(session.name.clone()), + None, + Some(session.working_dir.display().to_string()), + ) + .await + .internal_err()?; + + // Link the internal session to the thread. + self.session_manager + .update(&session.id) + .thread_id(Some(thread.id.clone())) + .apply() + .await + .internal_err()?; + + // Copy conversation messages into thread_messages so they appear in the thread. + if let Some(ref conversation) = session.conversation { + for msg in conversation.messages() { + self.thread_manager + .append_message(&thread.id, Some(&session.id), msg) + .await + .internal_err()?; + } + } + + // Re-fetch thread to get accurate message_count. + let thread = self + .thread_manager + .get_thread(&thread.id) + .await + .internal_err()?; + + Ok(ImportSessionResponse { + session_id: thread.id, + title: Some(thread.name), + updated_at: Some(thread.updated_at.to_rfc3339()), + message_count: thread.message_count as u64, + }) + } + + pub(super) async fn on_update_session_project( + &self, + req: UpdateSessionProjectRequest, + ) -> Result { + let project_id = req.project_id; + self.update_thread_metadata(&req.session_id, move |meta| { + meta.project_id = project_id; + }) + .await?; + Ok(EmptyResponse {}) + } + + pub(super) async fn on_rename_session( + &self, + req: RenameSessionRequest, + ) -> Result { + self.thread_manager + .update_thread(&req.session_id, Some(req.title), Some(true), None) + .await + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + Ok(EmptyResponse {}) + } + + pub(super) async fn on_archive_session( + &self, + req: ArchiveSessionRequest, + ) -> Result { + self.thread_manager + .archive_thread(&req.session_id) + .await + .internal_err()?; + self.sessions.lock().await.remove(&req.session_id); + Ok(EmptyResponse {}) + } + + pub(super) async fn on_unarchive_session( + &self, + req: UnarchiveSessionRequest, + ) -> Result { + self.thread_manager + .unarchive_thread(&req.session_id) + .await + .internal_err()?; + Ok(EmptyResponse {}) + } +} diff --git a/crates/goose/src/acp/server/sources.rs b/crates/goose/src/acp/server/sources.rs new file mode 100644 index 00000000..9dbfd738 --- /dev/null +++ b/crates/goose/src/acp/server/sources.rs @@ -0,0 +1,65 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_create_source( + &self, + req: CreateSourceRequest, + ) -> Result { + let source = crate::sources::create_source( + req.source_type, + &req.name, + &req.description, + &req.content, + req.global, + req.project_dir.as_deref(), + )?; + Ok(CreateSourceResponse { source }) + } + + pub(super) async fn on_list_sources( + &self, + req: ListSourcesRequest, + ) -> Result { + let sources = crate::sources::list_sources(req.source_type, req.project_dir.as_deref())?; + Ok(ListSourcesResponse { sources }) + } + + pub(super) async fn on_update_source( + &self, + req: UpdateSourceRequest, + ) -> Result { + let source = crate::sources::update_source( + req.source_type, + &req.path, + &req.name, + &req.description, + &req.content, + )?; + Ok(UpdateSourceResponse { source }) + } + + pub(super) async fn on_delete_source( + &self, + req: DeleteSourceRequest, + ) -> Result { + crate::sources::delete_source(req.source_type, &req.path)?; + Ok(EmptyResponse {}) + } + + pub(super) async fn on_export_source( + &self, + req: ExportSourceRequest, + ) -> Result { + let (json, filename) = crate::sources::export_source(req.source_type, &req.path)?; + Ok(ExportSourceResponse { json, filename }) + } + + pub(super) async fn on_import_sources( + &self, + req: ImportSourcesRequest, + ) -> Result { + let sources = + crate::sources::import_sources(&req.data, req.global, req.project_dir.as_deref())?; + Ok(ImportSourcesResponse { sources }) + } +} diff --git a/crates/goose/src/acp/server/tools.rs b/crates/goose/src/acp/server/tools.rs new file mode 100644 index 00000000..e91d4783 --- /dev/null +++ b/crates/goose/src/acp/server/tools.rs @@ -0,0 +1,18 @@ +use super::*; + +impl GooseAcpAgent { + pub(super) async fn on_get_tools( + &self, + req: GetToolsRequest, + ) -> Result { + let internal_id = self.internal_session_id(&req.session_id).await?; + let agent = self.get_session_agent(&req.session_id, None).await?; + let tools = agent.list_tools(&internal_id, None).await; + let tools_json = tools + .into_iter() + .map(|t| serde_json::to_value(&t)) + .collect::, _>>() + .internal_err()?; + Ok(GetToolsResponse { tools: tools_json }) + } +}