From ef73610cc1620b8675ada2709ccd44d0538fb417 Mon Sep 17 00:00:00 2001 From: Kalvin C Date: Thu, 30 Apr 2026 14:05:10 -0700 Subject: [PATCH] feat: goose2 add support for custom providers in ui & acp (#8924) --- crates/goose-cli/src/commands/configure.rs | 2 +- crates/goose-sdk/src/custom_requests.rs | 201 ++++++ .../src/routes/config_management.rs | 9 +- crates/goose/acp-meta.json | 30 + crates/goose/acp-schema.json | 637 +++++++++++++++++- .../goose/src/acp/server/custom_dispatch.rs | 48 ++ crates/goose/src/acp/server/providers.rs | 383 +++++++++++ .../goose/src/config/declarative_providers.rs | 189 +++++- .../tests/acp_custom_provider_methods_test.rs | 604 +++++++++++++++++ .../useAgentModelPickerState.test.ts | 51 ++ .../chat/hooks/useAgentModelPickerState.ts | 8 +- .../chat/hooks/useChatSessionController.ts | 22 +- ui/goose2/src/features/chat/types.ts | 3 +- .../src/features/chat/ui/AgentModelPicker.tsx | 299 +------- .../features/chat/ui/AgentModelPickerItem.tsx | 33 + .../chat/ui/AgentModelPickerLists.tsx | 285 ++++++++ ui/goose2/src/features/chat/ui/ChatInput.tsx | 2 + .../src/features/chat/ui/ChatInputToolbar.tsx | 5 +- ui/goose2/src/features/chat/ui/ChatView.tsx | 1 + .../ui/__tests__/AgentModelPicker.test.tsx | 103 ++- .../providers/api/customProviders.test.ts | 158 +++++ .../features/providers/api/customProviders.ts | 65 ++ .../hooks/useCustomProviders.test.tsx | 192 ++++++ .../providers/hooks/useCustomProviders.ts | 318 +++++++++ .../hooks/useProviderInventory.test.ts | 119 ++++ .../providers/hooks/useProviderInventory.ts | 25 +- .../providers/lib/customProviderDraft.test.ts | 208 ++++++ .../providers/lib/customProviderDraft.ts | 151 +++++ .../providers/lib/customProviderHeaders.ts | 120 ++++ .../providers/lib/customProviderModels.ts | 23 + .../providers/lib/customProviderTypes.ts | 58 ++ .../providers/lib/customProviderValidation.ts | 129 ++++ .../providers/providerCatalog.test.ts | 50 ++ .../src/features/providers/providerCatalog.ts | 433 +----------- .../providers/providerCatalogEntries.ts | 481 +++++++++++++ .../providers/ui/CustomHeadersEditor.tsx | 104 +++ .../providers/ui/CustomProviderChoice.tsx | 71 ++ .../providers/ui/CustomProviderDialog.tsx | 306 +++++++++ .../providers/ui/CustomProviderForm.tsx | 384 +++++++++++ .../providers/ui/ProviderModelListEditor.tsx | 103 +++ .../providers/ui/ProviderTemplatePicker.tsx | 133 ++++ .../settings/ui/AgentProviderCard.tsx | 12 +- .../features/settings/ui/ModelProviderRow.tsx | 12 +- .../settings/ui/ProvidersSettings.tsx | 319 ++++++++- .../features/settings/ui/SettingsModal.tsx | 5 +- .../ui/__tests__/ProvidersSettings.test.tsx | 150 +++++ .../settings/ui/customProviderFormAdapters.ts | 106 +++ .../src/shared/i18n/locales/en/settings.json | 88 +++ .../src/shared/i18n/locales/es/settings.json | 88 +++ ui/goose2/src/shared/ui/alert-dialog.tsx | 20 +- ui/goose2/src/shared/ui/dialog.tsx | 36 +- .../src/shared/ui/icons/ProviderIcons.tsx | 12 +- ui/goose2/src/shared/ui/select.tsx | 2 +- ui/sdk/src/generated/client.gen.ts | 88 +++ ui/sdk/src/generated/index.ts | 32 +- ui/sdk/src/generated/types.gen.ts | 175 ++++- ui/sdk/src/generated/zod.gen.ts | 210 +++++- 57 files changed, 7047 insertions(+), 854 deletions(-) create mode 100644 crates/goose/tests/acp_custom_provider_methods_test.rs create mode 100644 ui/goose2/src/features/chat/ui/AgentModelPickerItem.tsx create mode 100644 ui/goose2/src/features/chat/ui/AgentModelPickerLists.tsx create mode 100644 ui/goose2/src/features/providers/api/customProviders.test.ts create mode 100644 ui/goose2/src/features/providers/api/customProviders.ts create mode 100644 ui/goose2/src/features/providers/hooks/useCustomProviders.test.tsx create mode 100644 ui/goose2/src/features/providers/hooks/useCustomProviders.ts create mode 100644 ui/goose2/src/features/providers/hooks/useProviderInventory.test.ts create mode 100644 ui/goose2/src/features/providers/lib/customProviderDraft.test.ts create mode 100644 ui/goose2/src/features/providers/lib/customProviderDraft.ts create mode 100644 ui/goose2/src/features/providers/lib/customProviderHeaders.ts create mode 100644 ui/goose2/src/features/providers/lib/customProviderModels.ts create mode 100644 ui/goose2/src/features/providers/lib/customProviderTypes.ts create mode 100644 ui/goose2/src/features/providers/lib/customProviderValidation.ts create mode 100644 ui/goose2/src/features/providers/providerCatalogEntries.ts create mode 100644 ui/goose2/src/features/providers/ui/CustomHeadersEditor.tsx create mode 100644 ui/goose2/src/features/providers/ui/CustomProviderChoice.tsx create mode 100644 ui/goose2/src/features/providers/ui/CustomProviderDialog.tsx create mode 100644 ui/goose2/src/features/providers/ui/CustomProviderForm.tsx create mode 100644 ui/goose2/src/features/providers/ui/ProviderModelListEditor.tsx create mode 100644 ui/goose2/src/features/providers/ui/ProviderTemplatePicker.tsx create mode 100644 ui/goose2/src/features/settings/ui/customProviderFormAdapters.ts diff --git a/crates/goose-cli/src/commands/configure.rs b/crates/goose-cli/src/commands/configure.rs index ab54441f..c6827242 100644 --- a/crates/goose-cli/src/commands/configure.rs +++ b/crates/goose-cli/src/commands/configure.rs @@ -2093,7 +2093,7 @@ fn add_provider() -> anyhow::Result<()> { engine: provider_type.to_string(), display_name: display_name.clone(), api_url, - api_key, + api_key: requires_auth.then_some(api_key), models, supports_streaming: Some(supports_streaming), headers, diff --git a/crates/goose-sdk/src/custom_requests.rs b/crates/goose-sdk/src/custom_requests.rs index a3fc9e89..96ddf465 100644 --- a/crates/goose-sdk/src/custom_requests.rs +++ b/crates/goose-sdk/src/custom_requests.rs @@ -384,6 +384,207 @@ pub struct ProviderConfigChangeResponse { pub refresh: RefreshProviderInventoryResponse, } +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct ProviderCatalogEntryDto { + pub provider_id: String, + pub name: String, + pub format: String, + pub api_url: String, + pub model_count: usize, + pub doc_url: String, + pub env_var: String, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct ProviderTemplateCapabilitiesDto { + pub tool_call: bool, + pub reasoning: bool, + pub attachment: bool, + pub temperature: bool, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct ProviderTemplateModelDto { + pub id: String, + pub name: String, + pub context_limit: usize, + pub capabilities: ProviderTemplateCapabilitiesDto, + pub deprecated: bool, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct ProviderTemplateDto { + pub provider_id: String, + pub name: String, + pub format: String, + pub api_url: String, + pub models: Vec, + pub supports_streaming: bool, + pub env_var: String, + pub doc_url: String, +} + +/// List custom-provider catalog entries. Omit `format` to list all formats. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/providers/catalog/list", + response = ProviderCatalogListResponse +)] +#[serde(rename_all = "camelCase")] +pub struct ProviderCatalogListRequest { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub format: Option, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct ProviderCatalogListResponse { + pub providers: Vec, +} + +/// Return the editable template for one catalog provider. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/providers/catalog/template", + response = ProviderCatalogTemplateResponse +)] +#[serde(rename_all = "camelCase")] +pub struct ProviderCatalogTemplateRequest { + pub provider_id: String, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct ProviderCatalogTemplateResponse { + pub template: ProviderTemplateDto, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderConfigDto { + pub provider_id: String, + pub engine: String, + pub display_name: String, + pub api_url: String, + #[serde(default)] + pub models: Vec, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub supports_streaming: Option, + #[serde(default)] + pub headers: HashMap, + pub requires_auth: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub catalog_provider_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_path: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub api_key_env: Option, + pub api_key_set: bool, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderUpsertDto { + pub engine: String, + pub display_name: String, + pub api_url: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub api_key: Option, + #[serde(default)] + pub models: Vec, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub supports_streaming: Option, + #[serde(default)] + pub headers: HashMap, + pub requires_auth: bool, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub catalog_provider_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_path: Option, +} + +/// Create a custom provider backed by Goose's declarative provider store. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/providers/custom/create", + response = CustomProviderCreateResponse +)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderCreateRequest { + #[serde(flatten)] + pub provider: CustomProviderUpsertDto, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderCreateResponse { + pub provider_id: String, + pub status: ProviderConfigStatusDto, + pub refresh: RefreshProviderInventoryResponse, +} + +/// Read a declarative provider config. Custom configs are editable; bundled configs are read-only. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/providers/custom/read", + response = CustomProviderReadResponse +)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderReadRequest { + pub provider_id: String, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderReadResponse { + pub provider: CustomProviderConfigDto, + pub editable: bool, + pub status: ProviderConfigStatusDto, +} + +/// Update a custom provider backed by Goose's declarative provider store. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/providers/custom/update", + response = CustomProviderUpdateResponse +)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderUpdateRequest { + pub provider_id: String, + #[serde(flatten)] + pub provider: CustomProviderUpsertDto, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderUpdateResponse { + pub provider_id: String, + pub status: ProviderConfigStatusDto, + pub refresh: RefreshProviderInventoryResponse, +} + +/// Delete a custom provider from Goose's declarative provider store. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/providers/custom/delete", + response = CustomProviderDeleteResponse +)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderDeleteRequest { + pub provider_id: String, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +#[serde(rename_all = "camelCase")] +pub struct CustomProviderDeleteResponse { + pub provider_id: String, + pub refresh: RefreshProviderInventoryResponse, +} + /// The type of source entity. #[derive( Debug, Default, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, JsonSchema, diff --git a/crates/goose-server/src/routes/config_management.rs b/crates/goose-server/src/routes/config_management.rs index 6047f88b..1f93f5e2 100644 --- a/crates/goose-server/src/routes/config_management.rs +++ b/crates/goose-server/src/routes/config_management.rs @@ -108,6 +108,11 @@ fn default_requires_auth() -> bool { true } +fn normalize_custom_provider_api_key(api_key: String) -> Option { + let api_key = api_key.trim().to_string(); + (!api_key.is_empty()).then_some(api_key) +} + #[derive(Deserialize, ToSchema)] pub struct CheckProviderRequest { pub provider: String, @@ -583,7 +588,7 @@ pub async fn create_custom_provider( engine: request.engine, display_name: request.display_name, api_url: request.api_url, - api_key: request.api_key, + api_key: normalize_custom_provider_api_key(request.api_key), models: request.models, supports_streaming: request.supports_streaming, headers: request.headers, @@ -675,7 +680,7 @@ pub async fn update_custom_provider( engine: request.engine, display_name: request.display_name, api_url: request.api_url, - api_key: request.api_key, + api_key: normalize_custom_provider_api_key(request.api_key), models: request.models, supports_streaming: request.supports_streaming, headers: request.headers, diff --git a/crates/goose/acp-meta.json b/crates/goose/acp-meta.json index 1d13987b..ffce9ab5 100644 --- a/crates/goose/acp-meta.json +++ b/crates/goose/acp-meta.json @@ -60,6 +60,36 @@ "requestType": "ListProvidersRequest", "responseType": "ListProvidersResponse" }, + { + "method": "_goose/providers/catalog/list", + "requestType": "ProviderCatalogListRequest", + "responseType": "ProviderCatalogListResponse" + }, + { + "method": "_goose/providers/catalog/template", + "requestType": "ProviderCatalogTemplateRequest", + "responseType": "ProviderCatalogTemplateResponse" + }, + { + "method": "_goose/providers/custom/create", + "requestType": "CustomProviderCreateRequest", + "responseType": "CustomProviderCreateResponse" + }, + { + "method": "_goose/providers/custom/read", + "requestType": "CustomProviderReadRequest", + "responseType": "CustomProviderReadResponse" + }, + { + "method": "_goose/providers/custom/update", + "requestType": "CustomProviderUpdateRequest", + "responseType": "CustomProviderUpdateResponse" + }, + { + "method": "_goose/providers/custom/delete", + "requestType": "CustomProviderDeleteRequest", + "responseType": "CustomProviderDeleteResponse" + }, { "method": "_goose/providers/inventory/refresh", "requestType": "RefreshProviderInventoryRequest", diff --git a/crates/goose/acp-schema.json b/crates/goose/acp-schema.json index be922800..1c27c87a 100644 --- a/crates/goose/acp-schema.json +++ b/crates/goose/acp-schema.json @@ -471,21 +471,291 @@ ], "description": "A single model in provider inventory." }, - "RefreshProviderInventoryRequest": { + "ProviderCatalogListRequest": { "type": "object", "properties": { - "providerIds": { + "format": { + "type": [ + "string", + "null" + ] + } + }, + "description": "List custom-provider catalog entries. Omit `format` to list all formats.", + "x-side": "agent", + "x-method": "_goose/providers/catalog/list" + }, + "ProviderCatalogListResponse": { + "type": "object", + "properties": { + "providers": { + "type": "array", + "items": { + "$ref": "#/$defs/ProviderCatalogEntryDto" + } + } + }, + "required": [ + "providers" + ], + "x-side": "agent", + "x-method": "_goose/providers/catalog/list" + }, + "ProviderCatalogEntryDto": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "name": { + "type": "string" + }, + "format": { + "type": "string" + }, + "apiUrl": { + "type": "string" + }, + "modelCount": { + "type": "integer", + "minimum": 0 + }, + "docUrl": { + "type": "string" + }, + "envVar": { + "type": "string" + } + }, + "required": [ + "providerId", + "name", + "format", + "apiUrl", + "modelCount", + "docUrl", + "envVar" + ] + }, + "ProviderCatalogTemplateRequest": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + } + }, + "required": [ + "providerId" + ], + "description": "Return the editable template for one catalog provider.", + "x-side": "agent", + "x-method": "_goose/providers/catalog/template" + }, + "ProviderCatalogTemplateResponse": { + "type": "object", + "properties": { + "template": { + "$ref": "#/$defs/ProviderTemplateDto" + } + }, + "required": [ + "template" + ], + "x-side": "agent", + "x-method": "_goose/providers/catalog/template" + }, + "ProviderTemplateDto": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "name": { + "type": "string" + }, + "format": { + "type": "string" + }, + "apiUrl": { + "type": "string" + }, + "models": { + "type": "array", + "items": { + "$ref": "#/$defs/ProviderTemplateModelDto" + } + }, + "supportsStreaming": { + "type": "boolean" + }, + "envVar": { + "type": "string" + }, + "docUrl": { + "type": "string" + } + }, + "required": [ + "providerId", + "name", + "format", + "apiUrl", + "models", + "supportsStreaming", + "envVar", + "docUrl" + ] + }, + "ProviderTemplateModelDto": { + "type": "object", + "properties": { + "id": { + "type": "string" + }, + "name": { + "type": "string" + }, + "contextLimit": { + "type": "integer", + "minimum": 0 + }, + "capabilities": { + "$ref": "#/$defs/ProviderTemplateCapabilitiesDto" + }, + "deprecated": { + "type": "boolean" + } + }, + "required": [ + "id", + "name", + "contextLimit", + "capabilities", + "deprecated" + ] + }, + "ProviderTemplateCapabilitiesDto": { + "type": "object", + "properties": { + "toolCall": { + "type": "boolean" + }, + "reasoning": { + "type": "boolean" + }, + "attachment": { + "type": "boolean" + }, + "temperature": { + "type": "boolean" + } + }, + "required": [ + "toolCall", + "reasoning", + "attachment", + "temperature" + ] + }, + "CustomProviderCreateRequest": { + "type": "object", + "properties": { + "engine": { + "type": "string" + }, + "displayName": { + "type": "string" + }, + "apiUrl": { + "type": "string" + }, + "apiKey": { + "type": [ + "string", + "null" + ] + }, + "models": { "type": "array", "items": { "type": "string" }, - "description": "Which providers to refresh. Empty means all known providers.", "default": [] + }, + "supportsStreaming": { + "type": [ + "boolean", + "null" + ] + }, + "headers": { + "type": "object", + "additionalProperties": { + "type": "string" + }, + "default": {} + }, + "requiresAuth": { + "type": "boolean" + }, + "catalogProviderId": { + "type": [ + "string", + "null" + ] + }, + "basePath": { + "type": [ + "string", + "null" + ] } }, - "description": "Trigger a background refresh of provider inventories.", + "required": [ + "engine", + "displayName", + "apiUrl", + "requiresAuth" + ], + "description": "Create a custom provider backed by Goose's declarative provider store.", "x-side": "agent", - "x-method": "_goose/providers/inventory/refresh" + "x-method": "_goose/providers/custom/create" + }, + "CustomProviderCreateResponse": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "status": { + "$ref": "#/$defs/ProviderConfigStatusDto" + }, + "refresh": { + "$ref": "#/$defs/RefreshProviderInventoryResponse" + } + }, + "required": [ + "providerId", + "status", + "refresh" + ], + "x-side": "agent", + "x-method": "_goose/providers/custom/create" + }, + "ProviderConfigStatusDto": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "isConfigured": { + "type": "boolean" + } + }, + "required": [ + "providerId", + "isConfigured" + ] }, "RefreshProviderInventoryResponse": { "type": "object", @@ -537,6 +807,246 @@ "already_refreshing" ] }, + "CustomProviderReadRequest": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + } + }, + "required": [ + "providerId" + ], + "description": "Read a declarative provider config. Custom configs are editable; bundled configs are read-only.", + "x-side": "agent", + "x-method": "_goose/providers/custom/read" + }, + "CustomProviderReadResponse": { + "type": "object", + "properties": { + "provider": { + "$ref": "#/$defs/CustomProviderConfigDto" + }, + "editable": { + "type": "boolean" + }, + "status": { + "$ref": "#/$defs/ProviderConfigStatusDto" + } + }, + "required": [ + "provider", + "editable", + "status" + ], + "x-side": "agent", + "x-method": "_goose/providers/custom/read" + }, + "CustomProviderConfigDto": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "engine": { + "type": "string" + }, + "displayName": { + "type": "string" + }, + "apiUrl": { + "type": "string" + }, + "models": { + "type": "array", + "items": { + "type": "string" + }, + "default": [] + }, + "supportsStreaming": { + "type": [ + "boolean", + "null" + ] + }, + "headers": { + "type": "object", + "additionalProperties": { + "type": "string" + }, + "default": {} + }, + "requiresAuth": { + "type": "boolean" + }, + "catalogProviderId": { + "type": [ + "string", + "null" + ] + }, + "basePath": { + "type": [ + "string", + "null" + ] + }, + "apiKeyEnv": { + "type": [ + "string", + "null" + ] + }, + "apiKeySet": { + "type": "boolean" + } + }, + "required": [ + "providerId", + "engine", + "displayName", + "apiUrl", + "requiresAuth", + "apiKeySet" + ] + }, + "CustomProviderUpdateRequest": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "engine": { + "type": "string" + }, + "displayName": { + "type": "string" + }, + "apiUrl": { + "type": "string" + }, + "apiKey": { + "type": [ + "string", + "null" + ] + }, + "models": { + "type": "array", + "items": { + "type": "string" + }, + "default": [] + }, + "supportsStreaming": { + "type": [ + "boolean", + "null" + ] + }, + "headers": { + "type": "object", + "additionalProperties": { + "type": "string" + }, + "default": {} + }, + "requiresAuth": { + "type": "boolean" + }, + "catalogProviderId": { + "type": [ + "string", + "null" + ] + }, + "basePath": { + "type": [ + "string", + "null" + ] + } + }, + "required": [ + "providerId", + "engine", + "displayName", + "apiUrl", + "requiresAuth" + ], + "description": "Update a custom provider backed by Goose's declarative provider store.", + "x-side": "agent", + "x-method": "_goose/providers/custom/update" + }, + "CustomProviderUpdateResponse": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "status": { + "$ref": "#/$defs/ProviderConfigStatusDto" + }, + "refresh": { + "$ref": "#/$defs/RefreshProviderInventoryResponse" + } + }, + "required": [ + "providerId", + "status", + "refresh" + ], + "x-side": "agent", + "x-method": "_goose/providers/custom/update" + }, + "CustomProviderDeleteRequest": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + } + }, + "required": [ + "providerId" + ], + "description": "Delete a custom provider from Goose's declarative provider store.", + "x-side": "agent", + "x-method": "_goose/providers/custom/delete" + }, + "CustomProviderDeleteResponse": { + "type": "object", + "properties": { + "providerId": { + "type": "string" + }, + "refresh": { + "$ref": "#/$defs/RefreshProviderInventoryResponse" + } + }, + "required": [ + "providerId", + "refresh" + ], + "x-side": "agent", + "x-method": "_goose/providers/custom/delete" + }, + "RefreshProviderInventoryRequest": { + "type": "object", + "properties": { + "providerIds": { + "type": "array", + "items": { + "type": "string" + }, + "description": "Which providers to refresh. Empty means all known providers.", + "default": [] + } + }, + "description": "Trigger a background refresh of provider inventories.", + "x-side": "agent", + "x-method": "_goose/providers/inventory/refresh" + }, "ProviderConfigReadRequest": { "type": "object", "properties": { @@ -628,21 +1138,6 @@ "x-side": "agent", "x-method": "_goose/providers/config/status" }, - "ProviderConfigStatusDto": { - "type": "object", - "properties": { - "providerId": { - "type": "string" - }, - "isConfigured": { - "type": "boolean" - } - }, - "required": [ - "providerId", - "isConfigured" - ] - }, "ProviderConfigSaveRequest": { "type": "object", "properties": { @@ -1682,6 +2177,60 @@ "description": "Params for _goose/providers/list", "title": "ListProvidersRequest" }, + { + "allOf": [ + { + "$ref": "#/$defs/ProviderCatalogListRequest" + } + ], + "description": "Params for _goose/providers/catalog/list", + "title": "ProviderCatalogListRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/ProviderCatalogTemplateRequest" + } + ], + "description": "Params for _goose/providers/catalog/template", + "title": "ProviderCatalogTemplateRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderCreateRequest" + } + ], + "description": "Params for _goose/providers/custom/create", + "title": "CustomProviderCreateRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderReadRequest" + } + ], + "description": "Params for _goose/providers/custom/read", + "title": "CustomProviderReadRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderUpdateRequest" + } + ], + "description": "Params for _goose/providers/custom/update", + "title": "CustomProviderUpdateRequest" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderDeleteRequest" + } + ], + "description": "Params for _goose/providers/custom/delete", + "title": "CustomProviderDeleteRequest" + }, { "allOf": [ { @@ -2039,6 +2588,54 @@ ], "title": "ListProvidersResponse" }, + { + "allOf": [ + { + "$ref": "#/$defs/ProviderCatalogListResponse" + } + ], + "title": "ProviderCatalogListResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/ProviderCatalogTemplateResponse" + } + ], + "title": "ProviderCatalogTemplateResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderCreateResponse" + } + ], + "title": "CustomProviderCreateResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderReadResponse" + } + ], + "title": "CustomProviderReadResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderUpdateResponse" + } + ], + "title": "CustomProviderUpdateResponse" + }, + { + "allOf": [ + { + "$ref": "#/$defs/CustomProviderDeleteResponse" + } + ], + "title": "CustomProviderDeleteResponse" + }, { "allOf": [ { diff --git a/crates/goose/src/acp/server/custom_dispatch.rs b/crates/goose/src/acp/server/custom_dispatch.rs index 19a49843..b5e0462c 100644 --- a/crates/goose/src/acp/server/custom_dispatch.rs +++ b/crates/goose/src/acp/server/custom_dispatch.rs @@ -104,6 +104,54 @@ impl GooseAcpAgent { self.on_list_providers(req).await } + #[custom_method(ProviderCatalogListRequest)] + async fn dispatch_list_provider_catalog( + &self, + req: ProviderCatalogListRequest, + ) -> Result { + self.on_list_provider_catalog(req).await + } + + #[custom_method(ProviderCatalogTemplateRequest)] + async fn dispatch_get_provider_catalog_template( + &self, + req: ProviderCatalogTemplateRequest, + ) -> Result { + self.on_get_provider_catalog_template(req).await + } + + #[custom_method(CustomProviderCreateRequest)] + async fn dispatch_create_custom_provider( + &self, + req: CustomProviderCreateRequest, + ) -> Result { + self.on_create_custom_provider(req).await + } + + #[custom_method(CustomProviderReadRequest)] + async fn dispatch_read_custom_provider( + &self, + req: CustomProviderReadRequest, + ) -> Result { + self.on_read_custom_provider(req).await + } + + #[custom_method(CustomProviderUpdateRequest)] + async fn dispatch_update_custom_provider( + &self, + req: CustomProviderUpdateRequest, + ) -> Result { + self.on_update_custom_provider(req).await + } + + #[custom_method(CustomProviderDeleteRequest)] + async fn dispatch_delete_custom_provider( + &self, + req: CustomProviderDeleteRequest, + ) -> Result { + self.on_delete_custom_provider(req).await + } + #[custom_method(RefreshProviderInventoryRequest)] async fn dispatch_refresh_provider_inventory( &self, diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs index 7297e1c7..e61c8a4f 100644 --- a/crates/goose/src/acp/server/providers.rs +++ b/crates/goose/src/acp/server/providers.rs @@ -1,4 +1,6 @@ use super::*; +use crate::config::declarative_providers; +use std::str::FromStr; fn inventory_entry_to_dto(entry: ProviderInventoryEntry) -> ProviderInventoryEntryDto { let stale = ProviderInventoryService::is_stale(&entry); @@ -110,6 +112,197 @@ fn provider_config_field_value( } } +fn provider_catalog_entry_to_dto( + entry: crate::providers::catalog::ProviderCatalogEntry, +) -> ProviderCatalogEntryDto { + ProviderCatalogEntryDto { + provider_id: entry.id, + name: entry.name, + format: entry.format, + api_url: entry.api_url, + model_count: entry.model_count, + doc_url: entry.doc_url, + env_var: entry.env_var, + } +} + +fn provider_template_to_dto( + template: crate::providers::catalog::ProviderTemplate, +) -> ProviderTemplateDto { + ProviderTemplateDto { + provider_id: template.id, + name: template.name, + format: template.format, + api_url: template.api_url, + models: template + .models + .into_iter() + .map(|model| ProviderTemplateModelDto { + id: model.id, + name: model.name, + context_limit: model.context_limit, + capabilities: ProviderTemplateCapabilitiesDto { + tool_call: model.capabilities.tool_call, + reasoning: model.capabilities.reasoning, + attachment: model.capabilities.attachment, + temperature: model.capabilities.temperature, + }, + deprecated: model.deprecated, + }) + .collect(), + supports_streaming: template.supports_streaming, + env_var: template.env_var, + doc_url: template.doc_url, + } +} + +fn custom_provider_engine_to_dto(engine: &declarative_providers::ProviderEngine) -> &'static str { + match engine { + declarative_providers::ProviderEngine::OpenAI => "openai_compatible", + declarative_providers::ProviderEngine::Anthropic => "anthropic_compatible", + declarative_providers::ProviderEngine::Ollama => "ollama_compatible", + } +} + +fn normalize_custom_provider_engine(engine: &str) -> Result { + let engine = engine.trim().to_lowercase(); + if declarative_providers::ProviderEngine::from_str(&engine).is_err() { + return Err(sacp::Error::invalid_params() + .data(format!("Unsupported custom provider engine: {engine}"))); + } + + match engine.as_str() { + "openai" | "openai_compatible" => Ok("openai_compatible".to_string()), + "anthropic" | "anthropic_compatible" => Ok("anthropic_compatible".to_string()), + "ollama" | "ollama_compatible" => Ok("ollama_compatible".to_string()), + _ => unreachable!("provider engine was validated above"), + } +} + +fn non_empty_trimmed(value: String, field: &str) -> Result { + let value = value.trim().to_string(); + if value.is_empty() { + return Err(sacp::Error::invalid_params().data(format!("{field} cannot be empty"))); + } + Ok(value) +} + +fn normalize_optional_string(value: Option) -> Option { + value.and_then(|value| { + let value = value.trim().to_string(); + (!value.is_empty()).then_some(value) + }) +} + +fn normalize_custom_provider_upsert( + mut provider: CustomProviderUpsertDto, + require_api_key: bool, +) -> Result { + provider.engine = normalize_custom_provider_engine(&provider.engine)?; + provider.display_name = non_empty_trimmed(provider.display_name, "displayName")?; + provider.api_url = non_empty_trimmed(provider.api_url, "apiUrl")?; + let url = url::Url::parse(&provider.api_url) + .map_err(|_| sacp::Error::invalid_params().data("apiUrl must be a valid URL"))?; + if !matches!(url.scheme(), "http" | "https") { + return Err(sacp::Error::invalid_params().data("apiUrl must use HTTP or HTTPS")); + } + + provider.api_key = provider.api_key.and_then(|api_key| { + let api_key = api_key.trim().to_string(); + (!api_key.is_empty()).then_some(api_key) + }); + if require_api_key && provider.requires_auth && provider.api_key.is_none() { + return Err(sacp::Error::invalid_params().data("apiKey cannot be empty")); + } + provider.models = provider + .models + .into_iter() + .filter_map(|model| { + let model = model.trim().to_string(); + (!model.is_empty()).then_some(model) + }) + .collect(); + if provider.models.is_empty() { + return Err(sacp::Error::invalid_params().data("models cannot be empty")); + } + + provider.headers = provider + .headers + .into_iter() + .map(|(key, value)| { + let key = key.trim().to_string(); + let value = value.trim().to_string(); + if key.is_empty() { + return Ok(None); + } + reqwest::header::HeaderName::from_bytes(key.as_bytes()).map_err(|_| { + sacp::Error::invalid_params().data(format!("Invalid header name: {key}")) + })?; + reqwest::header::HeaderValue::from_str(&value).map_err(|_| { + sacp::Error::invalid_params().data(format!("Invalid header value for: {key}")) + })?; + Ok(Some((key, value))) + }) + .collect::, sacp::Error>>()? + .into_iter() + .flatten() + .collect(); + provider.catalog_provider_id = normalize_optional_string(provider.catalog_provider_id); + provider.base_path = normalize_optional_string(provider.base_path); + Ok(provider) +} + +fn custom_provider_headers(headers: HashMap) -> Option> { + (!headers.is_empty()).then_some(headers) +} + +fn load_declarative_provider_for_client( + provider_id: &str, +) -> Result { + declarative_providers::load_provider(provider_id).map_err(|error| { + if error.to_string().contains("Provider not found") { + sacp::Error::invalid_params().data(format!("Unknown provider: {provider_id}")) + } else if error.to_string().contains("Invalid provider id") { + sacp::Error::invalid_params().data(error.to_string()) + } else { + sacp::Error::internal_error().data(error.to_string()) + } + }) +} + +fn custom_provider_config_to_dto( + config: &declarative_providers::DeclarativeProviderConfig, +) -> CustomProviderConfigDto { + let api_key_env = normalize_optional_string(Some(config.api_key_env.clone())); + let api_key_set = api_key_env + .as_ref() + .map(|key| { + Config::global() + .get_secret::(key) + .is_ok() + }) + .unwrap_or(false); + + CustomProviderConfigDto { + provider_id: config.name.clone(), + engine: custom_provider_engine_to_dto(&config.engine).to_string(), + display_name: config.display_name.clone(), + api_url: config.base_url.clone(), + models: config + .models + .iter() + .map(|model| model.name.clone()) + .collect(), + supports_streaming: config.supports_streaming, + headers: config.headers.clone().unwrap_or_default(), + requires_auth: config.requires_auth, + catalog_provider_id: config.catalog_provider_id.clone(), + base_path: config.base_path.clone(), + api_key_env, + api_key_set, + } +} + fn refresh_skip_reason_to_dto(reason: RefreshSkipReason) -> RefreshProviderInventorySkipReasonDto { match reason { RefreshSkipReason::UnknownProvider => { @@ -154,6 +347,196 @@ impl GooseAcpAgent { }) } + pub(super) async fn on_list_provider_catalog( + &self, + req: ProviderCatalogListRequest, + ) -> Result { + let formats = match req.format { + Some(format) => vec![format + .parse::() + .map_err(|error| sacp::Error::invalid_params().data(error))?], + None => vec![ + crate::providers::catalog::ProviderFormat::OpenAI, + crate::providers::catalog::ProviderFormat::Anthropic, + crate::providers::catalog::ProviderFormat::Ollama, + ], + }; + + let mut providers = Vec::new(); + for format in formats { + providers.extend( + crate::providers::catalog::get_providers_by_format(format) + .await + .into_iter() + .map(provider_catalog_entry_to_dto), + ); + } + providers.sort_by(|a, b| { + a.name + .cmp(&b.name) + .then_with(|| a.provider_id.cmp(&b.provider_id)) + }); + + Ok(ProviderCatalogListResponse { providers }) + } + + pub(super) async fn on_get_provider_catalog_template( + &self, + req: ProviderCatalogTemplateRequest, + ) -> Result { + let template = crate::providers::catalog::get_provider_template(&req.provider_id) + .ok_or_else(|| { + sacp::Error::invalid_params() + .data(format!("Unknown catalog provider: {}", req.provider_id)) + })?; + Ok(ProviderCatalogTemplateResponse { + template: provider_template_to_dto(template), + }) + } + + pub(super) async fn on_create_custom_provider( + &self, + req: CustomProviderCreateRequest, + ) -> Result { + let provider = normalize_custom_provider_upsert(req.provider, true)?; + let config = declarative_providers::create_custom_provider( + declarative_providers::CreateCustomProviderParams { + engine: provider.engine, + display_name: provider.display_name, + api_url: provider.api_url, + api_key: provider.api_key, + models: provider.models, + supports_streaming: provider.supports_streaming, + headers: custom_provider_headers(provider.headers), + requires_auth: provider.requires_auth, + catalog_provider_id: provider.catalog_provider_id, + base_path: provider.base_path, + }, + ) + .internal_err_ctx("Failed to create custom provider")?; + + Config::global().invalidate_secrets_cache(); + crate::providers::refresh_custom_providers() + .await + .internal_err_ctx("Failed to refresh custom providers")?; + + let provider_id = config.name; + let provider_ids = [provider_id.clone()]; + let status = Self::provider_config_status(provider_id.clone()).await; + let refresh = self.start_provider_inventory_refresh(&provider_ids).await?; + Ok(CustomProviderCreateResponse { + provider_id, + status, + refresh, + }) + } + + pub(super) async fn on_read_custom_provider( + &self, + req: CustomProviderReadRequest, + ) -> Result { + let loaded = load_declarative_provider_for_client(&req.provider_id)?; + let status = Self::provider_config_status(req.provider_id).await; + Ok(CustomProviderReadResponse { + provider: custom_provider_config_to_dto(&loaded.config), + editable: loaded.is_editable, + status, + }) + } + + pub(super) async fn on_update_custom_provider( + &self, + req: CustomProviderUpdateRequest, + ) -> Result { + let loaded = load_declarative_provider_for_client(&req.provider_id)?; + if !loaded.is_editable { + return Err(sacp::Error::invalid_params() + .data(format!("Provider is not editable: {}", req.provider_id))); + } + + let provider = normalize_custom_provider_upsert(req.provider, false)?; + if provider.requires_auth && provider.api_key.is_none() { + let api_key_env = if loaded.config.api_key_env.is_empty() { + declarative_providers::generate_api_key_name(&req.provider_id) + } else { + loaded.config.api_key_env.clone() + }; + if Config::global().get_secret::(&api_key_env).is_err() { + return Err(sacp::Error::invalid_params() + .data("apiKey is required when auth is enabled and no secret is stored")); + } + } + declarative_providers::update_custom_provider( + declarative_providers::UpdateCustomProviderParams { + id: req.provider_id.clone(), + engine: provider.engine, + display_name: provider.display_name, + api_url: provider.api_url, + api_key: provider.api_key, + models: provider.models, + supports_streaming: provider.supports_streaming, + headers: Some(provider.headers), + requires_auth: provider.requires_auth, + catalog_provider_id: provider.catalog_provider_id, + base_path: provider.base_path, + }, + ) + .internal_err_ctx("Failed to update custom provider")?; + + Config::global().invalidate_secrets_cache(); + crate::providers::refresh_custom_providers() + .await + .internal_err_ctx("Failed to refresh custom providers")?; + + 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(CustomProviderUpdateResponse { + provider_id: req.provider_id, + status, + refresh, + }) + } + + pub(super) async fn on_delete_custom_provider( + &self, + req: CustomProviderDeleteRequest, + ) -> Result { + let loaded = load_declarative_provider_for_client(&req.provider_id)?; + if !loaded.is_editable { + return Err(sacp::Error::invalid_params() + .data(format!("Provider is not editable: {}", req.provider_id))); + } + + if Config::global() + .get_param::("GOOSE_PROVIDER") + .ok() + .as_deref() + == Some(req.provider_id.as_str()) + { + return Err(sacp::Error::invalid_params().data(format!( + "Cannot delete active provider: {}", + req.provider_id + ))); + } + + declarative_providers::remove_custom_provider(&req.provider_id) + .internal_err_ctx("Failed to delete custom provider")?; + + Config::global().invalidate_secrets_cache(); + crate::providers::refresh_custom_providers() + .await + .internal_err_ctx("Failed to refresh custom providers")?; + + Ok(CustomProviderDeleteResponse { + provider_id: req.provider_id, + refresh: RefreshProviderInventoryResponse { + started: Vec::new(), + skipped: Vec::new(), + }, + }) + } + 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) => { diff --git a/crates/goose/src/config/declarative_providers.rs b/crates/goose/src/config/declarative_providers.rs index b8a1dc58..35314401 100644 --- a/crates/goose/src/config/declarative_providers.rs +++ b/crates/goose/src/config/declarative_providers.rs @@ -9,6 +9,7 @@ use anyhow::Result; use include_dir::{include_dir, Dir}; use once_cell::sync::Lazy; use serde::{Deserialize, Deserializer, Serialize}; +use std::str::FromStr; /// Deserialize an optional string, treating empty/whitespace-only values as None. fn deserialize_non_empty_string<'de, D>(deserializer: D) -> Result, D::Error> @@ -19,7 +20,7 @@ where Ok(opt.filter(|s| !s.trim().is_empty())) } use std::collections::HashMap; -use std::path::Path; +use std::path::{Path, PathBuf}; use std::sync::Mutex; use utoipa::ToSchema; @@ -37,6 +38,19 @@ pub enum ProviderEngine { Anthropic, } +impl FromStr for ProviderEngine { + type Err = anyhow::Error; + + fn from_str(engine: &str) -> Result { + match engine.trim().to_lowercase().as_str() { + "openai" | "openai_compatible" => Ok(Self::OpenAI), + "anthropic" | "anthropic_compatible" => Ok(Self::Anthropic), + "ollama" | "ollama_compatible" => Ok(Self::Ollama), + _ => Err(anyhow::anyhow!("Invalid provider type: {}", engine)), + } + } +} + #[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] pub struct EnvVarConfig { pub name: String, @@ -147,7 +161,19 @@ static ID_GENERATION_LOCK: Lazy> = Lazy::new(|| Mutex::new(())); pub fn generate_id(display_name: &str) -> String { let _guard = ID_GENERATION_LOCK.lock().unwrap(); - let normalized = display_name.to_lowercase().replace(' ', "_"); + let normalized = display_name + .to_lowercase() + .chars() + .map(|ch| { + if ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '_' || ch == '-' { + ch + } else { + '_' + } + }) + .collect::() + .trim_matches('_') + .to_string(); let base_id = format!("custom_{}", normalized); let custom_dir = custom_providers_dir(); @@ -162,6 +188,40 @@ pub fn generate_id(display_name: &str) -> String { candidate_id } +pub fn validate_provider_id(id: &str) -> Result<()> { + let mut chars = id.chars(); + let Some(first) = chars.next() else { + return Err(anyhow::anyhow!( + "Invalid provider id: provider id cannot be empty" + )); + }; + + if !(first.is_ascii_lowercase() || first.is_ascii_digit() || first == '_') { + return Err(anyhow::anyhow!("Invalid provider id: {}", id)); + } + + if chars.all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '_' || ch == '-') { + Ok(()) + } else { + Err(anyhow::anyhow!("Invalid provider id: {}", id)) + } +} + +fn custom_provider_file_path(id: &str) -> Result { + if id.is_empty() + || id + .chars() + .any(|ch| ch == '/' || ch == '\\' || ch.is_control()) + { + return Err(anyhow::anyhow!( + "Invalid provider id: {}", + if id.is_empty() { "" } else { id } + )); + } + + Ok(custom_providers_dir().join(format!("{}.json", id))) +} + pub fn generate_api_key_name(id: &str) -> String { format!("{}_API_KEY", id.to_uppercase()) } @@ -171,7 +231,7 @@ pub struct CreateCustomProviderParams { pub engine: String, pub display_name: String, pub api_url: String, - pub api_key: String, + pub api_key: Option, pub models: Vec, pub supports_streaming: Option, pub headers: Option>, @@ -186,7 +246,7 @@ pub struct UpdateCustomProviderParams { pub engine: String, pub display_name: String, pub api_url: String, - pub api_key: String, + pub api_key: Option, pub models: Vec, pub supports_streaming: Option, pub headers: Option>, @@ -199,11 +259,17 @@ pub fn create_custom_provider( params: CreateCustomProviderParams, ) -> Result { let id = generate_id(¶ms.display_name); + validate_provider_id(&id)?; let api_key_env = if params.requires_auth { + let api_key = params + .api_key + .as_deref() + .filter(|api_key| !api_key.trim().is_empty()) + .ok_or_else(|| anyhow::anyhow!("apiKey cannot be empty"))?; let api_key_name = generate_api_key_name(&id); let config = Config::global(); - config.set_secret(&api_key_name, ¶ms.api_key)?; + config.set_secret(&api_key_name, &api_key)?; api_key_name } else { String::new() @@ -217,12 +283,7 @@ pub fn create_custom_provider( let provider_config = DeclarativeProviderConfig { name: id.clone(), - engine: match params.engine.as_str() { - "openai_compatible" => ProviderEngine::OpenAI, - "anthropic_compatible" => ProviderEngine::Anthropic, - "ollama_compatible" => ProviderEngine::Ollama, - _ => return Err(anyhow::anyhow!("Invalid provider type: {}", params.engine)), - }, + engine: ProviderEngine::from_str(¶ms.engine)?, display_name: params.display_name.clone(), description: Some(format!("Custom {} provider", params.display_name)), api_key_env, @@ -258,18 +319,24 @@ pub fn update_custom_provider(params: UpdateCustomProviderParams) -> Result<()> let editable = loaded_provider.is_editable; let config = Config::global(); - let api_key_env = if params.requires_auth { let api_key_name = if existing_config.api_key_env.is_empty() { generate_api_key_name(¶ms.id) } else { existing_config.api_key_env.clone() }; - if !params.api_key.is_empty() { - config.set_secret(&api_key_name, ¶ms.api_key)?; + if let Some(api_key) = params.api_key.as_deref() { + config.set_secret(&api_key_name, &api_key)?; + } else if config.get_secret::(&api_key_name).is_err() { + return Err(anyhow::anyhow!( + "apiKey is required when auth is enabled and no secret is stored" + )); } api_key_name } else { + if existing_config.api_key_env == generate_api_key_name(¶ms.id) { + config.delete_secret(&existing_config.api_key_env)?; + } String::new() }; @@ -282,12 +349,7 @@ pub fn update_custom_provider(params: UpdateCustomProviderParams) -> Result<()> let updated_config = DeclarativeProviderConfig { name: params.id.clone(), - engine: match params.engine.as_str() { - "openai_compatible" => ProviderEngine::OpenAI, - "anthropic_compatible" => ProviderEngine::Anthropic, - "ollama_compatible" => ProviderEngine::Ollama, - _ => return Err(anyhow::anyhow!("Invalid provider type: {}", params.engine)), - }, + engine: ProviderEngine::from_str(¶ms.engine)?, display_name: params.display_name, description: existing_config.description, api_key_env, @@ -311,7 +373,7 @@ pub fn update_custom_provider(params: UpdateCustomProviderParams) -> Result<()> fast_model: existing_config.fast_model.clone(), }; - let file_path = custom_providers_dir().join(format!("{}.json", updated_config.name)); + let file_path = custom_provider_file_path(&updated_config.name)?; let json_content = serde_json::to_string_pretty(&updated_config)?; std::fs::write(file_path, json_content)?; } @@ -320,11 +382,13 @@ pub fn update_custom_provider(params: UpdateCustomProviderParams) -> Result<()> pub fn remove_custom_provider(id: &str) -> Result<()> { let config = Config::global(); - let api_key_name = generate_api_key_name(id); - let _ = config.delete_secret(&api_key_name); + let loaded_provider = load_provider(id)?; + let api_key_env = loaded_provider.config.api_key_env; + if api_key_env == generate_api_key_name(id) { + let _ = config.delete_secret(&api_key_env); + } - let custom_providers_dir = custom_providers_dir(); - let file_path = custom_providers_dir.join(format!("{}.json", id)); + let file_path = custom_provider_file_path(id)?; if file_path.exists() { std::fs::remove_file(file_path)?; @@ -334,7 +398,7 @@ pub fn remove_custom_provider(id: &str) -> Result<()> { } pub fn load_provider(id: &str) -> Result { - let custom_file_path = custom_providers_dir().join(format!("{}.json", id)); + let custom_file_path = custom_provider_file_path(id)?; if custom_file_path.exists() { let content = std::fs::read_to_string(&custom_file_path)?; @@ -624,6 +688,79 @@ mod tests { assert_eq!(config.models[0].context_limit, 131072); } + #[test] + fn test_validate_provider_id_rejects_legacy_punctuation_for_new_ids() { + assert!(validate_provider_id("custom_z.ai").is_err()); + } + + fn write_legacy_provider_config(id: &str, display_name: &str) { + let custom_dir = custom_providers_dir(); + std::fs::create_dir_all(&custom_dir).unwrap(); + let content = format!( + r#"{{ + "name": "{id}", + "engine": "openai", + "display_name": "{display_name}", + "description": "legacy provider", + "api_key_env": "", + "base_url": "https://example.invalid/v1/chat/completions", + "models": [], + "requires_auth": false +}}"# + ); + std::fs::write(custom_dir.join(format!("{id}.json")), content).unwrap(); + } + + #[test] + fn test_load_provider_allows_legacy_custom_id_with_punctuation() { + let temp_dir = tempfile::tempdir().unwrap(); + let temp_root = temp_dir.path().display().to_string(); + let _guard = env_lock::lock_env([("GOOSE_PATH_ROOT", Some(temp_root.as_str()))]); + + write_legacy_provider_config("custom_z.ai", "Z.AI"); + + let loaded = load_provider("custom_z.ai").unwrap(); + assert!(loaded.is_editable); + assert_eq!(loaded.config.name, "custom_z.ai"); + } + + #[test] + fn test_update_and_remove_provider_allow_legacy_custom_id_with_punctuation() { + let temp_dir = tempfile::tempdir().unwrap(); + let temp_root = temp_dir.path().display().to_string(); + let _guard = env_lock::lock_env([("GOOSE_PATH_ROOT", Some(temp_root.as_str()))]); + + write_legacy_provider_config("custom_z.ai", "Z.AI"); + + update_custom_provider(UpdateCustomProviderParams { + id: "custom_z.ai".to_string(), + engine: "openai".to_string(), + display_name: "Z.AI Updated".to_string(), + api_url: "https://updated.example.invalid/v1/chat/completions".to_string(), + api_key: None, + models: vec!["z-model".to_string()], + supports_streaming: Some(true), + headers: None, + requires_auth: false, + catalog_provider_id: None, + base_path: None, + }) + .unwrap(); + + let updated = load_provider("custom_z.ai").unwrap(); + assert_eq!(updated.config.display_name, "Z.AI Updated"); + assert_eq!(updated.config.models[0].name, "z-model"); + + remove_custom_provider("custom_z.ai").unwrap(); + assert!(!custom_providers_dir().join("custom_z.ai.json").exists()); + } + + #[test] + fn test_load_provider_rejects_path_segments() { + assert!(load_provider("custom_../secret").is_err()); + assert!(load_provider("custom_..\\secret").is_err()); + } + #[test] fn test_expand_env_vars_replaces_placeholder() { let _guard = env_lock::lock_env([("TEST_EXPAND_HOST", Some("https://example.com/api"))]); diff --git a/crates/goose/tests/acp_custom_provider_methods_test.rs b/crates/goose/tests/acp_custom_provider_methods_test.rs new file mode 100644 index 00000000..7a0d7005 --- /dev/null +++ b/crates/goose/tests/acp_custom_provider_methods_test.rs @@ -0,0 +1,604 @@ +#[allow(dead_code)] +#[path = "acp_common_tests/mod.rs"] +mod common_tests; + +use common_tests::fixtures::server::AcpServerConnection; +use common_tests::fixtures::{run_test, send_custom, Connection, TestConnectionConfig}; +use goose::config::base::CONFIG_YAML_NAME; +use goose::config::declarative_providers::load_provider; +use goose::config::paths::Paths; +use goose::config::{Config, ConfigError, DeclarativeProviderConfig}; +use goose_test_support::EnforceSessionId; +use serial_test::serial; +use std::sync::Arc; + +fn write_config(config_dir: &std::path::Path, contents: &str) { + std::fs::create_dir_all(config_dir).unwrap(); + std::fs::write(config_dir.join(CONFIG_YAML_NAME), contents).unwrap(); +} + +fn write_secrets(config_dir: &std::path::Path, contents: &str) { + std::fs::write(config_dir.join("secrets.yaml"), contents).unwrap(); +} + +#[test] +#[serial] +fn acp_catalog_and_custom_provider_methods_use_core_provider_store() { + let root = tempfile::tempdir().unwrap(); + let root_path = root.path().to_string_lossy().to_string(); + let _env = env_lock::lock_env([ + ("GOOSE_PATH_ROOT", Some(root_path.as_str())), + ("GOOSE_DISABLE_KEYRING", Some("1")), + ("XAI_API_KEY", None), + ("XAI_HOST", None), + ("CUSTOM_STARK_ACP_PROVIDER_API_KEY", None), + ]); + + let config_dir = Paths::config_dir(); + write_config( + &config_dir, + "GOOSE_MODEL: gpt-4o\nGOOSE_PROVIDER: openai\nGOOSE_DISABLE_KEYRING: true\nXAI_HOST: https://api.x.ai/v1\n", + ); + write_secrets(&config_dir, "XAI_API_KEY: xai-configured-key\n"); + Config::global().invalidate_secrets_cache(); + + run_test(async move { + let openai = common_tests::fixtures::OpenAiFixture::new( + vec![], + Arc::new(EnforceSessionId::default()), + ) + .await; + let config = TestConnectionConfig { + data_root: config_dir.clone(), + ..Default::default() + }; + let conn = AcpServerConnection::new(config, openai).await; + + let catalog = send_custom( + conn.cx(), + "_goose/providers/catalog/list", + serde_json::json!({ "format": "openai" }), + ) + .await + .expect("provider catalog list should succeed"); + let catalog_providers = catalog + .get("providers") + .and_then(|providers| providers.as_array()) + .expect("catalog response should include providers"); + assert!( + catalog_providers + .iter() + .any(|provider| provider.get("providerId") == Some(&serde_json::json!("zai"))), + "OpenAI-compatible catalog should include z.ai" + ); + + let template = send_custom( + conn.cx(), + "_goose/providers/catalog/template", + serde_json::json!({ "providerId": "zai" }), + ) + .await + .expect("provider catalog template should succeed"); + assert_eq!( + template.pointer("/template/providerId"), + Some(&serde_json::json!("zai")) + ); + assert!( + template + .pointer("/template/models") + .and_then(|models| models.as_array()) + .is_some_and(|models| !models.is_empty()), + "provider template should expose model templates" + ); + + let configured_status = send_custom( + conn.cx(), + "_goose/providers/config/status", + serde_json::json!({ "providerIds": ["xai"] }), + ) + .await + .expect("provider config status should succeed"); + assert_eq!( + configured_status.pointer("/statuses/0"), + Some(&serde_json::json!({ + "providerId": "xai", + "isConfigured": true, + })), + "provider configured through core config should be configured through ACP" + ); + + let configured_read = send_custom( + conn.cx(), + "_goose/providers/config/read", + serde_json::json!({ "providerId": "xai" }), + ) + .await + .expect("provider config read should succeed"); + let fields = configured_read + .get("fields") + .and_then(|fields| fields.as_array()) + .expect("provider config read should include fields"); + let xai_key = fields + .iter() + .find(|field| field.get("key") == Some(&serde_json::json!("XAI_API_KEY"))) + .expect("provider config read should include XAI_API_KEY"); + assert_eq!(xai_key.get("isSet"), Some(&serde_json::json!(true))); + assert_ne!( + xai_key.get("value"), + Some(&serde_json::json!("xai-configured-key")), + "provider config read should not expose raw secret values" + ); + + Config::global().invalidate_secrets_cache(); + assert!(Config::global() + .get_secret::("CUSTOM_STARK_ACP_PROVIDER_API_KEY") + .is_err()); + + let created = send_custom( + conn.cx(), + "_goose/providers/custom/create", + serde_json::json!({ + "engine": "openai_compatible", + "displayName": "Stark ACP Provider", + "apiUrl": "https://stark.example/v1", + "apiKey": "created-custom-key", + "models": ["stark-1", "stark-2"], + "supportsStreaming": true, + "headers": { + "X-Stark": "enabled" + }, + "requiresAuth": true, + "catalogProviderId": "openai", + "basePath": "v1/chat/completions" + }), + ) + .await + .expect("custom provider create should succeed"); + let provider_id = created + .get("providerId") + .and_then(|provider_id| provider_id.as_str()) + .expect("custom provider create should return providerId") + .to_string(); + assert_eq!(provider_id, "custom_stark_acp_provider"); + assert_eq!( + created.get("status"), + Some(&serde_json::json!({ + "providerId": provider_id, + "isConfigured": true, + })), + "create should invalidate the secret cache before status checks" + ); + assert_eq!( + created.get("refresh"), + Some(&serde_json::json!({ + "started": [], + "skipped": [ + { + "providerId": provider_id, + "reason": "does_not_support_refresh", + }, + ], + })) + ); + + let custom_provider_path = Paths::config_dir() + .join("custom_providers") + .join(format!("{provider_id}.json")); + assert!( + custom_provider_path.exists(), + "custom provider should be saved in Goose's declarative provider store" + ); + let saved_provider: DeclarativeProviderConfig = + serde_json::from_str(&std::fs::read_to_string(&custom_provider_path).unwrap()) + .expect("saved provider should be core-compatible declarative config"); + assert_eq!(saved_provider.name, provider_id); + assert_eq!(saved_provider.display_name, "Stark ACP Provider"); + assert_eq!(saved_provider.base_url, "https://stark.example/v1"); + assert_eq!( + saved_provider + .models + .iter() + .map(|model| model.name.as_str()) + .collect::>(), + vec!["stark-1", "stark-2"] + ); + assert_eq!( + Config::global() + .get_secret::("CUSTOM_STARK_ACP_PROVIDER_API_KEY") + .unwrap(), + "created-custom-key", + "custom provider create should write through Goose's config store" + ); + assert!( + load_provider(&provider_id) + .expect("core should load the ACP-created custom provider") + .is_editable + ); + + let read = send_custom( + conn.cx(), + "_goose/providers/custom/read", + serde_json::json!({ "providerId": provider_id }), + ) + .await + .expect("custom provider read should succeed"); + assert_eq!(read.get("editable"), Some(&serde_json::json!(true))); + assert_eq!( + read.pointer("/provider"), + Some(&serde_json::json!({ + "providerId": provider_id, + "engine": "openai_compatible", + "displayName": "Stark ACP Provider", + "apiUrl": "https://stark.example/v1", + "models": ["stark-1", "stark-2"], + "supportsStreaming": true, + "headers": { + "X-Stark": "enabled" + }, + "requiresAuth": true, + "catalogProviderId": "openai", + "basePath": "v1/chat/completions", + "apiKeyEnv": "CUSTOM_STARK_ACP_PROVIDER_API_KEY", + "apiKeySet": true, + })) + ); + + let inventory = send_custom( + conn.cx(), + "_goose/providers/list", + serde_json::json!({ "providerIds": [provider_id] }), + ) + .await + .expect("provider inventory list should include custom provider"); + assert_eq!( + inventory.pointer("/entries/0/providerType"), + Some(&serde_json::json!("Custom")) + ); + assert_eq!( + inventory.pointer("/entries/0/providerId"), + Some(&serde_json::json!(provider_id)) + ); + + let updated = send_custom( + conn.cx(), + "_goose/providers/custom/update", + serde_json::json!({ + "providerId": provider_id, + "engine": "openai", + "displayName": "Stark ACP Provider Updated", + "apiUrl": "https://stark.example/openai", + "apiKey": "updated-custom-key", + "models": ["stark-3"], + "supportsStreaming": false, + "headers": {}, + "requiresAuth": true, + "catalogProviderId": "zai" + }), + ) + .await + .expect("custom provider update should succeed"); + assert_eq!( + updated.get("status"), + Some(&serde_json::json!({ + "providerId": provider_id, + "isConfigured": true, + })), + "update should invalidate the secret cache before status checks" + ); + assert_eq!( + Config::global() + .get_secret::("CUSTOM_STARK_ACP_PROVIDER_API_KEY") + .unwrap(), + "updated-custom-key", + "custom provider update should write through Goose's config store" + ); + let updated_provider: DeclarativeProviderConfig = + serde_json::from_str(&std::fs::read_to_string(&custom_provider_path).unwrap()) + .expect("updated provider should remain core-compatible"); + assert_eq!(updated_provider.display_name, "Stark ACP Provider Updated"); + assert_eq!(updated_provider.base_url, "https://stark.example/openai"); + assert_eq!( + updated_provider.catalog_provider_id, + Some("zai".to_string()) + ); + assert_eq!(updated_provider.base_path, None); + assert_eq!(updated_provider.headers, None); + assert_eq!( + updated_provider + .models + .iter() + .map(|model| model.name.as_str()) + .collect::>(), + vec!["stark-3"] + ); + + let auth_disabled = send_custom( + conn.cx(), + "_goose/providers/custom/update", + serde_json::json!({ + "providerId": provider_id, + "engine": "openai_compatible", + "displayName": "Stark ACP Provider No Auth", + "apiUrl": "https://stark.example/openai", + "apiKey": "", + "models": ["stark-3"], + "supportsStreaming": false, + "headers": {}, + "requiresAuth": false, + "catalogProviderId": "zai" + }), + ) + .await + .expect("custom provider auth disable should succeed"); + assert_eq!( + auth_disabled.get("status"), + Some(&serde_json::json!({ + "providerId": provider_id, + "isConfigured": true, + })), + "auth disable should invalidate the secret cache before status checks" + ); + let no_auth_provider: DeclarativeProviderConfig = + serde_json::from_str(&std::fs::read_to_string(&custom_provider_path).unwrap()) + .expect("no-auth provider should remain core-compatible"); + assert!(!no_auth_provider.requires_auth); + assert_eq!(no_auth_provider.api_key_env, ""); + assert!( + matches!( + Config::global().get_secret::("CUSTOM_STARK_ACP_PROVIDER_API_KEY"), + Err(ConfigError::NotFound(_)) + ), + "disabling auth should delete the previously stored API key" + ); + + let auth_reenabled_without_key = send_custom( + conn.cx(), + "_goose/providers/custom/update", + serde_json::json!({ + "providerId": provider_id, + "engine": "openai_compatible", + "displayName": "Stark ACP Provider Reauth", + "apiUrl": "https://stark.example/openai", + "apiKey": "", + "models": ["stark-3"], + "supportsStreaming": false, + "headers": {}, + "requiresAuth": true, + "catalogProviderId": "zai" + }), + ) + .await + .expect_err("re-enabling auth without a stored secret should fail"); + assert!( + auth_reenabled_without_key + .to_string() + .contains("apiKey is required"), + "unexpected error: {auth_reenabled_without_key}" + ); + assert!( + matches!( + Config::global().get_secret::("CUSTOM_STARK_ACP_PROVIDER_API_KEY"), + Err(ConfigError::NotFound(_)) + ), + "blank re-enable should not recreate the previous API key" + ); + + let deleted = send_custom( + conn.cx(), + "_goose/providers/custom/delete", + serde_json::json!({ "providerId": provider_id }), + ) + .await + .expect("custom provider delete should succeed"); + assert_eq!( + deleted.pointer("/providerId"), + Some(&serde_json::json!(provider_id)) + ); + assert_eq!( + deleted.get("refresh"), + Some(&serde_json::json!({ + "started": [], + "skipped": [], + })) + ); + assert!( + !custom_provider_path.exists(), + "custom provider delete should remove the declarative provider file" + ); + assert!( + matches!( + Config::global().get_secret::("CUSTOM_STARK_ACP_PROVIDER_API_KEY"), + Err(ConfigError::NotFound(_)) + ), + "custom provider delete should invalidate the secret cache before later reads" + ); + + let deleted_status = send_custom( + conn.cx(), + "_goose/providers/config/status", + serde_json::json!({ "providerIds": [provider_id] }), + ) + .await + .expect("provider config status should succeed after delete"); + assert_eq!( + deleted_status.pointer("/statuses/0"), + Some(&serde_json::json!({ + "providerId": provider_id, + "isConfigured": false, + })) + ); + + for invalid_id in [ + "../escape", + "foo/bar", + ".hidden", + "-bad", + "", + "Uppercase", + "has space", + ] { + let read = send_custom( + conn.cx(), + "_goose/providers/custom/read", + serde_json::json!({ "providerId": invalid_id }), + ) + .await; + assert!( + read.is_err(), + "invalid provider id should fail: {invalid_id:?}" + ); + } + + for valid_id in ["custom_openai", "openai-compat", "a1"] { + assert!( + goose::config::declarative_providers::validate_provider_id(valid_id).is_ok(), + "provider id should be valid: {valid_id}" + ); + } + + for (name, patch) in [ + ( + "ftp URL", + serde_json::json!({ "apiUrl": "ftp://example.com" }), + ), + ("relative URL", serde_json::json!({ "apiUrl": "/v1" })), + ("empty models", serde_json::json!({ "models": [] })), + ("blank models", serde_json::json!({ "models": [" ", "\n"] })), + ( + "invalid header name", + serde_json::json!({ "headers": { "Bad Header": "value" } }), + ), + ( + "invalid header value", + serde_json::json!({ "headers": { "X-Test": "bad\r\nvalue" } }), + ), + ( + "unsupported engine", + serde_json::json!({ "engine": "future_engine" }), + ), + ] { + let mut payload = serde_json::json!({ + "engine": "openai_compatible", + "displayName": format!("Invalid {name}"), + "apiUrl": "https://api.example.test/v1", + "apiKey": "secret", + "models": ["model-a"], + "headers": {}, + "requiresAuth": true + }); + let payload_obj = payload.as_object_mut().unwrap(); + for (key, value) in patch.as_object().unwrap() { + payload_obj.insert(key.clone(), value.clone()); + } + + let result = send_custom(conn.cx(), "_goose/providers/custom/create", payload).await; + assert!(result.is_err(), "{name} should be rejected"); + } + + Config::global() + .set_secret("SHARED_API_KEY", &"shared-secret") + .unwrap(); + + let shared = send_custom( + conn.cx(), + "_goose/providers/custom/create", + serde_json::json!({ + "engine": "openai_compatible", + "displayName": "Shared Secret Test", + "apiUrl": "https://api.example.test/v1", + "apiKey": "owned-secret", + "models": ["model-a"], + "headers": {}, + "requiresAuth": true + }), + ) + .await + .expect("shared-secret provider create should succeed"); + let shared_id = shared + .get("providerId") + .and_then(|provider_id| provider_id.as_str()) + .unwrap() + .to_string(); + let shared_path = Paths::config_dir() + .join("custom_providers") + .join(format!("{shared_id}.json")); + let mut shared_config: DeclarativeProviderConfig = + serde_json::from_str(&std::fs::read_to_string(&shared_path).unwrap()).unwrap(); + shared_config.api_key_env = "SHARED_API_KEY".to_string(); + std::fs::write( + &shared_path, + serde_json::to_string_pretty(&shared_config).unwrap(), + ) + .unwrap(); + Config::global().invalidate_secrets_cache(); + + send_custom( + conn.cx(), + "_goose/providers/custom/update", + serde_json::json!({ + "providerId": shared_id, + "engine": "openai_compatible", + "displayName": "Shared Secret Test", + "apiUrl": "https://api.example.test/v1", + "models": ["model-a"], + "headers": {}, + "requiresAuth": false + }), + ) + .await + .expect("disabling auth should preserve shared secrets"); + assert_eq!( + Config::global() + .get_secret::("SHARED_API_KEY") + .unwrap(), + "shared-secret" + ); + + let shared_delete = send_custom( + conn.cx(), + "_goose/providers/custom/create", + serde_json::json!({ + "engine": "openai_compatible", + "displayName": "Shared Secret Delete", + "apiUrl": "https://api.example.test/v1", + "apiKey": "owned-secret", + "models": ["model-a"], + "headers": {}, + "requiresAuth": true + }), + ) + .await + .expect("shared-delete provider create should succeed"); + let shared_delete_id = shared_delete + .get("providerId") + .and_then(|provider_id| provider_id.as_str()) + .unwrap() + .to_string(); + let shared_delete_path = Paths::config_dir() + .join("custom_providers") + .join(format!("{shared_delete_id}.json")); + let mut shared_delete_config: DeclarativeProviderConfig = + serde_json::from_str(&std::fs::read_to_string(&shared_delete_path).unwrap()).unwrap(); + shared_delete_config.api_key_env = "SHARED_API_KEY".to_string(); + std::fs::write( + &shared_delete_path, + serde_json::to_string_pretty(&shared_delete_config).unwrap(), + ) + .unwrap(); + Config::global().invalidate_secrets_cache(); + + send_custom( + conn.cx(), + "_goose/providers/custom/delete", + serde_json::json!({ "providerId": shared_delete_id }), + ) + .await + .expect("deleting provider should preserve shared secrets"); + assert_eq!( + Config::global() + .get_secret::("SHARED_API_KEY") + .unwrap(), + "shared-secret" + ); + }); +} diff --git a/ui/goose2/src/features/chat/hooks/__tests__/useAgentModelPickerState.test.ts b/ui/goose2/src/features/chat/hooks/__tests__/useAgentModelPickerState.test.ts index be091de5..110abb63 100644 --- a/ui/goose2/src/features/chat/hooks/__tests__/useAgentModelPickerState.test.ts +++ b/ui/goose2/src/features/chat/hooks/__tests__/useAgentModelPickerState.test.ts @@ -122,4 +122,55 @@ describe("useAgentModelPickerState", () => { recommended: true, }); }); + + it("uses the clicked model when multiple providers expose the same model id", () => { + const onModelSelected = vi.fn(); + const customModel = { + id: "llama3.2", + name: "llama3.2", + displayName: "llama3.2", + providerId: "custom_ollama", + providerName: "Custom Ollama", + }; + + mockUseProviderInventory.mockReturnValue({ + entries: new Map(), + getEntry: () => undefined, + configuredModelProviderEntries: [], + getModelsForAgent: () => [ + { + id: "llama3.2", + name: "llama3.2", + displayName: "llama3.2", + providerId: "ollama", + providerName: "Ollama", + }, + customModel, + ], + loading: false, + }); + + const { result } = renderHook(() => + useAgentModelPickerState({ + providers: [{ id: "goose", label: "Goose" }], + selectedProvider: "ollama", + onProviderSelected: vi.fn(), + onModelSelected, + }), + ); + + act(() => { + result.current.handleModelChange("llama3.2", customModel); + }); + + expect(onModelSelected).toHaveBeenCalledWith({ + id: "llama3.2", + name: "llama3.2", + displayName: "llama3.2", + provider: undefined, + providerId: "custom_ollama", + providerName: "Custom Ollama", + recommended: undefined, + }); + }); }); diff --git a/ui/goose2/src/features/chat/hooks/useAgentModelPickerState.ts b/ui/goose2/src/features/chat/hooks/useAgentModelPickerState.ts index f5f5fc47..d8a817af 100644 --- a/ui/goose2/src/features/chat/hooks/useAgentModelPickerState.ts +++ b/ui/goose2/src/features/chat/hooks/useAgentModelPickerState.ts @@ -145,10 +145,10 @@ export function useAgentModelPickerState({ ); const handleModelChange = useCallback( - (modelId: string) => { - const selectedModel = availableModels.find( - (model) => model.id === modelId, - ); + (modelId: string, selectedModelOverride?: ModelOption) => { + const selectedModel = + selectedModelOverride ?? + availableModels.find((model) => model.id === modelId); onModelSelected?.({ id: modelId, name: selectedModel?.name ?? modelId, diff --git a/ui/goose2/src/features/chat/hooks/useChatSessionController.ts b/ui/goose2/src/features/chat/hooks/useChatSessionController.ts index dc8f629d..efd2ff10 100644 --- a/ui/goose2/src/features/chat/hooks/useChatSessionController.ts +++ b/ui/goose2/src/features/chat/hooks/useChatSessionController.ts @@ -1,7 +1,6 @@ import { useCallback, useEffect, useMemo, useRef, useState } from "react"; -import type { ChatSkillDraft } from "../types"; import type { ChatAttachmentDraft } from "@/shared/types/messages"; -import type { ChatSendOptions } from "../types"; +import type { ChatSendOptions, ChatSkillDraft, ModelOption } from "../types"; import { INITIAL_TOKEN_STATE } from "@/shared/types/chat"; import { useChat } from "./useChat"; import { useAutoCompactPreferences } from "./useAutoCompactPreferences"; @@ -287,14 +286,24 @@ export function useChatSessionController({ ); const handleModelChangeWithContextReset = useCallback( - (modelId: string) => { - if (modelId === effectiveModelSelection?.id) { + (modelId: string, model?: ModelOption) => { + const nextProviderId = model?.providerId; + if ( + modelId === effectiveModelSelection?.id && + (!nextProviderId || + nextProviderId === effectiveModelSelection?.providerId) + ) { return; } useChatStore.getState().resetTokenState(stateSessionId); - handleModelChange(modelId); + handleModelChange(modelId, model); }, - [effectiveModelSelection?.id, handleModelChange, stateSessionId], + [ + effectiveModelSelection?.id, + effectiveModelSelection?.providerId, + handleModelChange, + stateSessionId, + ], ); const handleProjectChange = useCallback( @@ -816,6 +825,7 @@ export function useChatSessionController({ selectedProvider: selectedAgentId, handleProviderChange: handleProviderChangeWithContextReset, currentModelId: effectiveModelSelection?.id ?? null, + currentModelProviderId: effectiveModelSelection?.providerId ?? null, currentModelName: effectiveModelSelection?.name ?? null, availableModels, modelsLoading, diff --git a/ui/goose2/src/features/chat/types.ts b/ui/goose2/src/features/chat/types.ts index 287e4883..e084cb5a 100644 --- a/ui/goose2/src/features/chat/types.ts +++ b/ui/goose2/src/features/chat/types.ts @@ -60,11 +60,12 @@ export interface ChatInputProps { selectedProvider?: string; onProviderChange?: (providerId: string) => void; currentModelId?: string | null; + currentModelProviderId?: string | null; currentModel?: string; availableModels?: ModelOption[]; modelsLoading?: boolean; modelStatusMessage?: string | null; - onModelChange?: (modelId: string) => void; + onModelChange?: (modelId: string, model?: ModelOption) => void; onPickerOpen?: () => void; selectedProjectId?: string | null; availableProjects?: ProjectOption[]; diff --git a/ui/goose2/src/features/chat/ui/AgentModelPicker.tsx b/ui/goose2/src/features/chat/ui/AgentModelPicker.tsx index 99148839..5d82c273 100644 --- a/ui/goose2/src/features/chat/ui/AgentModelPicker.tsx +++ b/ui/goose2/src/features/chat/ui/AgentModelPicker.tsx @@ -1,16 +1,10 @@ -import { useEffect, useMemo, useRef, useState, type ReactNode } from "react"; -import { - IconCheck, - IconChevronDown, - IconChevronLeft, - IconSearch, -} from "@tabler/icons-react"; +import { useEffect, useState } from "react"; +import { IconCheck, IconChevronDown } from "@tabler/icons-react"; import { useTranslation } from "react-i18next"; import type { AcpProvider } from "@/shared/api/acp"; import { cn } from "@/shared/lib/cn"; import { Button } from "@/shared/ui/button"; import { Popover, PopoverContent, PopoverTrigger } from "@/shared/ui/popover"; -import { SearchBar } from "@/shared/ui/SearchBar"; import { ScrollArea } from "@/shared/ui/scroll-area"; import { Spinner } from "@/shared/ui/spinner"; import { @@ -18,300 +12,34 @@ import { getProviderIcon, } from "@/shared/ui/icons/ProviderIcons"; import type { ModelOption } from "../types"; +import { AllModelsList, RecommendedModelList } from "./AgentModelPickerLists"; +import { PickerItem } from "./AgentModelPickerItem"; interface AgentModelPickerProps { agents: AcpProvider[]; selectedAgentId: string; onAgentChange: (agentId: string) => void; currentModelId?: string | null; + currentModelProviderId?: string | null; currentModelName?: string | null; availableModels: ModelOption[]; modelsLoading?: boolean; modelStatusMessage?: string | null; - onModelChange?: (modelId: string) => void; + onModelChange?: (modelId: string, model?: ModelOption) => void; loading?: boolean; isCompact?: boolean; showSelectedModelInTrigger?: boolean; onOpen?: () => void; } -function getModelDisplayName(model: ModelOption) { - return model.displayName ?? model.name; -} - -function getGooseModelProviderLabel(model: ModelOption) { - if (model.providerName) { - return model.providerName; - } - - if (model.providerId) { - return formatProviderLabel(model.providerId); - } - - return null; -} - -function sortModels(models: ModelOption[], currentModelId: string | null) { - return [...models].sort((left, right) => { - if (left.id === currentModelId) return -1; - if (right.id === currentModelId) return 1; - - const leftProvider = getGooseModelProviderLabel(left) ?? ""; - const rightProvider = getGooseModelProviderLabel(right) ?? ""; - if (leftProvider !== rightProvider) { - return leftProvider.localeCompare(rightProvider); - } - - return getModelDisplayName(left).localeCompare(getModelDisplayName(right)); - }); -} - -function PickerItem({ - children, - onClick, - selected = false, - disabled = false, - className, -}: { - children: ReactNode; - onClick?: () => void; - selected?: boolean; - disabled?: boolean; - className?: string; -}) { - return ( - - ); -} - -// ── Model list views ──────────────────────────────────────────────── - type ModelView = "recommended" | "all"; -function RecommendedModelList({ - models, - currentModelId, - selectedAgentId, - onModelSelect, - onShowAll, - t, -}: { - models: ModelOption[]; - currentModelId: string | null; - selectedAgentId: string; - onModelSelect: (id: string) => void; - onShowAll: () => void; - t: (key: string) => string; -}) { - const recommended = useMemo(() => { - const rec = models.filter((m) => m.recommended); - // If the current model isn't in the recommended list, prepend it - // so the user can always see what's selected. - if ( - currentModelId && - rec.length > 0 && - !rec.some((m) => m.id === currentModelId) - ) { - const current = models.find((m) => m.id === currentModelId); - if (current) { - return [current, ...rec]; - } - } - // Fall back to full list if no recommendations exist (e.g. ACP agents). - return rec.length > 0 ? rec : models; - }, [models, currentModelId]); - - const sorted = useMemo( - () => sortModels(recommended, currentModelId), - [recommended, currentModelId], - ); - - const hasMore = models.length > recommended.length; - - return ( -
-
- {t("toolbar.model")} -
- -
- {sorted.map((model) => { - const providerLabel = getGooseModelProviderLabel(model); - return ( - onModelSelect(model.id)} - selected={model.id === currentModelId} - className="justify-between" - > -
- {selectedAgentId === "goose" && model.providerId ? ( - - {getProviderIcon(model.providerId, "size-3.5")} - - ) : null} -
- {getModelDisplayName(model)} -
-
- {model.id === currentModelId ? ( - - ) : null} -
- ); - })} -
-
- {hasMore ? ( -
- -
- ) : null} -
- ); -} - -function AllModelsList({ - models, - currentModelId, - selectedAgentId, - onModelSelect, - onBack, - t, -}: { - models: ModelOption[]; - currentModelId: string | null; - selectedAgentId: string; - onModelSelect: (id: string) => void; - onBack: () => void; - t: (key: string) => string; -}) { - const [query, setQuery] = useState(""); - const inputRef = useRef(null); - - useEffect(() => { - // Auto-focus search on mount. - inputRef.current?.focus(); - }, []); - - const filtered = useMemo(() => { - if (!query.trim()) { - return sortModels(models, currentModelId); - } - const q = query.toLowerCase(); - const matches = models.filter( - (m) => - m.name.toLowerCase().includes(q) || - m.id.toLowerCase().includes(q) || - m.displayName?.toLowerCase().includes(q) || - m.providerName?.toLowerCase().includes(q) || - m.providerId?.toLowerCase().includes(q), - ); - return sortModels(matches, currentModelId); - }, [models, query, currentModelId]); - - return ( -
-
- - -
- {filtered.length > 0 ? ( - -
- {filtered.map((model) => { - const providerLabel = getGooseModelProviderLabel(model); - const displayName = getModelDisplayName(model); - // Show the raw model_id as secondary text when it differs from name - const showModelId = - model.id !== model.name && model.id !== displayName; - - return ( - onModelSelect(model.id)} - selected={model.id === currentModelId} - className="justify-between" - > -
- {selectedAgentId === "goose" && model.providerId ? ( - - {getProviderIcon(model.providerId, "size-3.5")} - - ) : null} -
-
{displayName}
- {showModelId ? ( -
- {model.id} -
- ) : null} -
-
- {model.id === currentModelId ? ( - - ) : null} -
- ); - })} -
-
- ) : ( -
- {t("toolbar.noSearchResults")} -
- )} -
- ); -} - -// ── Main component ────────────────────────────────────────────────── - export function AgentModelPicker({ agents, selectedAgentId, onAgentChange, currentModelId = null, + currentModelProviderId = null, currentModelName = null, availableModels, modelsLoading = false, @@ -342,8 +70,8 @@ export function AgentModelPicker({ } }; - const handleModelSelect = (modelId: string) => { - onModelChange?.(modelId); + const handleModelSelect = (model: ModelOption) => { + onModelChange?.(model.id, model); setOpen(false); }; @@ -457,6 +185,7 @@ export function AgentModelPicker({
{agents.map((agent) => { const isSelected = agent.id === selectedAgentId; + const agentIcon = getProviderIcon(agent.id, "size-4"); return ( handleAgentSelect(agent.id)} selected={isSelected} > - - {getProviderIcon(agent.id, "size-4")} - + {agentIcon ? ( + {agentIcon} + ) : null} {agent.label} @@ -513,6 +242,7 @@ export function AgentModelPicker({ setModelView("all")} @@ -522,6 +252,7 @@ export function AgentModelPicker({ setModelView("recommended")} diff --git a/ui/goose2/src/features/chat/ui/AgentModelPickerItem.tsx b/ui/goose2/src/features/chat/ui/AgentModelPickerItem.tsx new file mode 100644 index 00000000..c09fb372 --- /dev/null +++ b/ui/goose2/src/features/chat/ui/AgentModelPickerItem.tsx @@ -0,0 +1,33 @@ +import type { ReactNode } from "react"; +import { cn } from "@/shared/lib/cn"; + +export function PickerItem({ + children, + onClick, + selected = false, + disabled = false, + className, +}: { + children: ReactNode; + onClick?: () => void; + selected?: boolean; + disabled?: boolean; + className?: string; +}) { + return ( + + ); +} diff --git a/ui/goose2/src/features/chat/ui/AgentModelPickerLists.tsx b/ui/goose2/src/features/chat/ui/AgentModelPickerLists.tsx new file mode 100644 index 00000000..41d5a4cd --- /dev/null +++ b/ui/goose2/src/features/chat/ui/AgentModelPickerLists.tsx @@ -0,0 +1,285 @@ +import { useEffect, useMemo, useRef, useState } from "react"; +import { IconCheck, IconChevronLeft, IconSearch } from "@tabler/icons-react"; +import { SearchBar } from "@/shared/ui/SearchBar"; +import { ScrollArea } from "@/shared/ui/scroll-area"; +import { + formatProviderLabel, + getProviderIcon, +} from "@/shared/ui/icons/ProviderIcons"; +import type { ModelOption } from "../types"; +import { PickerItem } from "./AgentModelPickerItem"; + +function getModelDisplayName(model: ModelOption) { + return model.displayName ?? model.name; +} + +function getGooseModelProviderLabel(model: ModelOption) { + if (model.providerName) { + return model.providerName; + } + + if (model.providerId) { + return formatProviderLabel(model.providerId); + } + + return null; +} + +function modelMatchesSelection( + model: ModelOption, + currentModelId: string | null, + currentModelProviderId: string | null, +) { + if (model.id !== currentModelId) { + return false; + } + + if (currentModelProviderId) { + return model.providerId === currentModelProviderId; + } + + // Providerless selections are ambiguous legacy/incomplete state, so fall back + // to model-ID-only matching until the user selects a concrete provider row. + return true; +} + +function sortModels( + models: ModelOption[], + currentModelId: string | null, + currentModelProviderId: string | null, +) { + return [...models].sort((left, right) => { + if (modelMatchesSelection(left, currentModelId, currentModelProviderId)) { + return -1; + } + if (modelMatchesSelection(right, currentModelId, currentModelProviderId)) { + return 1; + } + + const leftProvider = getGooseModelProviderLabel(left) ?? ""; + const rightProvider = getGooseModelProviderLabel(right) ?? ""; + if (leftProvider !== rightProvider) { + return leftProvider.localeCompare(rightProvider); + } + + return getModelDisplayName(left).localeCompare(getModelDisplayName(right)); + }); +} + +interface ModelListProps { + models: ModelOption[]; + currentModelId: string | null; + currentModelProviderId: string | null; + selectedAgentId: string; + onModelSelect: (model: ModelOption) => void; + t: (key: string) => string; +} + +export function RecommendedModelList({ + models, + currentModelId, + currentModelProviderId, + selectedAgentId, + onModelSelect, + onShowAll, + t, +}: ModelListProps & { onShowAll: () => void }) { + const recommended = useMemo(() => { + const rec = models.filter((m) => m.recommended); + if ( + currentModelId && + rec.length > 0 && + !rec.some((m) => + modelMatchesSelection(m, currentModelId, currentModelProviderId), + ) + ) { + const current = models.find((m) => + modelMatchesSelection(m, currentModelId, currentModelProviderId), + ); + if (current) { + return [current, ...rec]; + } + } + return rec.length > 0 ? rec : models; + }, [models, currentModelId, currentModelProviderId]); + + const sorted = useMemo( + () => sortModels(recommended, currentModelId, currentModelProviderId), + [recommended, currentModelId, currentModelProviderId], + ); + + const hasMore = models.length > recommended.length; + + return ( +
+
+ {t("toolbar.model")} +
+ +
+ {sorted.map((model) => { + const providerLabel = getGooseModelProviderLabel(model); + const providerIcon = + selectedAgentId === "goose" && model.providerId + ? getProviderIcon(model.providerId, "size-3.5") + : null; + const isSelected = modelMatchesSelection( + model, + currentModelId, + currentModelProviderId, + ); + return ( + onModelSelect(model)} + selected={isSelected} + className="justify-between" + > +
+ {providerIcon ? ( + + {providerIcon} + + ) : null} +
+ {getModelDisplayName(model)} +
+
+ {isSelected ? ( + + ) : null} +
+ ); + })} +
+
+ {hasMore ? ( +
+ +
+ ) : null} +
+ ); +} + +export function AllModelsList({ + models, + currentModelId, + currentModelProviderId, + selectedAgentId, + onModelSelect, + onBack, + t, +}: ModelListProps & { onBack: () => void }) { + const [query, setQuery] = useState(""); + const inputRef = useRef(null); + + useEffect(() => { + inputRef.current?.focus(); + }, []); + + const filtered = useMemo(() => { + if (!query.trim()) { + return sortModels(models, currentModelId, currentModelProviderId); + } + const q = query.toLowerCase(); + const matches = models.filter( + (m) => + m.name.toLowerCase().includes(q) || + m.id.toLowerCase().includes(q) || + m.displayName?.toLowerCase().includes(q) || + m.providerName?.toLowerCase().includes(q) || + m.providerId?.toLowerCase().includes(q), + ); + return sortModels(matches, currentModelId, currentModelProviderId); + }, [models, query, currentModelId, currentModelProviderId]); + + return ( +
+
+ + +
+ {filtered.length > 0 ? ( + +
+ {filtered.map((model) => { + const providerLabel = getGooseModelProviderLabel(model); + const providerIcon = + selectedAgentId === "goose" && model.providerId + ? getProviderIcon(model.providerId, "size-3.5") + : null; + const displayName = getModelDisplayName(model); + const showModelId = + model.id !== model.name && model.id !== displayName; + const isSelected = modelMatchesSelection( + model, + currentModelId, + currentModelProviderId, + ); + + return ( + onModelSelect(model)} + selected={isSelected} + className="justify-between" + > +
+ {providerIcon ? ( + + {providerIcon} + + ) : null} +
+
{displayName}
+ {showModelId ? ( +
+ {model.id} +
+ ) : null} +
+
+ {isSelected ? ( + + ) : null} +
+ ); + })} +
+
+ ) : ( +
+ {t("toolbar.noSearchResults")} +
+ )} +
+ ); +} diff --git a/ui/goose2/src/features/chat/ui/ChatInput.tsx b/ui/goose2/src/features/chat/ui/ChatInput.tsx index 26abd164..2c74caa5 100644 --- a/ui/goose2/src/features/chat/ui/ChatInput.tsx +++ b/ui/goose2/src/features/chat/ui/ChatInput.tsx @@ -46,6 +46,7 @@ export function ChatInput({ selectedProvider = "goose", onProviderChange, currentModelId = null, + currentModelProviderId = null, currentModel, availableModels = [], modelsLoading = false, @@ -457,6 +458,7 @@ export function ChatInput({ selectedProvider={selectedProvider} onProviderChange={(id) => onProviderChange?.(id)} currentModelId={currentModelId} + currentModelProviderId={currentModelProviderId} currentModel={resolvedCurrentModel} availableModels={availableModels} modelsLoading={modelsLoading} diff --git a/ui/goose2/src/features/chat/ui/ChatInputToolbar.tsx b/ui/goose2/src/features/chat/ui/ChatInputToolbar.tsx index 26250037..8594af55 100644 --- a/ui/goose2/src/features/chat/ui/ChatInputToolbar.tsx +++ b/ui/goose2/src/features/chat/ui/ChatInputToolbar.tsx @@ -46,11 +46,12 @@ interface ChatInputToolbarProps { onProviderChange: (providerId: string) => void; // Model currentModelId?: string | null; + currentModelProviderId?: string | null; currentModel?: string; availableModels: ModelOption[]; modelsLoading?: boolean; modelStatusMessage?: string | null; - onModelChange?: (modelId: string) => void; + onModelChange?: (modelId: string, model?: ModelOption) => void; onPickerOpen?: () => void; // Project selectedProjectId: string | null; @@ -92,6 +93,7 @@ export function ChatInputToolbar({ selectedProvider, onProviderChange, currentModelId, + currentModelProviderId, currentModel, availableModels, modelsLoading = false, @@ -216,6 +218,7 @@ export function ChatInputToolbar({ selectedAgentId={selectedProvider} onAgentChange={onProviderChange} currentModelId={currentModelId} + currentModelProviderId={currentModelProviderId} currentModelName={currentModel ?? null} availableModels={availableModels} modelsLoading={modelsLoading} diff --git a/ui/goose2/src/features/chat/ui/ChatView.tsx b/ui/goose2/src/features/chat/ui/ChatView.tsx index b5041f94..e227cbfa 100644 --- a/ui/goose2/src/features/chat/ui/ChatView.tsx +++ b/ui/goose2/src/features/chat/ui/ChatView.tsx @@ -134,6 +134,7 @@ export function ChatView({ selectedProvider={controller.selectedProvider} onProviderChange={controller.handleProviderChange} currentModelId={controller.currentModelId} + currentModelProviderId={controller.currentModelProviderId} currentModel={controller.currentModelName ?? undefined} availableModels={controller.availableModels} modelsLoading={controller.modelsLoading} diff --git a/ui/goose2/src/features/chat/ui/__tests__/AgentModelPicker.test.tsx b/ui/goose2/src/features/chat/ui/__tests__/AgentModelPicker.test.tsx index 0c42694a..ebdcecff 100644 --- a/ui/goose2/src/features/chat/ui/__tests__/AgentModelPicker.test.tsx +++ b/ui/goose2/src/features/chat/ui/__tests__/AgentModelPicker.test.tsx @@ -61,7 +61,108 @@ describe("AgentModelPicker", () => { await user.click(screen.getByRole("button", { name: "GPT-4o" })); - expect(onModelChange).toHaveBeenCalledWith("gpt-4o"); + expect(onModelChange).toHaveBeenCalledWith( + "gpt-4o", + expect.objectContaining({ id: "gpt-4o" }), + ); + }); + + it("passes the clicked model option through for duplicate model ids", async () => { + const user = userEvent.setup(); + const onModelChange = vi.fn(); + + render( + , + ); + + await user.click( + screen.getByRole("button", { name: /choose agent and model/i }), + ); + + const duplicateModelRows = screen.getAllByRole("button", { + name: "llama3.2", + }); + + const selectedDuplicateRows = duplicateModelRows.filter((row) => + row.classList.contains("bg-muted/60"), + ); + expect(selectedDuplicateRows).toHaveLength(1); + + await user.click(selectedDuplicateRows[0]); + + expect(onModelChange).toHaveBeenCalledWith( + "llama3.2", + expect.objectContaining({ + name: "llama3.2", + providerId: "custom_ollama", + }), + ); + }); + + it("does not select providerless duplicate rows when the current provider is known", async () => { + const user = userEvent.setup(); + + render( + , + ); + + await user.click( + screen.getByRole("button", { name: /choose agent and model/i }), + ); + + const duplicateModelRows = screen.getAllByRole("button", { + name: "llama3.2", + }); + + expect( + duplicateModelRows.filter((row) => row.classList.contains("bg-muted/60")), + ).toHaveLength(1); }); it("auto-expands the group containing the selected model", async () => { diff --git a/ui/goose2/src/features/providers/api/customProviders.test.ts b/ui/goose2/src/features/providers/api/customProviders.test.ts new file mode 100644 index 00000000..992713b7 --- /dev/null +++ b/ui/goose2/src/features/providers/api/customProviders.test.ts @@ -0,0 +1,158 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { + createCustomProvider, + deleteCustomProvider, + getCustomProviderTemplate, + listCustomProviderCatalog, + readCustomProvider, + updateCustomProvider, +} from "./customProviders"; + +const mocks = vi.hoisted(() => ({ + catalogList: vi.fn(), + catalogTemplate: vi.fn(), + customCreate: vi.fn(), + customRead: vi.fn(), + customUpdate: vi.fn(), + customDelete: vi.fn(), + getClient: vi.fn(), +})); + +vi.mock("@/shared/api/acpConnection", () => ({ + getClient: () => mocks.getClient(), +})); + +describe("custom provider API", () => { + const input = { + engine: "openai_compatible" as const, + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + apiKey: "secret", + models: ["acme-large"], + supportsStreaming: true, + headers: { + "X-Acme": "goose", + }, + requiresAuth: true, + catalogProviderId: "acme", + basePath: "/v1", + }; + + beforeEach(() => { + vi.clearAllMocks(); + mocks.getClient.mockResolvedValue({ + goose: { + GooseProvidersCatalogList: mocks.catalogList, + GooseProvidersCatalogTemplate: mocks.catalogTemplate, + GooseProvidersCustomCreate: mocks.customCreate, + GooseProvidersCustomRead: mocks.customRead, + GooseProvidersCustomUpdate: mocks.customUpdate, + GooseProvidersCustomDelete: mocks.customDelete, + }, + }); + }); + + it("lists catalog providers with the planned typed ACP method", async () => { + const providers = [ + { + providerId: "acme", + name: "Acme AI", + format: "openai", + apiUrl: "https://api.acme.test/v1", + modelCount: 1, + docUrl: "https://acme.test/docs", + envVar: "ACME_API_KEY", + }, + ]; + mocks.catalogList.mockResolvedValue({ providers }); + + await expect(listCustomProviderCatalog("openai")).resolves.toEqual( + providers, + ); + + expect(mocks.catalogList).toHaveBeenCalledWith({ format: "openai" }); + }); + + it("reads a catalog template through the planned typed ACP method", async () => { + const template = { + providerId: "acme", + name: "Acme AI", + format: "openai", + apiUrl: "https://api.acme.test/v1", + models: [], + supportsStreaming: true, + envVar: "ACME_API_KEY", + docUrl: "https://acme.test/docs", + }; + mocks.catalogTemplate.mockResolvedValue({ template }); + + await expect(getCustomProviderTemplate("acme")).resolves.toEqual(template); + + expect(mocks.catalogTemplate).toHaveBeenCalledWith({ + providerId: "acme", + }); + }); + + it("creates, reads, updates, and deletes custom providers by generated method name", async () => { + const createResponse = { + providerId: "acme_ai", + status: { providerId: "acme_ai", isConfigured: true }, + refresh: { started: ["acme_ai"], skipped: [] }, + }; + const readResponse = { + provider: { + providerId: "acme_ai", + ...input, + headers: input.headers ?? {}, + apiKeyEnv: "ACME_AI_API_KEY", + apiKeySet: true, + }, + editable: true, + status: { providerId: "acme_ai", isConfigured: true }, + }; + const updateResponse = createResponse; + const deleteResponse = { + providerId: "acme_ai", + refresh: { started: [], skipped: [] }, + }; + mocks.customCreate.mockResolvedValue(createResponse); + mocks.customRead.mockResolvedValue(readResponse); + mocks.customUpdate.mockResolvedValue(updateResponse); + mocks.customDelete.mockResolvedValue(deleteResponse); + + await expect(createCustomProvider(input)).resolves.toEqual(createResponse); + await expect(readCustomProvider("acme_ai")).resolves.toEqual(readResponse); + await expect(updateCustomProvider("acme_ai", input)).resolves.toEqual( + updateResponse, + ); + await expect(deleteCustomProvider("acme_ai")).resolves.toEqual( + deleteResponse, + ); + + expect(mocks.customCreate).toHaveBeenCalledWith(input); + expect(mocks.customRead).toHaveBeenCalledWith({ providerId: "acme_ai" }); + expect(mocks.customUpdate).toHaveBeenCalledWith({ + ...input, + providerId: "acme_ai", + }); + expect(mocks.customDelete).toHaveBeenCalledWith({ providerId: "acme_ai" }); + }); + + it("lets the explicit update target override a conflicting runtime provider id", async () => { + mocks.customUpdate.mockResolvedValue({ + providerId: "acme_ai", + status: { providerId: "acme_ai", isConfigured: true }, + refresh: { started: [], skipped: [] }, + }); + + await updateCustomProvider("acme_ai", { + ...input, + providerId: "wrong_id", + } as typeof input & { providerId: string }); + + expect(mocks.customUpdate).toHaveBeenCalledWith({ + ...input, + providerId: "acme_ai", + }); + }); +}); diff --git a/ui/goose2/src/features/providers/api/customProviders.ts b/ui/goose2/src/features/providers/api/customProviders.ts new file mode 100644 index 00000000..88c8dd55 --- /dev/null +++ b/ui/goose2/src/features/providers/api/customProviders.ts @@ -0,0 +1,65 @@ +import { getClient } from "@/shared/api/acpConnection"; +import type { + CustomProviderCreateResponse, + CustomProviderDeleteResponse, + CustomProviderReadResponse, + CustomProviderUpdateResponse, + ProviderCatalogEntryDto, + ProviderTemplateDto, +} from "@aaif/goose-sdk"; +import type { + CustomProviderFormat, + CustomProviderUpsertRequest, +} from "../lib/customProviderTypes"; + +async function getProviderClient() { + const client = await getClient(); + return client.goose; +} + +export async function listCustomProviderCatalog( + format?: CustomProviderFormat, +): Promise { + const client = await getProviderClient(); + const response = await client.GooseProvidersCatalogList( + format ? { format } : {}, + ); + return response.providers; +} + +export async function getCustomProviderTemplate( + providerId: string, +): Promise { + const client = await getProviderClient(); + const response = await client.GooseProvidersCatalogTemplate({ providerId }); + return response.template; +} + +export async function createCustomProvider( + input: CustomProviderUpsertRequest, +): Promise { + const client = await getProviderClient(); + return client.GooseProvidersCustomCreate(input); +} + +export async function readCustomProvider( + providerId: string, +): Promise { + const client = await getProviderClient(); + return client.GooseProvidersCustomRead({ providerId }); +} + +export async function updateCustomProvider( + providerId: string, + input: CustomProviderUpsertRequest, +): Promise { + const client = await getProviderClient(); + return client.GooseProvidersCustomUpdate({ ...input, providerId }); +} + +export async function deleteCustomProvider( + providerId: string, +): Promise { + const client = await getProviderClient(); + return client.GooseProvidersCustomDelete({ providerId }); +} diff --git a/ui/goose2/src/features/providers/hooks/useCustomProviders.test.tsx b/ui/goose2/src/features/providers/hooks/useCustomProviders.test.tsx new file mode 100644 index 00000000..f13ffe76 --- /dev/null +++ b/ui/goose2/src/features/providers/hooks/useCustomProviders.test.tsx @@ -0,0 +1,192 @@ +import { act, renderHook, waitFor } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useProviderInventoryStore } from "../stores/providerInventoryStore"; +import { useCustomProviders } from "./useCustomProviders"; + +const mocks = vi.hoisted(() => ({ + createCustomProvider: vi.fn(), + deleteCustomProvider: vi.fn(), + getCustomProviderTemplate: vi.fn(), + listCustomProviderCatalog: vi.fn(), + readCustomProvider: vi.fn(), + updateCustomProvider: vi.fn(), + syncProviderInventory: vi.fn(), +})); + +vi.mock("../api/customProviders", () => ({ + createCustomProvider: mocks.createCustomProvider, + deleteCustomProvider: mocks.deleteCustomProvider, + getCustomProviderTemplate: mocks.getCustomProviderTemplate, + listCustomProviderCatalog: mocks.listCustomProviderCatalog, + readCustomProvider: mocks.readCustomProvider, + updateCustomProvider: mocks.updateCustomProvider, +})); + +vi.mock("../api/inventorySync", () => ({ + syncProviderInventory: mocks.syncProviderInventory, +})); + +function providerEntry(providerId: string) { + return { + providerId, + providerName: "Acme AI", + description: "", + defaultModel: "acme-large", + configured: true, + providerType: "Custom", + configKeys: [], + setupSteps: [], + supportsRefresh: true, + refreshing: false, + models: [], + stale: false, + }; +} + +describe("useCustomProviders", () => { + const input = { + engine: "openai_compatible" as const, + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + apiKey: "secret", + models: ["acme-large"], + supportsStreaming: true, + requiresAuth: true, + }; + + beforeEach(() => { + vi.clearAllMocks(); + useProviderInventoryStore.getState().setEntries([]); + mocks.createCustomProvider.mockResolvedValue({ + providerId: "acme_ai", + status: { providerId: "acme_ai", isConfigured: true }, + refresh: { started: ["acme_ai"], skipped: [] }, + }); + mocks.updateCustomProvider.mockResolvedValue({ + providerId: "acme_ai", + status: { providerId: "acme_ai", isConfigured: true }, + refresh: { started: ["acme_ai"], skipped: [] }, + }); + mocks.deleteCustomProvider.mockResolvedValue({ + providerId: "acme_ai", + refresh: { started: [], skipped: [] }, + }); + mocks.readCustomProvider.mockResolvedValue({ + provider: { + providerId: "acme_ai", + engine: "openai_compatible", + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + models: ["acme-large"], + supportsStreaming: true, + headers: {}, + requiresAuth: true, + apiKeySet: true, + }, + editable: true, + status: { providerId: "acme_ai", isConfigured: true }, + }); + mocks.listCustomProviderCatalog.mockResolvedValue([]); + mocks.syncProviderInventory.mockImplementation( + async (_providerIds, options) => { + const entries = [providerEntry("acme_ai")]; + options?.onEntries?.(entries); + return { + entries, + refresh: { started: ["acme_ai"], skipped: [] }, + settled: true, + polledProviderIds: ["acme_ai"], + }; + }, + ); + }); + + it("loads catalog providers into hook state", async () => { + const providers = [ + { + providerId: "acme", + name: "Acme AI", + format: "openai", + apiUrl: "https://api.acme.test/v1", + modelCount: 1, + docUrl: "https://acme.test/docs", + envVar: "ACME_API_KEY", + }, + ]; + mocks.listCustomProviderCatalog.mockResolvedValue(providers); + const { result } = renderHook(() => useCustomProviders()); + + await act(async () => { + await result.current.loadCatalog("openai"); + }); + + expect(result.current.catalog).toEqual(providers); + expect(mocks.listCustomProviderCatalog).toHaveBeenCalledWith("openai"); + }); + + it("creates a provider, tracks configured status, and merges inventory entries", async () => { + const { result } = renderHook(() => useCustomProviders()); + + await act(async () => { + await result.current.create(input); + }); + + expect(mocks.createCustomProvider).toHaveBeenCalledWith(input); + expect(result.current.configuredIds.has("acme_ai")).toBe(true); + await waitFor(() => + expect( + useProviderInventoryStore.getState().entries.get("acme_ai"), + ).toEqual(providerEntry("acme_ai")), + ); + }); + + it("reads a provider and merges its status", async () => { + const { result } = renderHook(() => useCustomProviders()); + + await act(async () => { + await result.current.read("acme_ai"); + }); + + expect(mocks.readCustomProvider).toHaveBeenCalledWith("acme_ai"); + expect(result.current.configuredIds.has("acme_ai")).toBe(true); + }); + + it("updates from a validated draft", async () => { + const { result } = renderHook(() => useCustomProviders()); + + await act(async () => { + await result.current.saveDraft({ + providerId: "acme_ai", + editable: true, + ...input, + apiKeySet: false, + basePath: "", + modelsInput: "acme-large", + headers: [], + authInitiallyEnabled: true, + }); + }); + + expect(mocks.updateCustomProvider).toHaveBeenCalledWith("acme_ai", input); + }); + + it("removes stale inventory entries after deleting a custom provider", async () => { + mocks.syncProviderInventory.mockResolvedValueOnce({ + entries: [], + refresh: { started: [], skipped: [] }, + settled: true, + polledProviderIds: ["acme_ai"], + }); + useProviderInventoryStore.getState().setEntries([providerEntry("acme_ai")]); + const { result } = renderHook(() => useCustomProviders()); + + await act(async () => { + await result.current.remove("acme_ai"); + }); + + expect(mocks.deleteCustomProvider).toHaveBeenCalledWith("acme_ai"); + expect(useProviderInventoryStore.getState().entries.has("acme_ai")).toBe( + false, + ); + }); +}); diff --git a/ui/goose2/src/features/providers/hooks/useCustomProviders.ts b/ui/goose2/src/features/providers/hooks/useCustomProviders.ts new file mode 100644 index 00000000..1913b25e --- /dev/null +++ b/ui/goose2/src/features/providers/hooks/useCustomProviders.ts @@ -0,0 +1,318 @@ +import { useCallback, useMemo, useRef, useState } from "react"; +import type { ProviderConfigStatusDto } from "@aaif/goose-sdk"; +import { + createCustomProvider, + deleteCustomProvider, + getCustomProviderTemplate, + listCustomProviderCatalog, + readCustomProvider, + updateCustomProvider, +} from "../api/customProviders"; +import { + syncProviderInventory, + type SyncProviderInventoryResult, +} from "../api/inventorySync"; +import { + assertValidCustomProviderDraft, + type CustomProviderValidationOptions, +} from "../lib/customProviderValidation"; +import { customProviderDraftToUpsertRequest } from "../lib/customProviderDraft"; +import type { + CustomProviderCreateResponse, + CustomProviderDeleteResponse, + CustomProviderDraft, + CustomProviderFormat, + CustomProviderReadResponse, + CustomProviderUpdateResponse, + CustomProviderUpsertRequest, + ProviderCatalogEntryDto, + ProviderTemplateDto, +} from "../lib/customProviderTypes"; +import { useProviderInventoryStore } from "../stores/providerInventoryStore"; + +interface SaveDraftOptions extends CustomProviderValidationOptions { + providerId?: string; +} + +interface UseCustomProvidersReturn { + catalog: ProviderCatalogEntryDto[]; + catalogLoading: boolean; + saving: boolean; + savingProviderIds: Set; + deletingProviderIds: Set; + syncingProviderIds: Set; + inventoryWarnings: Map; + statusByProviderId: Map; + configuredIds: Set; + loadCatalog: ( + format?: CustomProviderFormat, + ) => Promise; + getTemplate: (providerId: string) => Promise; + read: (providerId: string) => Promise; + create: ( + input: CustomProviderUpsertRequest, + ) => Promise; + update: ( + providerId: string, + input: CustomProviderUpsertRequest, + ) => Promise; + remove: (providerId: string) => Promise; + saveDraft: ( + draft: CustomProviderDraft, + options?: SaveDraftOptions, + ) => Promise; +} + +function errorMessage(error: unknown): string { + return error instanceof Error ? error.message : String(error); +} + +function inventoryWarning( + providerId: string, + result: SyncProviderInventoryResult, +): string | null { + const entry = result.entries.find((item) => item.providerId === providerId); + const skipped = result.refresh.skipped?.find( + (item) => item.providerId === providerId, + ); + + if (skipped?.reason === "unknown_provider") { + return "Provider inventory is unavailable."; + } + + if (entry?.lastRefreshError) { + return entry.lastRefreshError; + } + + if (!result.settled && entry?.refreshing) { + return "Model inventory is still refreshing."; + } + + return null; +} + +function useSetMembershipState() { + const [state, setState] = useState>(() => new Set()); + + const setMembership = useCallback((providerId: string, present: boolean) => { + setState((current) => { + const next = new Set(current); + if (present) { + next.add(providerId); + } else { + next.delete(providerId); + } + return next; + }); + }, []); + + return [state, setMembership] as const; +} + +export function useCustomProviders(): UseCustomProvidersReturn { + const catalogRequestIdRef = useRef(0); + const operationIdRef = useRef(0); + const deletedProviderIdsRef = useRef(new Set()); + const [catalog, setCatalog] = useState([]); + const [catalogLoading, setCatalogLoading] = useState(false); + const [savingProviderIds, setProviderSaving] = useSetMembershipState(); + const [deletingProviderIds, setProviderDeleting] = useSetMembershipState(); + const [syncingProviderIds, setProviderSyncing] = useSetMembershipState(); + const [statusByProviderId, setStatusByProviderId] = useState< + Map + >(() => new Map()); + const [inventoryWarnings, setInventoryWarnings] = useState< + Map + >(() => new Map()); + + const saving = savingProviderIds.size > 0 || deletingProviderIds.size > 0; + + const configuredIds = useMemo( + () => + new Set( + [...statusByProviderId.values()] + .filter((status) => status.isConfigured) + .map((status) => status.providerId), + ), + [statusByProviderId], + ); + + const setProviderInventoryWarning = useCallback( + (providerId: string, warning: string | null) => { + setInventoryWarnings((current) => { + const next = new Map(current); + if (warning) { + next.set(providerId, warning); + } else { + next.delete(providerId); + } + return next; + }); + }, + [], + ); + + const updateStatus = useCallback((status: ProviderConfigStatusDto) => { + setStatusByProviderId((current) => { + const next = new Map(current); + next.set(status.providerId, status); + return next; + }); + }, []); + + const removeInventoryEntry = useCallback((providerId: string) => { + const store = useProviderInventoryStore.getState(); + store.setEntries( + [...store.entries.values()].filter( + (entry) => entry.providerId !== providerId, + ), + ); + }, []); + + const startInventorySync = useCallback( + (providerId: string, result: SyncProviderInventoryResult["refresh"]) => { + setProviderSyncing(providerId, true); + setProviderInventoryWarning(providerId, null); + + void syncProviderInventory([providerId], { + initialRefresh: result, + onEntries: (entries) => { + const visibleEntries = entries.filter( + (entry) => !deletedProviderIdsRef.current.has(entry.providerId), + ); + if (visibleEntries.length > 0) { + useProviderInventoryStore.getState().mergeEntries(visibleEntries); + } + }, + }) + .then((syncResult) => { + setProviderInventoryWarning( + providerId, + inventoryWarning(providerId, syncResult), + ); + }) + .catch((error) => { + setProviderInventoryWarning(providerId, errorMessage(error)); + }) + .finally(() => setProviderSyncing(providerId, false)); + }, + [setProviderInventoryWarning, setProviderSyncing], + ); + + const loadCatalog = useCallback(async (format?: CustomProviderFormat) => { + const requestId = catalogRequestIdRef.current + 1; + catalogRequestIdRef.current = requestId; + setCatalogLoading(true); + try { + const nextCatalog = await listCustomProviderCatalog(format); + if (catalogRequestIdRef.current === requestId) { + setCatalog(nextCatalog); + } + return nextCatalog; + } finally { + if (catalogRequestIdRef.current === requestId) { + setCatalogLoading(false); + } + } + }, []); + + const read = useCallback( + async (providerId: string) => { + const result = await readCustomProvider(providerId); + updateStatus(result.status); + return result; + }, + [updateStatus], + ); + + const create = useCallback( + async (input: CustomProviderUpsertRequest) => { + const pendingId = `create-${operationIdRef.current + 1}`; + operationIdRef.current += 1; + setProviderSaving(pendingId, true); + try { + const result = await createCustomProvider(input); + deletedProviderIdsRef.current.delete(result.providerId); + updateStatus(result.status); + startInventorySync(result.providerId, result.refresh); + return result; + } finally { + setProviderSaving(pendingId, false); + } + }, + [setProviderSaving, startInventorySync, updateStatus], + ); + + const update = useCallback( + async (providerId: string, input: CustomProviderUpsertRequest) => { + setProviderSaving(providerId, true); + try { + const result = await updateCustomProvider(providerId, input); + deletedProviderIdsRef.current.delete(result.providerId); + updateStatus(result.status); + startInventorySync(result.providerId, result.refresh); + return result; + } finally { + setProviderSaving(providerId, false); + } + }, + [setProviderSaving, startInventorySync, updateStatus], + ); + + const remove = useCallback( + async (providerId: string) => { + setProviderDeleting(providerId, true); + deletedProviderIdsRef.current.add(providerId); + try { + const result = await deleteCustomProvider(providerId); + setStatusByProviderId((current) => { + const next = new Map(current); + next.set(providerId, { providerId, isConfigured: false }); + return next; + }); + removeInventoryEntry(providerId); + return result; + } catch (error) { + deletedProviderIdsRef.current.delete(providerId); + throw error; + } finally { + setProviderDeleting(providerId, false); + } + }, + [removeInventoryEntry, setProviderDeleting], + ); + + const saveDraft = useCallback( + async (draft: CustomProviderDraft, options: SaveDraftOptions = {}) => { + assertValidCustomProviderDraft(draft, options); + const providerId = options.providerId ?? draft.providerId; + const input = customProviderDraftToUpsertRequest(draft); + + if (providerId) { + return update(providerId, input); + } + + return create(input); + }, + [create, update], + ); + + return { + catalog, + catalogLoading, + saving, + savingProviderIds, + deletingProviderIds, + syncingProviderIds, + inventoryWarnings, + statusByProviderId, + configuredIds, + loadCatalog, + getTemplate: getCustomProviderTemplate, + read, + create, + update, + remove, + saveDraft, + }; +} diff --git a/ui/goose2/src/features/providers/hooks/useProviderInventory.test.ts b/ui/goose2/src/features/providers/hooks/useProviderInventory.test.ts new file mode 100644 index 00000000..677d427b --- /dev/null +++ b/ui/goose2/src/features/providers/hooks/useProviderInventory.test.ts @@ -0,0 +1,119 @@ +import { renderHook } from "@testing-library/react"; +import type { ProviderInventoryEntryDto } from "@aaif/goose-sdk"; +import { beforeEach, describe, expect, it } from "vitest"; +import { useProviderInventoryStore } from "../stores/providerInventoryStore"; +import { useProviderInventory } from "./useProviderInventory"; + +function providerEntry( + overrides: Partial, +): ProviderInventoryEntryDto { + const providerId = overrides.providerId ?? "openai"; + + return { + providerId, + providerName: overrides.providerName ?? providerId, + description: "", + defaultModel: "", + configured: true, + providerType: "Preferred", + configKeys: [], + setupSteps: [], + supportsRefresh: true, + refreshing: false, + models: [], + stale: false, + ...overrides, + }; +} + +describe("useProviderInventory", () => { + beforeEach(() => { + useProviderInventoryStore.setState({ + entries: new Map(), + loading: false, + }); + }); + + it("shows configured static, custom, and curated declarative model providers", () => { + useProviderInventoryStore.getState().setEntries([ + providerEntry({ + providerId: "openai", + providerName: "OpenAI", + providerType: "Preferred", + }), + providerEntry({ + providerId: "custom_acme_openai", + providerName: "Acme OpenAI", + providerType: "Custom", + }), + providerEntry({ + providerId: "custom_deepseek", + providerName: "DeepSeek", + providerType: "Declarative", + }), + providerEntry({ + providerId: "internal_declarative", + providerName: "Internal Declarative", + providerType: "Declarative", + }), + providerEntry({ + providerId: "unconfigured_custom", + providerName: "Unconfigured Custom", + providerType: "Custom", + configured: false, + }), + providerEntry({ + providerId: "local", + providerName: "Local", + providerType: "Custom", + }), + providerEntry({ + providerId: "local_inference", + providerName: "Local Inference", + providerType: "Custom", + }), + ]); + + const { result } = renderHook(() => useProviderInventory()); + + expect( + result.current.configuredModelProviderEntries.map( + (entry) => entry.providerId, + ), + ).toEqual(["openai", "custom_acme_openai", "custom_deepseek"]); + }); + + it("aggregates custom provider models under Goose", () => { + useProviderInventoryStore.getState().setEntries([ + providerEntry({ + providerId: "custom_acme_openai", + providerName: "Acme OpenAI", + providerType: "Custom", + models: [ + { + id: "acme-gpt-5", + name: "Acme GPT-5", + family: "acme", + contextLimit: 128000, + recommended: true, + }, + ], + }), + ]); + + const { result } = renderHook(() => useProviderInventory()); + + expect(result.current.getModelsForAgent("goose")).toEqual([ + { + id: "acme-gpt-5", + name: "Acme GPT-5", + displayName: "Acme GPT-5", + provider: "acme", + providerId: "custom_acme_openai", + providerName: "Acme OpenAI", + contextLimit: 128000, + recommended: true, + }, + ]); + }); +}); diff --git a/ui/goose2/src/features/providers/hooks/useProviderInventory.ts b/ui/goose2/src/features/providers/hooks/useProviderInventory.ts index ddf0b013..1251ab2c 100644 --- a/ui/goose2/src/features/providers/hooks/useProviderInventory.ts +++ b/ui/goose2/src/features/providers/hooks/useProviderInventory.ts @@ -9,6 +9,26 @@ import { getModelProviders } from "../providerCatalog"; const MODEL_PROVIDER_IDS = new Set(getModelProviders().map((p) => p.id)); +function isConfiguredGooseModelProvider( + entry: ProviderInventoryEntryDto, +): boolean { + if (!entry.configured) { + return false; + } + + const isCuratedModelProvider = MODEL_PROVIDER_IDS.has(entry.providerId); + + if (entry.providerType === "Custom") { + return entry.providerId.startsWith("custom_"); + } + + if (entry.providerType === "Declarative") { + return isCuratedModelProvider; + } + + return isCuratedModelProvider; +} + function inventoryModelToOption( model: ProviderInventoryModelDto, provider?: Pick, @@ -44,10 +64,7 @@ export function useProviderInventory() { ); const configuredModelProviderEntries = useMemo( - () => - [...entries.values()].filter( - (entry) => entry.configured && MODEL_PROVIDER_IDS.has(entry.providerId), - ), + () => [...entries.values()].filter(isConfiguredGooseModelProvider), [entries], ); diff --git a/ui/goose2/src/features/providers/lib/customProviderDraft.test.ts b/ui/goose2/src/features/providers/lib/customProviderDraft.test.ts new file mode 100644 index 00000000..6c642c4b --- /dev/null +++ b/ui/goose2/src/features/providers/lib/customProviderDraft.test.ts @@ -0,0 +1,208 @@ +import { describe, expect, it } from "vitest"; +import { + createEmptyCustomProviderDraft, + customProviderDraftToUpsertRequest, + readToCustomProviderDraft, + templateToCustomProviderDraft, +} from "./customProviderDraft"; +import { + headerDraftsToRecord, + recordToHeaderDrafts, + validateCustomProviderHeaders, +} from "./customProviderHeaders"; +import { + formatCustomProviderModels, + parseCustomProviderModels, +} from "./customProviderModels"; +import { validateCustomProviderDraft } from "./customProviderValidation"; + +describe("custom provider helper functions", () => { + it("parses model input from comma and newline separated values", () => { + expect( + parseCustomProviderModels("claude-3-5-sonnet, gpt-4.1\n gpt-4.1"), + ).toEqual(["claude-3-5-sonnet", "gpt-4.1"]); + expect(formatCustomProviderModels(["a", "a", "b"])).toBe("a, b"); + }); + + it("converts header records to drafts and ignores blank draft rows on submit", () => { + expect(recordToHeaderDrafts({ Authorization: "Bearer token" })).toEqual([ + { + id: "server-header-0", + key: "Authorization", + value: "Bearer token", + }, + ]); + expect( + headerDraftsToRecord([ + { id: "a", key: " X-Test ", value: " enabled " }, + { id: "b", key: "", value: "" }, + ]), + ).toEqual({ + "X-Test": "enabled", + }); + }); + + it("reports header validation issues with stable i18n keys", () => { + const issues = validateCustomProviderHeaders([ + { id: "a", key: "Bad Header", value: "value" }, + { id: "b", key: "X-Test", value: "" }, + { id: "c", key: "x-test", value: "duplicate" }, + ]); + + expect(issues.map((issue) => issue.key)).toEqual([ + "settings.providers.custom.validation.headerNameInvalid", + "settings.providers.custom.validation.headerValueRequired", + "settings.providers.custom.validation.headerDuplicate", + ]); + }); + + it("builds a draft from a catalog template", () => { + const draft = templateToCustomProviderDraft({ + providerId: "acme", + name: "Acme AI", + format: "openai", + apiUrl: "https://api.acme.test/v1", + models: [ + { + id: "acme-large", + name: "Acme Large", + contextLimit: 128000, + capabilities: { + toolCall: true, + reasoning: false, + attachment: false, + temperature: true, + }, + deprecated: false, + }, + { + id: "acme-old", + name: "Acme Old", + contextLimit: 8192, + capabilities: { + toolCall: false, + reasoning: false, + attachment: false, + temperature: true, + }, + deprecated: true, + }, + ], + supportsStreaming: true, + envVar: "ACME_API_KEY", + docUrl: "https://acme.test/docs", + }); + + expect(draft).toMatchObject({ + engine: "openai_compatible", + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + models: ["acme-large"], + modelsInput: "acme-large", + catalogProviderId: "acme", + }); + }); + + it("builds a draft from an editable read response", () => { + const draft = readToCustomProviderDraft({ + provider: { + providerId: "acme_ai", + engine: "openai_compatible", + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + models: ["acme-large"], + supportsStreaming: true, + headers: { "X-Test": "enabled" }, + requiresAuth: true, + catalogProviderId: "acme", + basePath: "/v1", + apiKeyEnv: "ACME_AI_API_KEY", + apiKeySet: true, + }, + editable: true, + status: { providerId: "acme_ai", isConfigured: true }, + }); + + expect(draft).toMatchObject({ + providerId: "acme_ai", + headers: [{ id: "server-header-0", key: "X-Test", value: "enabled" }], + basePath: "/v1", + }); + }); + + it("validates required fields and maps draft fields to ACP upsert input", () => { + const emptyIssues = validateCustomProviderDraft( + createEmptyCustomProviderDraft(), + ); + expect(emptyIssues.map((issue) => issue.key)).toContain( + "settings.providers.custom.validation.displayNameRequired", + ); + + const draft = { + ...createEmptyCustomProviderDraft(), + displayName: " Acme AI ", + apiUrl: " https://api.acme.test/v1 ", + apiKey: " secret ", + modelsInput: "acme-large, acme-small", + headers: [{ id: "a", key: " X-Test ", value: " enabled " }], + basePath: " /v1 ", + catalogProviderId: "acme", + }; + + expect(validateCustomProviderDraft(draft)).toEqual([]); + expect(customProviderDraftToUpsertRequest(draft)).toEqual({ + engine: "openai_compatible", + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + apiKey: "secret", + models: ["acme-large", "acme-small"], + supportsStreaming: true, + headers: { + "X-Test": "enabled", + }, + requiresAuth: true, + catalogProviderId: "acme", + basePath: "/v1", + }); + }); + + it("omits unchanged API keys and preserves stable header ids", () => { + const draft = { + ...createEmptyCustomProviderDraft(), + providerId: "acme_ai", + displayName: "Acme AI", + apiUrl: "https://api.acme.test/v1", + apiKeySet: true, + models: ["acme-large"], + headers: [ + { id: "stable", key: "X-Original", value: "enabled" }, + { id: "empty", key: "", value: "" }, + ], + }; + + expect(validateCustomProviderDraft(draft)).toEqual([]); + expect(customProviderDraftToUpsertRequest(draft)).not.toHaveProperty( + "apiKey", + ); + + const nextHeaders = draft.headers.map((header) => + header.id === "stable" ? { ...header, key: "X-Renamed" } : header, + ); + expect(nextHeaders[0].id).toBe("stable"); + }); + + it("surfaces unknown engines as invalid instead of normalizing them", () => { + const draft = { + ...createEmptyCustomProviderDraft(), + engine: "future_engine", + displayName: "Future AI", + apiUrl: "https://api.future.test/v1", + apiKey: "secret", + models: ["future-large"], + }; + + expect( + validateCustomProviderDraft(draft).map((issue) => issue.field), + ).toContain("engine"); + }); +}); diff --git a/ui/goose2/src/features/providers/lib/customProviderDraft.ts b/ui/goose2/src/features/providers/lib/customProviderDraft.ts new file mode 100644 index 00000000..e8f6f7bc --- /dev/null +++ b/ui/goose2/src/features/providers/lib/customProviderDraft.ts @@ -0,0 +1,151 @@ +import { + formatCustomProviderModels, + parseCustomProviderModels, +} from "./customProviderModels"; +import { + headerDraftsToRecord, + recordToHeaderDrafts, +} from "./customProviderHeaders"; +import type { + CustomProviderDraft, + CustomProviderEngine, + CustomProviderReadResponse, + CustomProviderUpsertRequest, + ProviderTemplateDto, +} from "./customProviderTypes"; + +const FORMAT_ENGINE_MAP: Record = { + openai: "openai_compatible", + anthropic: "anthropic_compatible", + ollama: "ollama_compatible", +}; + +const ENGINE_MAP: Record = { + openai: "openai_compatible", + openai_compatible: "openai_compatible", + anthropic: "anthropic_compatible", + anthropic_compatible: "anthropic_compatible", + ollama: "ollama_compatible", + ollama_compatible: "ollama_compatible", +}; + +export function isCustomProviderEngine( + engine: string | undefined, +): engine is CustomProviderEngine { + if (!engine) { + return false; + } + const normalized = engine.trim().toLowerCase(); + return Boolean(ENGINE_MAP[normalized]); +} + +export function normalizeCustomProviderEngine( + engine: string | undefined, +): string { + if (!engine) { + return ""; + } + const normalized = engine.trim().toLowerCase(); + return ENGINE_MAP[normalized] ?? normalized; +} + +export function engineForCustomProviderFormat( + format: string | undefined, +): CustomProviderEngine { + return FORMAT_ENGINE_MAP[format ?? ""] ?? "openai_compatible"; +} + +export function createEmptyCustomProviderDraft(): CustomProviderDraft { + return { + editable: true, + engine: "openai_compatible", + displayName: "", + apiUrl: "", + basePath: "", + apiKey: "", + apiKeySet: false, + modelsInput: "", + models: [], + authInitiallyEnabled: true, + requiresAuth: true, + supportsStreaming: true, + headers: [], + }; +} + +export function templateToCustomProviderDraft( + template: ProviderTemplateDto, +): CustomProviderDraft { + const models = (template.models ?? []) + .filter((model) => !model.deprecated) + .map((model) => model.id); + + return { + editable: true, + engine: engineForCustomProviderFormat(template.format), + displayName: template.name, + apiUrl: template.apiUrl, + basePath: "", + apiKey: "", + apiKeySet: false, + modelsInput: formatCustomProviderModels(models), + models, + authInitiallyEnabled: true, + requiresAuth: true, + supportsStreaming: template.supportsStreaming, + headers: [], + catalogProviderId: template.providerId, + }; +} + +export function readToCustomProviderDraft( + response: CustomProviderReadResponse, +): CustomProviderDraft { + const provider = response.provider; + const models = parseCustomProviderModels(provider.models ?? []); + + return { + providerId: provider.providerId, + editable: response.editable, + engine: normalizeCustomProviderEngine(provider.engine), + displayName: provider.displayName, + apiUrl: provider.apiUrl, + basePath: provider.basePath ?? "", + apiKey: "", + apiKeySet: provider.apiKeySet, + modelsInput: formatCustomProviderModels(models), + models, + authInitiallyEnabled: provider.requiresAuth, + requiresAuth: provider.requiresAuth, + supportsStreaming: provider.supportsStreaming ?? true, + headers: recordToHeaderDrafts(provider.headers), + catalogProviderId: provider.catalogProviderId ?? undefined, + }; +} + +export function customProviderDraftToUpsertRequest( + draft: CustomProviderDraft, +): CustomProviderUpsertRequest { + const models = parseCustomProviderModels( + draft.models.length > 0 ? draft.models : draft.modelsInput, + ); + + const apiKey = draft.requiresAuth ? draft.apiKey.trim() : ""; + const request: CustomProviderUpsertRequest = { + engine: normalizeCustomProviderEngine(draft.engine) as CustomProviderEngine, + displayName: draft.displayName.trim(), + apiUrl: draft.apiUrl.trim(), + models, + supportsStreaming: draft.supportsStreaming, + headers: headerDraftsToRecord(draft.headers), + requiresAuth: draft.requiresAuth, + catalogProviderId: draft.catalogProviderId, + basePath: draft.basePath.trim() || undefined, + }; + + if (apiKey) { + request.apiKey = apiKey; + } + + return request; +} diff --git a/ui/goose2/src/features/providers/lib/customProviderHeaders.ts b/ui/goose2/src/features/providers/lib/customProviderHeaders.ts new file mode 100644 index 00000000..c8876e93 --- /dev/null +++ b/ui/goose2/src/features/providers/lib/customProviderHeaders.ts @@ -0,0 +1,120 @@ +import type { CustomProviderHeaderDraft } from "./customProviderTypes"; + +export interface CustomProviderHeaderIssue { + field: "headers"; + key: + | "settings.providers.custom.validation.headerNameRequired" + | "settings.providers.custom.validation.headerValueRequired" + | "settings.providers.custom.validation.headerNameInvalid" + | "settings.providers.custom.validation.headerDuplicate"; + message: string; + index?: number; +} + +const HEADER_TOKEN_RE = /^[!#$%&'*+.^_`|~0-9A-Za-z-]+$/; +let nextHeaderId = 0; + +export function createCustomProviderHeaderDraft( + key = "", + value = "", +): CustomProviderHeaderDraft { + nextHeaderId += 1; + return { + id: `header-${nextHeaderId}`, + key, + value, + }; +} + +export function normalizeHeaderName(name: string): string { + return name.trim(); +} + +export function normalizeHeaderValue(value: string): string { + return value.trim(); +} + +export function recordToHeaderDrafts( + headers?: Record | null, +): CustomProviderHeaderDraft[] { + return Object.entries(headers ?? {}).map(([key, value], index) => ({ + id: `server-header-${index}`, + key, + value, + })); +} + +export function headerDraftsToRecord( + headers: CustomProviderHeaderDraft[], +): Record | undefined { + const record: Record = {}; + + for (const header of headers) { + const key = normalizeHeaderName(header.key); + const value = normalizeHeaderValue(header.value); + if (key && value) { + record[key] = value; + } + } + + return Object.keys(record).length > 0 ? record : undefined; +} + +export function validateCustomProviderHeaders( + headers: CustomProviderHeaderDraft[], +): CustomProviderHeaderIssue[] { + const issues: CustomProviderHeaderIssue[] = []; + const seen = new Map(); + + headers.forEach((header, index) => { + const key = normalizeHeaderName(header.key); + const value = normalizeHeaderValue(header.value); + const normalizedKey = key.toLowerCase(); + + if (!key && !value) { + return; + } + + if (!key) { + issues.push({ + field: "headers", + key: "settings.providers.custom.validation.headerNameRequired", + message: "Header name is required.", + index, + }); + return; + } + + if (!HEADER_TOKEN_RE.test(key)) { + issues.push({ + field: "headers", + key: "settings.providers.custom.validation.headerNameInvalid", + message: "Header names can only contain valid HTTP token characters.", + index, + }); + } + + if (!value) { + issues.push({ + field: "headers", + key: "settings.providers.custom.validation.headerValueRequired", + message: "Header value is required.", + index, + }); + } + + if (seen.has(normalizedKey)) { + issues.push({ + field: "headers", + key: "settings.providers.custom.validation.headerDuplicate", + message: "Header names must be unique.", + index, + }); + return; + } + + seen.set(normalizedKey, index); + }); + + return issues; +} diff --git a/ui/goose2/src/features/providers/lib/customProviderModels.ts b/ui/goose2/src/features/providers/lib/customProviderModels.ts new file mode 100644 index 00000000..ef0eb4e6 --- /dev/null +++ b/ui/goose2/src/features/providers/lib/customProviderModels.ts @@ -0,0 +1,23 @@ +export function parseCustomProviderModels(input: string | string[]): string[] { + const rawModels = Array.isArray(input) + ? input + : input.split(/[\n,]/).map((value) => value.trim()); + + const seen = new Set(); + const models: string[] = []; + + for (const rawModel of rawModels) { + const model = rawModel.trim(); + if (!model || seen.has(model)) { + continue; + } + seen.add(model); + models.push(model); + } + + return models; +} + +export function formatCustomProviderModels(models: string[]): string { + return parseCustomProviderModels(models).join(", "); +} diff --git a/ui/goose2/src/features/providers/lib/customProviderTypes.ts b/ui/goose2/src/features/providers/lib/customProviderTypes.ts new file mode 100644 index 00000000..1f082696 --- /dev/null +++ b/ui/goose2/src/features/providers/lib/customProviderTypes.ts @@ -0,0 +1,58 @@ +import type { + CustomProviderConfigDto, + CustomProviderCreateRequest, + CustomProviderCreateResponse, + CustomProviderDeleteResponse, + CustomProviderReadResponse, + CustomProviderUpdateResponse, + ProviderCatalogEntryDto, + ProviderTemplateDto, +} from "@aaif/goose-sdk"; + +export type CustomProviderFormat = "openai" | "anthropic" | "ollama"; + +export type CustomProviderEngine = + | "openai_compatible" + | "anthropic_compatible" + | "ollama_compatible"; + +export interface CustomProviderHeaderDraft { + id: string; + key: string; + value: string; +} + +export interface CustomProviderDraft { + providerId?: string; + editable: boolean; + engine: string; + displayName: string; + apiUrl: string; + basePath: string; + apiKey: string; + apiKeySet: boolean; + modelsInput: string; + models: string[]; + authInitiallyEnabled: boolean; + requiresAuth: boolean; + supportsStreaming: boolean; + headers: CustomProviderHeaderDraft[]; + catalogProviderId?: string; +} + +export type CustomProviderUpsertRequest = Omit< + CustomProviderCreateRequest, + "providerId" +> & { + engine: CustomProviderEngine; +}; + +export type { + CustomProviderConfigDto, + CustomProviderCreateResponse, + CustomProviderDeleteResponse, + CustomProviderReadResponse, + CustomProviderUpdateResponse, + ProviderCatalogEntryDto, + ProviderTemplateDto, +}; diff --git a/ui/goose2/src/features/providers/lib/customProviderValidation.ts b/ui/goose2/src/features/providers/lib/customProviderValidation.ts new file mode 100644 index 00000000..8a5a1821 --- /dev/null +++ b/ui/goose2/src/features/providers/lib/customProviderValidation.ts @@ -0,0 +1,129 @@ +import { parseCustomProviderModels } from "./customProviderModels"; +import { validateCustomProviderHeaders } from "./customProviderHeaders"; +import { + isCustomProviderEngine, + normalizeCustomProviderEngine, +} from "./customProviderDraft"; +import type { CustomProviderDraft } from "./customProviderTypes"; + +export type CustomProviderValidationField = + | "displayName" + | "engine" + | "apiUrl" + | "apiKey" + | "models" + | "headers"; + +export interface CustomProviderValidationIssue { + field: CustomProviderValidationField; + key: + | "settings.providers.custom.validation.displayNameRequired" + | "settings.providers.custom.validation.engineRequired" + | "settings.providers.custom.validation.apiUrlRequired" + | "settings.providers.custom.validation.apiUrlInvalid" + | "settings.providers.custom.validation.apiKeyRequired" + | "settings.providers.custom.validation.modelsRequired" + | "settings.providers.custom.validation.headerNameRequired" + | "settings.providers.custom.validation.headerValueRequired" + | "settings.providers.custom.validation.headerNameInvalid" + | "settings.providers.custom.validation.headerDuplicate"; + message: string; + index?: number; +} + +export interface CustomProviderValidationOptions { + requireApiKey?: boolean; +} + +export class CustomProviderValidationError extends Error { + readonly issues: CustomProviderValidationIssue[]; + + constructor(issues: CustomProviderValidationIssue[]) { + super("Custom provider validation failed."); + this.name = "CustomProviderValidationError"; + this.issues = issues; + } +} + +function isProbablyUrl(value: string): boolean { + try { + const url = new URL(value); + return url.protocol === "http:" || url.protocol === "https:"; + } catch { + return false; + } +} + +export function validateCustomProviderDraft( + draft: CustomProviderDraft, + options: CustomProviderValidationOptions = {}, +): CustomProviderValidationIssue[] { + const issues: CustomProviderValidationIssue[] = []; + const apiUrl = draft.apiUrl.trim(); + const models = parseCustomProviderModels( + draft.models.length > 0 ? draft.models : draft.modelsInput, + ); + const requireApiKey = + options.requireApiKey ?? (draft.requiresAuth && !draft.apiKeySet); + const engine = normalizeCustomProviderEngine(draft.engine); + + if (!draft.displayName.trim()) { + issues.push({ + field: "displayName", + key: "settings.providers.custom.validation.displayNameRequired", + message: "Display name is required.", + }); + } + + if (!isCustomProviderEngine(engine)) { + issues.push({ + field: "engine", + key: "settings.providers.custom.validation.engineRequired", + message: "Choose a provider engine.", + }); + } + + if (!apiUrl) { + issues.push({ + field: "apiUrl", + key: "settings.providers.custom.validation.apiUrlRequired", + message: "API URL is required.", + }); + } else if (!isProbablyUrl(apiUrl)) { + issues.push({ + field: "apiUrl", + key: "settings.providers.custom.validation.apiUrlInvalid", + message: "Enter a valid HTTP or HTTPS URL.", + }); + } + + if (requireApiKey && !draft.apiKey.trim()) { + issues.push({ + field: "apiKey", + key: "settings.providers.custom.validation.apiKeyRequired", + message: "API key is required.", + }); + } + + if (models.length === 0) { + issues.push({ + field: "models", + key: "settings.providers.custom.validation.modelsRequired", + message: "Add at least one model.", + }); + } + + issues.push(...validateCustomProviderHeaders(draft.headers)); + + return issues; +} + +export function assertValidCustomProviderDraft( + draft: CustomProviderDraft, + options?: CustomProviderValidationOptions, +): void { + const issues = validateCustomProviderDraft(draft, options); + if (issues.length > 0) { + throw new CustomProviderValidationError(issues); + } +} diff --git a/ui/goose2/src/features/providers/providerCatalog.test.ts b/ui/goose2/src/features/providers/providerCatalog.test.ts index f150ffdb..244e8973 100644 --- a/ui/goose2/src/features/providers/providerCatalog.test.ts +++ b/ui/goose2/src/features/providers/providerCatalog.test.ts @@ -1,6 +1,7 @@ import { describe, expect, it } from "vitest"; import { getCatalogEntry, + getModelProviders, resolveAgentProviderCatalogId, } from "./providerCatalog"; @@ -20,6 +21,55 @@ describe("provider catalog", () => { }, ]); }); + + it("uses backend model provider ids for the curated catalog", () => { + const ids = getModelProviders().map((provider) => provider.id); + + expect(ids).toEqual([ + "anthropic", + "google", + "chatgpt_codex", + "openai", + "mistral", + "ollama", + "openrouter", + "databricks", + "github_copilot", + "custom_deepseek", + "xai", + "groq", + "azure_openai", + "aws_bedrock", + "gcp_vertex_ai", + "litellm", + "lmstudio", + "nvidia", + "cerebras", + "snowflake", + ]); + expect(ids).not.toContain("azure"); + expect(ids).not.toContain("bedrock"); + expect(ids).not.toContain("deepseek"); + expect(ids).not.toContain("local_inference"); + }); + + it("marks the planned promoted model providers", () => { + const promotedIds = getModelProviders() + .filter((provider) => provider.tier === "promoted") + .map((provider) => provider.id); + + expect(promotedIds).toEqual([ + "anthropic", + "google", + "chatgpt_codex", + "openai", + "mistral", + "ollama", + "openrouter", + "databricks", + "github_copilot", + ]); + }); }); describe("resolveAgentProviderCatalogId", () => { diff --git a/ui/goose2/src/features/providers/providerCatalog.ts b/ui/goose2/src/features/providers/providerCatalog.ts index eb2bc1db..4a103a67 100644 --- a/ui/goose2/src/features/providers/providerCatalog.ts +++ b/ui/goose2/src/features/providers/providerCatalog.ts @@ -4,431 +4,14 @@ import { AGENT_PROVIDER_FUZZY_MATCHERS, normalizeProviderKey, } from "./providerCatalogAliases"; +import { + AGENT_PROVIDER_CATALOG, + MODEL_PROVIDER_CATALOG, +} from "./providerCatalogEntries"; export const PROVIDER_CATALOG: ProviderCatalogEntry[] = [ - { - id: "goose", - displayName: "Goose", - category: "agent", - description: "Block's open-source coding agent", - setupMethod: "none", - tier: "promoted", - }, - { - id: "claude-acp", - displayName: "Claude Code", - category: "agent", - description: "Anthropic's agentic coding tool", - setupMethod: "cli_auth", - binaryName: "claude-agent-acp", - installCommand: - "npm install -g @anthropic-ai/claude-code @agentclientprotocol/claude-agent-acp", - authCommand: "claude auth login", - authStatusCommand: "claude auth status", - docsUrl: "https://docs.anthropic.com/en/docs/claude-code", - tier: "promoted", - }, - { - id: "codex-acp", - displayName: "Codex", - category: "agent", - description: "OpenAI's coding agent", - setupMethod: "cli_auth", - binaryName: "codex-acp", - installCommand: "npm install -g @openai/codex @zed-industries/codex-acp", - authCommand: "codex login", - authStatusCommand: "codex login status", - docsUrl: "https://github.com/openai/codex", - tier: "promoted", - }, - { - id: "copilot-acp", - displayName: "GitHub Copilot", - category: "agent", - description: "GitHub's AI pair programmer", - setupMethod: "cli_auth", - binaryName: "copilot", - installCommand: "npm install -g @github/copilot", - authCommand: "copilot login", - docsUrl: "https://docs.github.com/en/copilot/github-copilot-in-the-cli", - tier: "promoted", - }, - { - id: "amp-acp", - displayName: "Amp", - category: "agent", - description: "Sourcegraph's coding agent", - setupMethod: "cli_auth", - binaryName: "amp-acp", - installCommand: "npm install -g @sourcegraph/amp@latest amp-acp", - authCommand: "amp login", - authStatusCommand: "amp usage", - docsUrl: "https://ampcode.com", - tier: "standard", - }, - { - id: "cursor-agent", - displayName: "Cursor Agent", - category: "agent", - description: "Cursor's AI agent", - setupMethod: "cli_auth", - binaryName: "cursor-agent", - installCommand: "curl -fsSL https://cursor.com/install | bash", - authCommand: "cursor-agent login", - authStatusCommand: "cursor-agent status", - docsUrl: "https://docs.cursor.com/en/cli/overview", - tier: "standard", - }, - { - id: "pi-acp", - displayName: "Pi", - category: "agent", - description: "Open-source AI coding agent", - setupMethod: "cli_auth", - binaryName: "pi-acp", - docsUrl: "https://github.com/badlogic/pi-mono", - tier: "standard", - showOnlyWhenInstalled: true, - }, - - { - id: "anthropic", - displayName: "Anthropic", - category: "model", - description: "Claude models", - setupMethod: "single_api_key", - envVar: "ANTHROPIC_API_KEY", - fields: [ - { - key: "ANTHROPIC_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - docsUrl: "https://console.anthropic.com/settings/keys", - tier: "promoted", - }, - { - id: "google", - displayName: "Google Gemini", - category: "model", - description: "Gemini models", - setupMethod: "single_api_key", - envVar: "GOOGLE_API_KEY", - fields: [ - { - key: "GOOGLE_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - docsUrl: "https://aistudio.google.com/apikey", - tier: "promoted", - }, - { - id: "chatgpt_codex", - displayName: "ChatGPT Codex", - category: "model", - description: "OpenAI via ChatGPT subscription", - setupMethod: "oauth_device_code", - nativeConnectQuery: "ChatGPT Codex", - docsUrl: "https://chatgpt.com", - tier: "standard", - }, - { - id: "openai", - displayName: "OpenAI", - category: "model", - description: "GPT and o-series models", - setupMethod: "config_fields", - envVar: "OPENAI_API_KEY", - fields: [ - { - key: "OPENAI_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - docsUrl: "https://platform.openai.com/api-keys", - tier: "promoted", - }, - { - id: "ollama", - displayName: "Ollama", - category: "model", - description: "Run local or self-hosted models", - setupMethod: "config_fields", - fields: [ - { - key: "OLLAMA_HOST", - label: "Host", - secret: false, - required: true, - placeholder: "localhost or http://localhost:11434", - defaultValue: "http://localhost:11434", - }, - ], - docsUrl: "https://ollama.com", - tier: "promoted", - }, - { - id: "openrouter", - displayName: "OpenRouter", - category: "model", - description: "Unified API for many models", - setupMethod: "single_api_key", - envVar: "OPENROUTER_API_KEY", - fields: [ - { - key: "OPENROUTER_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - docsUrl: "https://openrouter.ai/keys", - tier: "promoted", - }, - { - id: "databricks", - displayName: "Databricks", - category: "model", - description: "Databricks Foundation Models", - setupMethod: "host_with_oauth_fallback", - fields: [ - { - key: "DATABRICKS_HOST", - label: "Host URL", - secret: false, - required: true, - placeholder: "https://dbc-...cloud.databricks.com", - }, - { - key: "DATABRICKS_TOKEN", - label: "Access Token", - secret: true, - required: false, - placeholder: "Paste your access token", - }, - ], - tier: "standard", - }, - { - id: "github_copilot", - displayName: "GitHub Copilot Models", - category: "model", - description: "Models via GitHub Copilot subscription", - setupMethod: "oauth_device_code", - nativeConnectQuery: "GitHub Copilot", - tier: "standard", - }, - { - id: "xai", - displayName: "xAI", - category: "model", - description: "Grok models", - setupMethod: "single_api_key", - envVar: "XAI_API_KEY", - fields: [ - { - key: "XAI_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - tier: "standard", - }, - { - id: "azure", - displayName: "Azure OpenAI", - category: "model", - description: "OpenAI models on Azure", - setupMethod: "config_fields", - fields: [ - { - key: "AZURE_OPENAI_ENDPOINT", - label: "Endpoint", - secret: false, - required: true, - placeholder: "https://your-resource.openai.azure.com", - }, - { - key: "AZURE_OPENAI_DEPLOYMENT_NAME", - label: "Deployment", - secret: false, - required: true, - placeholder: "gpt-4o", - }, - { - key: "AZURE_OPENAI_API_KEY", - label: "API Key", - secret: true, - required: false, - placeholder: "Paste your API key", - }, - ], - tier: "advanced", - }, - { - id: "bedrock", - displayName: "AWS Bedrock", - category: "model", - description: "Models on AWS", - setupMethod: "cloud_credentials", - fields: [ - { - key: "AWS_REGION", - label: "AWS Region", - secret: false, - required: false, - placeholder: "us-west-2", - }, - ], - tier: "advanced", - }, - { - id: "gcp_vertex_ai", - displayName: "GCP Vertex AI", - category: "model", - description: "Models on Google Cloud", - setupMethod: "cloud_credentials", - fields: [ - { - key: "GCP_PROJECT_ID", - label: "Project ID", - secret: false, - required: true, - placeholder: "my-gcp-project", - }, - { - key: "GCP_LOCATION", - label: "Location", - secret: false, - required: true, - placeholder: "us-central1", - }, - ], - tier: "advanced", - }, - { - id: "litellm", - displayName: "LiteLLM", - category: "model", - description: "LiteLLM proxy gateway", - setupMethod: "config_fields", - envVar: "LITELLM_API_KEY", - fields: [ - { - key: "LITELLM_HOST", - label: "Host URL", - secret: false, - required: true, - placeholder: "https://your-proxy.example.com", - }, - { - key: "LITELLM_API_KEY", - label: "API Key", - secret: true, - required: false, - placeholder: "Paste your API key", - }, - ], - tier: "advanced", - }, - { - id: "nanogpt", - displayName: "NanoGPT", - category: "model", - description: "NanoGPT inference", - setupMethod: "single_api_key", - envVar: "NANOGPT_API_KEY", - fields: [ - { - key: "NANOGPT_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - tier: "advanced", - }, - { - id: "tetrate", - displayName: "Tetrate", - category: "model", - description: "Tetrate AI gateway", - setupMethod: "single_api_key", - fields: [ - { - key: "TETRATE_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - tier: "advanced", - }, - { - id: "venice", - displayName: "Venice", - category: "model", - description: "Venice AI", - setupMethod: "single_api_key", - envVar: "VENICE_API_KEY", - fields: [ - { - key: "VENICE_API_KEY", - label: "API Key", - secret: true, - required: true, - placeholder: "Paste your API key", - }, - ], - tier: "advanced", - }, - { - id: "snowflake", - displayName: "Snowflake", - category: "model", - description: "Snowflake Cortex", - setupMethod: "config_fields", - fields: [ - { - key: "SNOWFLAKE_HOST", - label: "Host URL", - secret: false, - required: true, - placeholder: "https://your-account.snowflakecomputing.com", - }, - { - key: "SNOWFLAKE_TOKEN", - label: "Access Token", - secret: true, - required: true, - placeholder: "Paste your access token", - }, - ], - tier: "advanced", - }, - { - id: "local_inference", - displayName: "Local Inference", - category: "model", - description: "Custom local model server", - setupMethod: "local", - tier: "advanced", - }, + ...AGENT_PROVIDER_CATALOG, + ...MODEL_PROVIDER_CATALOG, ]; export function getCatalogEntry( @@ -438,11 +21,11 @@ export function getCatalogEntry( } export function getAgentProviders(): ProviderCatalogEntry[] { - return PROVIDER_CATALOG.filter((p) => p.category === "agent"); + return AGENT_PROVIDER_CATALOG; } export function getModelProviders(): ProviderCatalogEntry[] { - return PROVIDER_CATALOG.filter((p) => p.category === "model"); + return MODEL_PROVIDER_CATALOG; } export function resolveAgentProviderCatalogIdStrict( diff --git a/ui/goose2/src/features/providers/providerCatalogEntries.ts b/ui/goose2/src/features/providers/providerCatalogEntries.ts new file mode 100644 index 00000000..a3c1f017 --- /dev/null +++ b/ui/goose2/src/features/providers/providerCatalogEntries.ts @@ -0,0 +1,481 @@ +import type { ProviderCatalogEntry } from "@/shared/types/providers"; + +export const AGENT_PROVIDER_CATALOG: ProviderCatalogEntry[] = [ + { + id: "goose", + displayName: "Goose", + category: "agent", + description: "Block's open-source coding agent", + setupMethod: "none", + tier: "promoted", + }, + { + id: "claude-acp", + displayName: "Claude Code", + category: "agent", + description: "Anthropic's agentic coding tool", + setupMethod: "cli_auth", + binaryName: "claude-agent-acp", + installCommand: + "npm install -g @anthropic-ai/claude-code @agentclientprotocol/claude-agent-acp", + authCommand: "claude auth login", + authStatusCommand: "claude auth status", + docsUrl: "https://docs.anthropic.com/en/docs/claude-code", + tier: "promoted", + }, + { + id: "codex-acp", + displayName: "Codex", + category: "agent", + description: "OpenAI's coding agent", + setupMethod: "cli_auth", + binaryName: "codex-acp", + installCommand: "npm install -g @openai/codex @zed-industries/codex-acp", + authCommand: "codex login", + authStatusCommand: "codex login status", + docsUrl: "https://github.com/openai/codex", + tier: "promoted", + }, + { + id: "copilot-acp", + displayName: "GitHub Copilot", + category: "agent", + description: "GitHub's AI pair programmer", + setupMethod: "cli_auth", + binaryName: "copilot", + installCommand: "npm install -g @github/copilot", + authCommand: "copilot login", + docsUrl: "https://docs.github.com/en/copilot/github-copilot-in-the-cli", + tier: "promoted", + }, + { + id: "amp-acp", + displayName: "Amp", + category: "agent", + description: "Sourcegraph's coding agent", + setupMethod: "cli_auth", + binaryName: "amp-acp", + installCommand: "npm install -g @sourcegraph/amp@latest amp-acp", + authCommand: "amp login", + authStatusCommand: "amp usage", + docsUrl: "https://ampcode.com", + tier: "standard", + }, + { + id: "cursor-agent", + displayName: "Cursor Agent", + category: "agent", + description: "Cursor's AI agent", + setupMethod: "cli_auth", + binaryName: "cursor-agent", + installCommand: "curl -fsSL https://cursor.com/install | bash", + authCommand: "cursor-agent login", + authStatusCommand: "cursor-agent status", + docsUrl: "https://docs.cursor.com/en/cli/overview", + tier: "standard", + }, + { + id: "pi-acp", + displayName: "Pi", + category: "agent", + description: "Open-source AI coding agent", + setupMethod: "cli_auth", + binaryName: "pi-acp", + docsUrl: "https://github.com/badlogic/pi-mono", + tier: "standard", + showOnlyWhenInstalled: true, + }, +]; + +export const MODEL_PROVIDER_CATALOG: ProviderCatalogEntry[] = [ + { + id: "anthropic", + displayName: "Anthropic", + category: "model", + description: "Claude models", + setupMethod: "single_api_key", + envVar: "ANTHROPIC_API_KEY", + fields: [ + { + key: "ANTHROPIC_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://console.anthropic.com/settings/keys", + tier: "promoted", + }, + { + id: "google", + displayName: "Google Gemini", + category: "model", + description: "Gemini models", + setupMethod: "single_api_key", + envVar: "GOOGLE_API_KEY", + fields: [ + { + key: "GOOGLE_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://aistudio.google.com/apikey", + tier: "promoted", + }, + { + id: "chatgpt_codex", + displayName: "ChatGPT", + category: "model", + description: "OpenAI via ChatGPT subscription", + setupMethod: "oauth_device_code", + nativeConnectQuery: "ChatGPT Codex", + docsUrl: "https://chatgpt.com", + tier: "promoted", + }, + { + id: "openai", + displayName: "OpenAI", + category: "model", + description: "GPT and o-series models", + setupMethod: "config_fields", + envVar: "OPENAI_API_KEY", + fields: [ + { + key: "OPENAI_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://platform.openai.com/api-keys", + tier: "promoted", + }, + { + id: "mistral", + displayName: "Mistral AI", + category: "model", + description: "Frontier models from Mistral AI", + setupMethod: "single_api_key", + envVar: "MISTRAL_API_KEY", + fields: [ + { + key: "MISTRAL_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://console.mistral.ai/api-keys", + tier: "promoted", + }, + { + id: "ollama", + displayName: "Ollama", + category: "model", + description: "Run local or self-hosted models", + setupMethod: "config_fields", + fields: [ + { + key: "OLLAMA_HOST", + label: "Host", + secret: false, + required: true, + placeholder: "localhost or http://localhost:11434", + defaultValue: "http://localhost:11434", + }, + ], + docsUrl: "https://ollama.com", + tier: "promoted", + }, + { + id: "openrouter", + displayName: "OpenRouter", + category: "model", + description: "Unified API for many models", + setupMethod: "single_api_key", + envVar: "OPENROUTER_API_KEY", + fields: [ + { + key: "OPENROUTER_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://openrouter.ai/keys", + tier: "promoted", + }, + { + id: "databricks", + displayName: "Databricks", + category: "model", + description: "Databricks Foundation Models", + setupMethod: "host_with_oauth_fallback", + fields: [ + { + key: "DATABRICKS_HOST", + label: "Host URL", + secret: false, + required: true, + placeholder: "https://dbc-...cloud.databricks.com", + }, + { + key: "DATABRICKS_TOKEN", + label: "Access Token", + secret: true, + required: false, + placeholder: "Paste your access token", + }, + ], + tier: "promoted", + }, + { + id: "github_copilot", + displayName: "GitHub Copilot Models", + category: "model", + description: "Models via GitHub Copilot subscription", + setupMethod: "oauth_device_code", + nativeConnectQuery: "GitHub Copilot", + tier: "promoted", + }, + { + id: "custom_deepseek", + displayName: "DeepSeek", + category: "model", + description: "DeepSeek chat and reasoning models", + setupMethod: "single_api_key", + envVar: "DEEPSEEK_API_KEY", + fields: [ + { + key: "DEEPSEEK_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://platform.deepseek.com/api_keys", + tier: "advanced", + }, + { + id: "xai", + displayName: "xAI", + category: "model", + description: "Grok models", + setupMethod: "single_api_key", + envVar: "XAI_API_KEY", + fields: [ + { + key: "XAI_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + tier: "advanced", + }, + { + id: "groq", + displayName: "Groq", + category: "model", + description: "Fast inference with Groq hardware", + setupMethod: "single_api_key", + envVar: "GROQ_API_KEY", + fields: [ + { + key: "GROQ_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://console.groq.com/keys", + tier: "advanced", + }, + { + id: "azure_openai", + displayName: "Azure OpenAI", + category: "model", + description: "OpenAI models on Azure", + setupMethod: "config_fields", + fields: [ + { + key: "AZURE_OPENAI_ENDPOINT", + label: "Endpoint", + secret: false, + required: true, + placeholder: "https://your-resource.openai.azure.com", + }, + { + key: "AZURE_OPENAI_DEPLOYMENT_NAME", + label: "Deployment", + secret: false, + required: true, + placeholder: "gpt-4o", + }, + { + key: "AZURE_OPENAI_API_KEY", + label: "API Key", + secret: true, + required: false, + placeholder: "Paste your API key", + }, + ], + tier: "advanced", + }, + { + id: "aws_bedrock", + displayName: "AWS Bedrock", + category: "model", + description: "Models on AWS", + setupMethod: "cloud_credentials", + fields: [ + { + key: "AWS_REGION", + label: "AWS Region", + secret: false, + required: false, + placeholder: "us-west-2", + }, + ], + tier: "advanced", + }, + { + id: "gcp_vertex_ai", + displayName: "GCP Vertex AI", + category: "model", + description: "Models on Google Cloud", + setupMethod: "cloud_credentials", + fields: [ + { + key: "GCP_PROJECT_ID", + label: "Project ID", + secret: false, + required: true, + placeholder: "my-gcp-project", + }, + { + key: "GCP_LOCATION", + label: "Location", + secret: false, + required: true, + placeholder: "us-central1", + }, + ], + tier: "advanced", + }, + { + id: "litellm", + displayName: "LiteLLM", + category: "model", + description: "LiteLLM proxy gateway", + setupMethod: "config_fields", + envVar: "LITELLM_API_KEY", + fields: [ + { + key: "LITELLM_HOST", + label: "Host URL", + secret: false, + required: true, + placeholder: "https://your-proxy.example.com", + }, + { + key: "LITELLM_API_KEY", + label: "API Key", + secret: true, + required: false, + placeholder: "Paste your API key", + }, + ], + tier: "advanced", + }, + { + id: "lmstudio", + displayName: "LM Studio", + category: "model", + description: "Run local models with LM Studio", + setupMethod: "config_fields", + fields: [ + { + key: "LMSTUDIO_HOST", + label: "Host URL", + secret: false, + required: false, + placeholder: "http://localhost:1234/v1/chat/completions", + }, + ], + docsUrl: "https://lmstudio.ai/docs/app/api", + tier: "advanced", + }, + { + id: "nvidia", + displayName: "NVIDIA", + category: "model", + description: "Hosted NVIDIA NIM models", + setupMethod: "single_api_key", + envVar: "NVIDIA_API_KEY", + fields: [ + { + key: "NVIDIA_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://build.nvidia.com/models", + tier: "advanced", + }, + { + id: "cerebras", + displayName: "Cerebras", + category: "model", + description: "Fast inference on Cerebras wafer-scale engines", + setupMethod: "single_api_key", + envVar: "CEREBRAS_API_KEY", + fields: [ + { + key: "CEREBRAS_API_KEY", + label: "API Key", + secret: true, + required: true, + placeholder: "Paste your API key", + }, + ], + docsUrl: "https://cloud.cerebras.ai/platform", + tier: "advanced", + }, + { + id: "snowflake", + displayName: "Snowflake", + category: "model", + description: "Snowflake Cortex", + setupMethod: "config_fields", + fields: [ + { + key: "SNOWFLAKE_HOST", + label: "Host URL", + secret: false, + required: true, + placeholder: "https://your-account.snowflakecomputing.com", + }, + { + key: "SNOWFLAKE_TOKEN", + label: "Access Token", + secret: true, + required: true, + placeholder: "Paste your access token", + }, + ], + tier: "advanced", + }, +]; diff --git a/ui/goose2/src/features/providers/ui/CustomHeadersEditor.tsx b/ui/goose2/src/features/providers/ui/CustomHeadersEditor.tsx new file mode 100644 index 00000000..fb88beea --- /dev/null +++ b/ui/goose2/src/features/providers/ui/CustomHeadersEditor.tsx @@ -0,0 +1,104 @@ +import { useTranslation } from "react-i18next"; +import { Button } from "@/shared/ui/button"; +import { Input } from "@/shared/ui/input"; +import { IconPlus, IconTrash } from "@tabler/icons-react"; +import { createCustomProviderHeaderDraft } from "@/features/providers/lib/customProviderHeaders"; + +export interface CustomHeader { + id: string; + key: string; + value: string; +} + +interface CustomHeadersEditorProps { + headers: CustomHeader[]; + onChange: (headers: CustomHeader[]) => void; + disabled?: boolean; +} + +export function CustomHeadersEditor({ + headers, + onChange, + disabled = false, +}: CustomHeadersEditorProps) { + const { t } = useTranslation("settings"); + + function updateHeader( + index: number, + field: keyof CustomHeader, + value: string, + ) { + onChange( + headers.map((header, currentIndex) => + currentIndex === index ? { ...header, [field]: value } : header, + ), + ); + } + + function removeHeader(index: number) { + onChange(headers.filter((_, currentIndex) => currentIndex !== index)); + } + + return ( +
+ {headers.length > 0 ? ( +
+ {headers.map((header, index) => ( +
+ + updateHeader(index, "key", event.target.value) + } + placeholder={t("providers.custom.fields.headerKey")} + disabled={disabled} + spellCheck={false} + className="h-8 text-xs" + /> + + updateHeader(index, "value", event.target.value) + } + placeholder={t("providers.custom.fields.headerValue")} + disabled={disabled} + spellCheck={false} + className="h-8 text-xs" + /> + +
+ ))} +
+ ) : ( +

+ {t("providers.custom.emptyHeaders")} +

+ )} + + +
+ ); +} diff --git a/ui/goose2/src/features/providers/ui/CustomProviderChoice.tsx b/ui/goose2/src/features/providers/ui/CustomProviderChoice.tsx new file mode 100644 index 00000000..a1ae5a62 --- /dev/null +++ b/ui/goose2/src/features/providers/ui/CustomProviderChoice.tsx @@ -0,0 +1,71 @@ +import { useTranslation } from "react-i18next"; +import { Button } from "@/shared/ui/button"; +import { IconPencil, IconTrash } from "@tabler/icons-react"; + +export interface CustomProviderChoiceInfo { + providerId: string; + displayName: string; + description?: string; + modelCount: number; + configured: boolean; +} + +interface CustomProviderChoiceProps { + provider: CustomProviderChoiceInfo; + onEdit: () => void; + onDelete: () => void; + deleting?: boolean; +} + +export function CustomProviderChoice({ + provider, + onEdit, + onDelete, + deleting = false, +}: CustomProviderChoiceProps) { + const { t } = useTranslation("settings"); + + return ( +
+
+ {provider.displayName.charAt(0).toUpperCase()} +
+ +
+
+

{provider.displayName}

+ {!provider.configured ? ( + + {t("providers.custom.notConfigured")} + + ) : null} +
+
+ + + +
+ ); +} diff --git a/ui/goose2/src/features/providers/ui/CustomProviderDialog.tsx b/ui/goose2/src/features/providers/ui/CustomProviderDialog.tsx new file mode 100644 index 00000000..91ebcbf8 --- /dev/null +++ b/ui/goose2/src/features/providers/ui/CustomProviderDialog.tsx @@ -0,0 +1,306 @@ +import { useEffect, useMemo, useRef, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { + Dialog, + DialogContent, + DialogDescription, + DialogHeader, + DialogTitle, +} from "@/shared/ui/dialog"; +import { Button } from "@/shared/ui/button"; +import { + IconArrowLeft, + IconLayoutGrid, + IconSettings, +} from "@tabler/icons-react"; +import { + CustomProviderForm, + type CustomProviderFormValues, + type ProviderTemplate, +} from "./CustomProviderForm"; +import { ProviderTemplatePicker } from "./ProviderTemplatePicker"; + +export type CustomProviderMutationInput = Omit< + CustomProviderFormValues, + "providerId" +> & { + providerId?: string; +}; + +interface CustomProviderDialogProps { + open: boolean; + mode: "create" | "edit"; + provider?: CustomProviderFormValues | null; + templates?: ProviderTemplate[]; + onOpenChange: (open: boolean) => void; + onCreate: (input: CustomProviderMutationInput) => Promise; + onUpdate: ( + providerId: string, + input: CustomProviderMutationInput, + ) => Promise; + onDelete?: (providerId: string) => Promise; +} + +const EMPTY_FORM: CustomProviderFormValues = { + displayName: "", + engine: "openai_compatible", + apiUrl: "", + basePath: "", + requiresAuth: true, + apiKey: "", + apiKeySet: false, + models: [], + authInitiallyEnabled: true, + supportsStreaming: true, + headers: [], +}; + +type CreateStep = "choice" | "template" | "form"; + +function valueFromTemplate( + template: ProviderTemplate, +): CustomProviderFormValues { + return { + ...EMPTY_FORM, + displayName: template.displayName, + engine: template.engine, + apiUrl: template.apiUrl, + basePath: template.basePath ?? "", + requiresAuth: template.requiresAuth, + authInitiallyEnabled: template.requiresAuth, + models: template.models, + supportsStreaming: template.supportsStreaming, + headers: template.headers, + catalogProviderId: template.id, + }; +} + +export function CustomProviderDialog({ + open, + mode, + provider, + templates = [], + onOpenChange, + onCreate, + onUpdate, + onDelete, +}: CustomProviderDialogProps) { + const { t } = useTranslation("settings"); + const [value, setValue] = useState(EMPTY_FORM); + const [selectedTemplateId, setSelectedTemplateId] = useState( + null, + ); + const [createStep, setCreateStep] = useState("choice"); + const [saving, setSaving] = useState(false); + const [deleting, setDeleting] = useState(false); + const [error, setError] = useState(""); + const openStateKeyRef = useRef(null); + const templateById = useMemo( + () => new Map(templates.map((template) => [template.id, template])), + [templates], + ); + + useEffect(() => { + if (!open) { + openStateKeyRef.current = null; + return; + } + const openStateKey = `${mode}:${provider?.providerId ?? "new"}`; + if (openStateKeyRef.current === openStateKey) { + return; + } + openStateKeyRef.current = openStateKey; + setValue(provider ?? EMPTY_FORM); + setSelectedTemplateId(provider?.catalogProviderId ?? null); + setCreateStep(mode === "create" ? "choice" : "form"); + setSaving(false); + setDeleting(false); + setError(""); + }, [mode, open, provider]); + + function handleStartManual() { + setSelectedTemplateId(null); + setValue(EMPTY_FORM); + setCreateStep("form"); + } + + function handleSelectTemplate(templateId: string) { + setSelectedTemplateId(templateId); + const template = templateById.get(templateId); + setValue(template ? valueFromTemplate(template) : EMPTY_FORM); + setCreateStep("form"); + } + + function handleBack() { + setError(""); + if (createStep === "template") { + setCreateStep("choice"); + return; + } + if (selectedTemplateId) { + setCreateStep("template"); + return; + } + setCreateStep("choice"); + } + + async function handleSubmit() { + setSaving(true); + setError(""); + try { + if (mode === "edit" && value.providerId) { + await onUpdate(value.providerId, value); + } else { + await onCreate(value); + } + onOpenChange(false); + } catch (nextError) { + setError( + nextError instanceof Error + ? nextError.message + : t("providers.custom.errors.saveFailed"), + ); + } finally { + setSaving(false); + } + } + + async function handleDelete() { + if (!value.providerId || !onDelete) { + return; + } + + setDeleting(true); + setError(""); + try { + const deleted = await onDelete(value.providerId); + if (deleted === false) { + return; + } + onOpenChange(false); + } catch (nextError) { + setError( + nextError instanceof Error + ? nextError.message + : t("providers.custom.errors.deleteFailed"), + ); + } finally { + setDeleting(false); + } + } + + function renderCreateChoice() { + return ( +
+ + + +
+ ); + } + + function renderBackButton() { + if (mode !== "create" || createStep === "choice") { + return null; + } + + return ( +
+ +
+ ); + } + + function renderContent() { + if (mode === "create" && createStep === "choice") { + return renderCreateChoice(); + } + + if (mode === "create" && createStep === "template") { + return ( + <> + {renderBackButton()} + + + ); + } + + return ( + <> + {renderBackButton()} + void handleSubmit()} + onDelete={ + mode === "edit" && onDelete ? () => void handleDelete() : undefined + } + /> + + ); + } + + return ( + + + + + {mode === "edit" + ? t("providers.custom.editTitle") + : t("providers.custom.addTitle")} + + + {t("providers.custom.description")} + + + + {renderContent()} + + + ); +} diff --git a/ui/goose2/src/features/providers/ui/CustomProviderForm.tsx b/ui/goose2/src/features/providers/ui/CustomProviderForm.tsx new file mode 100644 index 00000000..ab43d373 --- /dev/null +++ b/ui/goose2/src/features/providers/ui/CustomProviderForm.tsx @@ -0,0 +1,384 @@ +import { useMemo, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { Button } from "@/shared/ui/button"; +import { Input } from "@/shared/ui/input"; +import { Label } from "@/shared/ui/label"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/shared/ui/select"; +import { Switch } from "@/shared/ui/switch"; +import { + IconDeviceFloppy, + IconEye, + IconEyeOff, + IconTrash, +} from "@tabler/icons-react"; +import type { CustomProviderEngine } from "@/features/providers/lib/customProviderTypes"; +import { + validateCustomProviderDraft, + type CustomProviderValidationField, + type CustomProviderValidationIssue, +} from "@/features/providers/lib/customProviderValidation"; +import { CustomHeadersEditor, type CustomHeader } from "./CustomHeadersEditor"; +import { ProviderModelListEditor } from "./ProviderModelListEditor"; + +export interface ProviderTemplate { + id: string; + displayName: string; + description?: string; + engine: string; + apiUrl: string; + basePath?: string; + requiresAuth: boolean; + supportsStreaming: boolean; + models: string[]; + headers: CustomHeader[]; +} + +export interface CustomProviderFormValues { + providerId?: string; + displayName: string; + engine: string; + apiUrl: string; + basePath: string; + requiresAuth: boolean; + apiKey: string; + apiKeySet: boolean; + models: string[]; + authInitiallyEnabled: boolean; + supportsStreaming: boolean; + headers: CustomHeader[]; + catalogProviderId?: string; +} + +interface CustomProviderFormProps { + value: CustomProviderFormValues; + mode: "create" | "edit"; + saving?: boolean; + deleting?: boolean; + error?: string; + onChange: (value: CustomProviderFormValues) => void; + onSubmit: () => void; + onDelete?: () => void; +} + +const ENGINE_OPTIONS: CustomProviderEngine[] = [ + "openai_compatible", + "anthropic_compatible", + "ollama_compatible", +]; + +function translationKey(key: string) { + return key.replace(/^settings\./, ""); +} + +function fieldIssues( + issues: CustomProviderValidationIssue[], + field: CustomProviderValidationField, +) { + return issues.filter((issue) => issue.field === field); +} + +export function CustomProviderForm({ + value, + mode, + saving = false, + deleting = false, + error = "", + onChange, + onSubmit, + onDelete, +}: CustomProviderFormProps) { + const { t } = useTranslation(["settings", "common"]); + const [apiKeyVisible, setApiKeyVisible] = useState(false); + const disabled = saving || deleting; + const validationIssues = useMemo( + () => + validateCustomProviderDraft({ + providerId: value.providerId, + editable: true, + engine: value.engine, + displayName: value.displayName, + apiUrl: value.apiUrl, + basePath: value.basePath, + apiKey: value.apiKey, + apiKeySet: value.apiKeySet, + modelsInput: value.models.join(", "), + models: value.models, + authInitiallyEnabled: value.authInitiallyEnabled, + requiresAuth: value.requiresAuth, + supportsStreaming: value.supportsStreaming, + headers: value.headers, + catalogProviderId: value.catalogProviderId, + }), + [value], + ); + const isValid = validationIssues.length === 0; + + function update(patch: Partial) { + onChange({ ...value, ...patch }); + } + + function renderFieldErrors(field: CustomProviderValidationField) { + const issues = fieldIssues(validationIssues, field); + if (issues.length === 0) { + return null; + } + + return ( +
+ {issues.map((issue) => ( +

+ {t(translationKey(issue.key))} +

+ ))} +
+ ); + } + + return ( +
{ + event.preventDefault(); + onSubmit(); + }} + > +
+
+ + update({ displayName: event.target.value })} + placeholder={t("providers.custom.fields.displayNamePlaceholder")} + disabled={disabled} + spellCheck={false} + className="h-8 text-xs" + /> + {renderFieldErrors("displayName")} +
+ +
+ + + {renderFieldErrors("engine")} +
+ +
+ + update({ apiUrl: event.target.value })} + placeholder={t("providers.custom.fields.apiUrlPlaceholder")} + disabled={disabled} + spellCheck={false} + className="h-8 text-xs" + /> + {renderFieldErrors("apiUrl")} +
+ +
+ + update({ basePath: event.target.value })} + placeholder={t("providers.custom.fields.basePathPlaceholder")} + disabled={disabled} + spellCheck={false} + className="h-8 text-xs" + /> +
+
+ +
+
+
+ +

+ {t("providers.custom.fields.requiresAuthDescription")} +

+
+ update({ requiresAuth })} + disabled={disabled} + /> +
+ + {value.requiresAuth ? ( +
+ +
+ update({ apiKey: event.target.value })} + placeholder={ + mode === "edit" && value.apiKeySet + ? t("providers.custom.fields.apiKeyEditPlaceholder") + : t("providers.custom.fields.apiKeyPlaceholder") + } + disabled={disabled} + spellCheck={false} + autoComplete="new-password" + data-1p-ignore + data-lpignore + className="h-8 text-xs" + /> + +
+ {renderFieldErrors("apiKey")} +
+ ) : null} +
+ +
+ + {t("providers.custom.fields.models")} + + update({ models })} + disabled={disabled} + /> + {renderFieldErrors("models")} +
+ +
+
+
+ +

+ {t("providers.custom.fields.supportsStreamingDescription")} +

+
+ + update({ supportsStreaming }) + } + disabled={disabled} + /> +
+
+ +
+ + {t("providers.custom.fields.headers")} + + update({ headers })} + disabled={disabled} + /> + {renderFieldErrors("headers")} +
+ + {error ? ( +

+ {error} +

+ ) : null} + +
+ {mode === "edit" && onDelete ? ( + + ) : ( + + )} + + +
+
+ ); +} diff --git a/ui/goose2/src/features/providers/ui/ProviderModelListEditor.tsx b/ui/goose2/src/features/providers/ui/ProviderModelListEditor.tsx new file mode 100644 index 00000000..a401d70b --- /dev/null +++ b/ui/goose2/src/features/providers/ui/ProviderModelListEditor.tsx @@ -0,0 +1,103 @@ +import { useState } from "react"; +import { useTranslation } from "react-i18next"; +import { Button } from "@/shared/ui/button"; +import { Input } from "@/shared/ui/input"; +import { IconPlus, IconX } from "@tabler/icons-react"; + +interface ProviderModelListEditorProps { + models: string[]; + onChange: (models: string[]) => void; + disabled?: boolean; +} + +function normalizeModels(values: string[]) { + return [...new Set(values.map((value) => value.trim()).filter(Boolean))]; +} + +function splitModelInput(value: string) { + return value.split(/[\n,]/); +} + +export function ProviderModelListEditor({ + models, + onChange, + disabled = false, +}: ProviderModelListEditorProps) { + const { t } = useTranslation("settings"); + const [draft, setDraft] = useState(""); + + function addModels(value: string) { + const nextModels = normalizeModels([...models, ...splitModelInput(value)]); + onChange(nextModels); + setDraft(""); + } + + function removeModel(model: string) { + onChange(models.filter((item) => item !== model)); + } + + return ( +
+
+ setDraft(event.target.value)} + onKeyDown={(event) => { + if (event.nativeEvent.isComposing) { + return; + } + if (event.key === "Enter" || event.key === ",") { + event.preventDefault(); + addModels(draft); + } + }} + onPaste={(event) => { + const pasted = event.clipboardData.getData("text"); + if (/[\n,]/.test(pasted)) { + event.preventDefault(); + addModels(pasted); + } + }} + placeholder={t("providers.custom.fields.modelsPlaceholder")} + disabled={disabled} + spellCheck={false} + className="h-8 text-xs" + /> + +
+ + {models.length > 0 ? ( +
+ {models.map((model) => ( + + {model} + + + ))} +
+ ) : null} +
+ ); +} diff --git a/ui/goose2/src/features/providers/ui/ProviderTemplatePicker.tsx b/ui/goose2/src/features/providers/ui/ProviderTemplatePicker.tsx new file mode 100644 index 00000000..2fca22e9 --- /dev/null +++ b/ui/goose2/src/features/providers/ui/ProviderTemplatePicker.tsx @@ -0,0 +1,133 @@ +import { useMemo, useState } from "react"; +import { useTranslation } from "react-i18next"; +import { SearchBar } from "@/shared/ui/SearchBar"; +import { ScrollArea } from "@/shared/ui/scroll-area"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/shared/ui/select"; +import { IconLayoutGrid } from "@tabler/icons-react"; +import type { CustomProviderEngine } from "@/features/providers/lib/customProviderTypes"; +import type { ProviderTemplate } from "./CustomProviderForm"; + +interface ProviderTemplatePickerProps { + templates: ProviderTemplate[]; + onSelect: (templateId: string) => void; + disabled?: boolean; +} + +type CompatibilityFilter = "all" | CustomProviderEngine; + +function matchesTemplateSearch(template: ProviderTemplate, query: string) { + const needle = query.trim().toLowerCase(); + if (!needle) { + return true; + } + + return [ + template.displayName, + template.description ?? "", + template.engine, + template.id, + ...template.models, + ] + .join(" ") + .toLowerCase() + .includes(needle); +} + +export function ProviderTemplatePicker({ + templates, + onSelect, + disabled = false, +}: ProviderTemplatePickerProps) { + const { t } = useTranslation("settings"); + const [query, setQuery] = useState(""); + const [compatibility, setCompatibility] = + useState("all"); + const filteredTemplates = useMemo( + () => + templates.filter((template) => { + const matchesCompatibility = + compatibility === "all" || template.engine === compatibility; + return matchesCompatibility && matchesTemplateSearch(template, query); + }), + [compatibility, query, templates], + ); + + return ( +
+
+ + +
+ + +
+ {filteredTemplates.map((template) => ( + + ))} + + {filteredTemplates.length === 0 ? ( +

+ {t("providers.custom.templates.empty")} +

+ ) : null} +
+
+
+ ); +} diff --git a/ui/goose2/src/features/settings/ui/AgentProviderCard.tsx b/ui/goose2/src/features/settings/ui/AgentProviderCard.tsx index adf8dc3a..cc093dc4 100644 --- a/ui/goose2/src/features/settings/ui/AgentProviderCard.tsx +++ b/ui/goose2/src/features/settings/ui/AgentProviderCard.tsx @@ -425,10 +425,14 @@ export function AgentProviderCard({ provider }: AgentProviderCardProps) { >
-
- {icon} -
- {provider.displayName} + {icon ? ( +
+ {icon} +
+ ) : null} + + {provider.displayName} +

{provider.description}

diff --git a/ui/goose2/src/features/settings/ui/ModelProviderRow.tsx b/ui/goose2/src/features/settings/ui/ModelProviderRow.tsx index 1c9c7ed8..fc00c0e5 100644 --- a/ui/goose2/src/features/settings/ui/ModelProviderRow.tsx +++ b/ui/goose2/src/features/settings/ui/ModelProviderRow.tsx @@ -468,13 +468,17 @@ export function ModelProviderRow({ disabled={authenticating} className="flex w-full items-center gap-3 rounded-lg border border-border px-3 py-2.5 text-left transition-colors hover:bg-accent/30 focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring disabled:cursor-default disabled:hover:bg-transparent" > -
- {icon || ( + {icon ? ( +
+ {icon} +
+ ) : ( +
{formatProviderLabel(provider.id).charAt(0)} - )} -
+
+ )} {provider.displayName} diff --git a/ui/goose2/src/features/settings/ui/ProvidersSettings.tsx b/ui/goose2/src/features/settings/ui/ProvidersSettings.tsx index ca007742..76818217 100644 --- a/ui/goose2/src/features/settings/ui/ProvidersSettings.tsx +++ b/ui/goose2/src/features/settings/ui/ProvidersSettings.tsx @@ -1,16 +1,46 @@ import { useEffect, useMemo, useState } from "react"; import { useTranslation } from "react-i18next"; -import { Button } from "@/shared/ui/button"; +import { Button, buttonVariants } from "@/shared/ui/button"; +import { + AlertDialog, + AlertDialogAction, + AlertDialogCancel, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "@/shared/ui/alert-dialog"; import { Separator } from "@/shared/ui/separator"; import { Spinner } from "@/shared/ui/spinner"; -import { IconChevronDown } from "@tabler/icons-react"; +import { IconChevronDown, IconPlus } from "@tabler/icons-react"; import { getAgentProviders, getModelProviders, } from "@/features/providers/providerCatalog"; import { useCredentials } from "@/features/providers/hooks/useCredentials"; +import { useCustomProviders } from "@/features/providers/hooks/useCustomProviders"; +import { + CustomProviderChoice, + type CustomProviderChoiceInfo, +} from "@/features/providers/ui/CustomProviderChoice"; +import { + CustomProviderDialog, + type CustomProviderMutationInput, +} from "@/features/providers/ui/CustomProviderDialog"; +import type { + CustomProviderFormValues, + ProviderTemplate, +} from "@/features/providers/ui/CustomProviderForm"; +import { useProviderInventoryStore } from "@/features/providers/stores/providerInventoryStore"; import { AgentProviderCard } from "./AgentProviderCard"; import { ModelProviderRow } from "./ModelProviderRow"; +import { + catalogEntryToTemplate, + formValueToDraft, + readResponseToFormValue, + templateToFormValue, +} from "./customProviderFormAdapters"; import type { ProviderDisplayInfo, ProviderSetupStatus, @@ -37,10 +67,55 @@ function toDisplayInfo( })); } +function isCustomProviderEntry(entry: { + providerId: string; + providerType?: string; +}) { + return entry.providerType === "Custom"; +} + +function toCustomProviderChoiceInfo(entry: { + providerId: string; + providerName: string; + description?: string; + configured: boolean; + models: unknown[]; +}): CustomProviderChoiceInfo { + return { + providerId: entry.providerId, + displayName: entry.providerName, + description: entry.description || undefined, + configured: entry.configured, + modelCount: entry.models.length, + }; +} + +interface PendingCustomProviderDelete { + providerId: string; + displayName: string; + resolve: (deleted: boolean) => void; + reject: (error: unknown) => void; +} + export function ProvidersSettings() { const { t } = useTranslation(["settings", "common"]); const [showAllModels, setShowAllModels] = useState(false); const [modelOrder, setModelOrder] = useState(null); + const [customDialogOpen, setCustomDialogOpen] = useState(false); + const [customDialogMode, setCustomDialogMode] = useState<"create" | "edit">( + "create", + ); + const [customProviderDraft, setCustomProviderDraft] = + useState(null); + const [customProviderTemplates, setCustomProviderTemplates] = useState< + ProviderTemplate[] + >([]); + const [customProviderError, setCustomProviderError] = useState(""); + const [customProviderDeleteError, setCustomProviderDeleteError] = + useState(""); + const [pendingCustomProviderDelete, setPendingCustomProviderDelete] = + useState(null); + const inventoryEntries = useProviderInventoryStore((state) => state.entries); const { configuredIds, @@ -53,6 +128,7 @@ export function ProvidersSettings() { remove, completeNativeSetup, } = useCredentials(); + const customProvidersApi = useCustomProviders(); const agents = useMemo( () => toDisplayInfo(getAgentProviders(), configuredIds), @@ -112,6 +188,123 @@ export function ProvidersSettings() { const advancedModels = orderedModels.filter((m) => m.tier === "advanced"); const visibleModels = showAllModels ? orderedModels : promotedModels; + const customProviders = useMemo( + () => + [...inventoryEntries.values()] + .filter(isCustomProviderEntry) + .map(toCustomProviderChoiceInfo) + .sort((a, b) => a.displayName.localeCompare(b.displayName)), + [inventoryEntries], + ); + + async function loadTemplates() { + try { + setCustomProviderError(""); + const catalog = await customProvidersApi.loadCatalog(); + const templates = await Promise.all( + catalog.map(async (entry) => { + try { + return templateToFormValue( + await customProvidersApi.getTemplate(entry.providerId), + ); + } catch { + return catalogEntryToTemplate(entry); + } + }), + ); + setCustomProviderTemplates(templates); + } catch (error) { + setCustomProviderTemplates([]); + setCustomProviderError( + error instanceof Error + ? error.message + : t("providers.custom.errors.templatesFailed"), + ); + } + } + + async function openCreateCustomProvider() { + setCustomProviderError(""); + setCustomProviderDeleteError(""); + setCustomDialogMode("create"); + setCustomProviderDraft(null); + setCustomDialogOpen(true); + await loadTemplates(); + } + + async function openEditCustomProvider(providerId: string) { + setCustomProviderError(""); + setCustomProviderDeleteError(""); + try { + const provider = readResponseToFormValue( + await customProvidersApi.read(providerId), + ); + setCustomDialogMode("edit"); + setCustomProviderDraft(provider); + setCustomDialogOpen(true); + await loadTemplates(); + } catch (error) { + setCustomProviderError( + error instanceof Error + ? error.message + : t("providers.custom.errors.loadFailed"), + ); + } + } + + async function createCustomProvider(input: CustomProviderMutationInput) { + await customProvidersApi.saveDraft(formValueToDraft(input)); + } + + async function updateCustomProvider( + providerId: string, + input: CustomProviderMutationInput, + ) { + await customProvidersApi.saveDraft(formValueToDraft(input), { providerId }); + } + + async function deleteCustomProvider(providerId: string) { + const providerName = + customProviders.find((provider) => provider.providerId === providerId) + ?.displayName ?? providerId; + + return new Promise((resolve, reject) => { + setPendingCustomProviderDelete({ + providerId, + displayName: providerName, + resolve, + reject, + }); + }); + } + + function cancelCustomProviderDelete() { + pendingCustomProviderDelete?.resolve(false); + setPendingCustomProviderDelete(null); + } + + async function confirmCustomProviderDelete() { + const pendingDelete = pendingCustomProviderDelete; + if (!pendingDelete) { + return; + } + + setCustomProviderDeleteError(""); + try { + await customProvidersApi.remove(pendingDelete.providerId); + pendingDelete.resolve(true); + setPendingCustomProviderDelete(null); + } catch (error) { + setCustomProviderDeleteError( + error instanceof Error + ? error.message + : t("providers.custom.errors.deleteFailed"), + ); + pendingDelete.reject(error); + setPendingCustomProviderDelete(null); + } + } + return (

@@ -141,22 +334,71 @@ export function ProvidersSettings() {
-
-

- {t("providers.models.title")} -

- {loading ? ( - - - {t("providers.models.checkingStatus")} - - ) : null} +
+
+
+

+ {t("providers.models.title")} +

+ {loading ? ( + + + {t("providers.models.checkingStatus")} + + ) : null} +
+

+ {t("providers.models.description")} +

+
+
-

- {t("providers.models.description")} -

+ {customProviderError ? ( +

+ {customProviderError} +

+ ) : null} + {customProviderDeleteError ? ( +

+ {customProviderDeleteError} +

+ ) : null} + + {customProviders.length > 0 ? ( +
+ {customProviders.map((provider) => ( + void openEditCustomProvider(provider.providerId)} + onDelete={() => + void deleteCustomProvider(provider.providerId).catch(() => {}) + } + deleting={customProvidersApi.deletingProviderIds.has( + provider.providerId, + )} + /> + ))} +
+ ) : null} +
{visibleModels.map((model) => ( )}
+ + + + { + if (!open) { + cancelCustomProviderDelete(); + } + }} + > + + + + {t("providers.custom.confirmDeleteTitle", { + name: pendingCustomProviderDelete?.displayName ?? "", + })} + + + {t("providers.custom.confirmDelete", { + name: pendingCustomProviderDelete?.displayName ?? "", + })} + + + + {t("common:actions.cancel")} + { + event.preventDefault(); + void confirmCustomProviderDelete(); + }} + > + {t("common:actions.delete")} + + + +

); } diff --git a/ui/goose2/src/features/settings/ui/SettingsModal.tsx b/ui/goose2/src/features/settings/ui/SettingsModal.tsx index a44358f7..f1a71f42 100644 --- a/ui/goose2/src/features/settings/ui/SettingsModal.tsx +++ b/ui/goose2/src/features/settings/ui/SettingsModal.tsx @@ -154,9 +154,8 @@ export function SettingsModal({ {/* biome-ignore lint/a11y/noStaticElementInteractions: click handler only prevents backdrop dismiss propagation */}
e.stopPropagation()} > diff --git a/ui/goose2/src/features/settings/ui/__tests__/ProvidersSettings.test.tsx b/ui/goose2/src/features/settings/ui/__tests__/ProvidersSettings.test.tsx index 66407c25..618cb499 100644 --- a/ui/goose2/src/features/settings/ui/__tests__/ProvidersSettings.test.tsx +++ b/ui/goose2/src/features/settings/ui/__tests__/ProvidersSettings.test.tsx @@ -1,18 +1,48 @@ import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import type { ProviderInventoryEntryDto } from "@aaif/goose-sdk"; import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useProviderInventoryStore } from "@/features/providers/stores/providerInventoryStore"; import { ProvidersSettings } from "../ProvidersSettings"; const mocks = vi.hoisted(() => ({ useCredentials: vi.fn(), + useCustomProviders: vi.fn(), })); vi.mock("@/features/providers/hooks/useCredentials", () => ({ useCredentials: () => mocks.useCredentials(), })); +vi.mock("@/features/providers/hooks/useCustomProviders", () => ({ + useCustomProviders: () => mocks.useCustomProviders(), +})); + +function providerEntry( + overrides: Partial, +): ProviderInventoryEntryDto { + return { + providerId: "custom_openai", + providerName: "Custom OpenAI", + description: "", + defaultModel: "", + configured: true, + providerType: "Custom", + configKeys: [], + setupSteps: [], + supportsRefresh: true, + refreshing: false, + models: [], + stale: false, + ...overrides, + }; +} + describe("ProvidersSettings", () => { beforeEach(() => { + vi.restoreAllMocks(); vi.clearAllMocks(); + useProviderInventoryStore.getState().setEntries([]); mocks.useCredentials.mockReturnValue({ configuredIds: new Set(), loading: false, @@ -25,6 +55,27 @@ describe("ProvidersSettings", () => { remove: vi.fn(), completeNativeSetup: vi.fn(), }); + mocks.useCustomProviders.mockReturnValue({ + catalog: [], + catalogLoading: false, + saving: false, + savingProviderIds: new Set(), + deletingProviderIds: new Set(), + syncingProviderIds: new Set(), + inventoryWarnings: new Map(), + statusByProviderId: new Map(), + configuredIds: new Set(), + loadCatalog: vi.fn().mockResolvedValue([]), + getTemplate: vi.fn(), + read: vi.fn(), + create: vi.fn(), + update: vi.fn(), + remove: vi.fn().mockResolvedValue({ + providerId: "custom_acme", + refresh: { started: [], skipped: [] }, + }), + saveDraft: vi.fn(), + }); }); it("does not show the restart banner for provider credential changes", () => { @@ -88,4 +139,103 @@ describe("ProvidersSettings", () => { Node.DOCUMENT_POSITION_FOLLOWING, ).toBeTruthy(); }); + + it("shows the custom provider entry point near model providers", async () => { + const user = userEvent.setup(); + render(); + + await user.click( + screen.getByRole("button", { name: /add custom provider/i }), + ); + + expect( + screen.getByRole("dialog", { name: /add custom provider/i }), + ).toBeInTheDocument(); + expect(screen.getByText(/fully custom/i)).toBeInTheDocument(); + expect(screen.getByText(/use a template/i)).toBeInTheDocument(); + }); + + it("shows custom inventory providers with edit and delete actions", () => { + useProviderInventoryStore.getState().setEntries([ + providerEntry({ + providerId: "custom_acme", + providerName: "Acme Models", + models: [ + { + id: "acme-fast", + name: "acme-fast", + }, + ], + }), + ]); + + render(); + + expect(screen.getByText("Acme Models")).toBeInTheDocument(); + expect(screen.queryByText("1 model")).not.toBeInTheDocument(); + expect( + screen.getByRole("button", { name: /edit acme models/i }), + ).toBeInTheDocument(); + expect( + screen.getByRole("button", { name: /delete acme models/i }), + ).toBeInTheDocument(); + }); + + it("confirms before deleting a custom provider", async () => { + const user = userEvent.setup(); + const remove = vi.fn().mockResolvedValue({ + providerId: "custom_acme", + refresh: { started: [], skipped: [] }, + }); + mocks.useCustomProviders.mockReturnValue({ + ...mocks.useCustomProviders(), + remove, + }); + useProviderInventoryStore.getState().setEntries([ + providerEntry({ + providerId: "custom_acme", + providerName: "Acme Models", + }), + ]); + + render(); + + await user.click( + screen.getByRole("button", { name: /delete acme models/i }), + ); + expect( + screen.getByRole("alertdialog", { name: /delete acme models/i }), + ).toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: /cancel/i })); + + expect(remove).not.toHaveBeenCalled(); + }); + + it("keeps a provider visible and shows an error when delete fails", async () => { + const user = userEvent.setup(); + const remove = vi.fn().mockRejectedValue(new Error("delete exploded")); + mocks.useCustomProviders.mockReturnValue({ + ...mocks.useCustomProviders(), + remove, + }); + useProviderInventoryStore.getState().setEntries([ + providerEntry({ + providerId: "custom_acme", + providerName: "Acme Models", + }), + ]); + + render(); + + await user.click( + screen.getByRole("button", { name: /delete acme models/i }), + ); + await user.click(screen.getByRole("button", { name: /^delete$/i })); + + expect(await screen.findByRole("alert")).toHaveTextContent( + "delete exploded", + ); + expect(screen.getByText("Acme Models")).toBeInTheDocument(); + }); }); diff --git a/ui/goose2/src/features/settings/ui/customProviderFormAdapters.ts b/ui/goose2/src/features/settings/ui/customProviderFormAdapters.ts new file mode 100644 index 00000000..7d847391 --- /dev/null +++ b/ui/goose2/src/features/settings/ui/customProviderFormAdapters.ts @@ -0,0 +1,106 @@ +import { normalizeCustomProviderEngine } from "@/features/providers/lib/customProviderDraft"; +import { recordToHeaderDrafts } from "@/features/providers/lib/customProviderHeaders"; +import { + formatCustomProviderModels, + parseCustomProviderModels, +} from "@/features/providers/lib/customProviderModels"; +import type { + CustomProviderDraft, + CustomProviderEngine, + CustomProviderReadResponse, + ProviderCatalogEntryDto, + ProviderTemplateDto, +} from "@/features/providers/lib/customProviderTypes"; +import type { CustomProviderMutationInput } from "@/features/providers/ui/CustomProviderDialog"; +import type { + CustomProviderFormValues, + ProviderTemplate, +} from "@/features/providers/ui/CustomProviderForm"; + +function engineForCustomProviderFormat(format: string): CustomProviderEngine { + if (format === "anthropic") { + return "anthropic_compatible"; + } + if (format === "ollama") { + return "ollama_compatible"; + } + return "openai_compatible"; +} + +export function templateToFormValue( + template: ProviderTemplateDto, +): ProviderTemplate { + const models = (template.models ?? []) + .filter((model) => !model.deprecated) + .map((model) => model.id); + + return { + id: template.providerId, + displayName: template.name, + engine: engineForCustomProviderFormat(template.format), + apiUrl: template.apiUrl, + requiresAuth: true, + supportsStreaming: template.supportsStreaming, + models, + headers: [], + }; +} + +export function catalogEntryToTemplate( + entry: ProviderCatalogEntryDto, +): ProviderTemplate { + return { + id: entry.providerId, + displayName: entry.name, + engine: engineForCustomProviderFormat(entry.format), + apiUrl: entry.apiUrl, + requiresAuth: true, + supportsStreaming: true, + models: [], + headers: [], + }; +} + +export function readResponseToFormValue( + response: CustomProviderReadResponse, +): CustomProviderFormValues { + const provider = response.provider; + return { + providerId: provider.providerId, + displayName: provider.displayName, + engine: normalizeCustomProviderEngine(provider.engine), + apiUrl: provider.apiUrl, + basePath: provider.basePath ?? "", + requiresAuth: provider.requiresAuth, + apiKey: "", + apiKeySet: provider.apiKeySet, + models: parseCustomProviderModels(provider.models ?? []), + authInitiallyEnabled: provider.requiresAuth, + supportsStreaming: provider.supportsStreaming ?? true, + headers: recordToHeaderDrafts(provider.headers), + catalogProviderId: provider.catalogProviderId ?? undefined, + }; +} + +export function formValueToDraft( + input: CustomProviderMutationInput, +): CustomProviderDraft { + const models = parseCustomProviderModels(input.models); + return { + providerId: input.providerId, + editable: true, + engine: input.engine, + displayName: input.displayName, + apiUrl: input.apiUrl, + basePath: input.basePath, + apiKey: input.apiKey, + apiKeySet: input.apiKeySet, + modelsInput: formatCustomProviderModels(models), + models, + authInitiallyEnabled: input.authInitiallyEnabled, + requiresAuth: input.requiresAuth, + supportsStreaming: input.supportsStreaming, + headers: input.headers, + catalogProviderId: input.catalogProviderId, + }; +} diff --git a/ui/goose2/src/shared/i18n/locales/en/settings.json b/ui/goose2/src/shared/i18n/locales/en/settings.json index 93a5aff9..9a5ec35a 100644 --- a/ui/goose2/src/shared/i18n/locales/en/settings.json +++ b/ui/goose2/src/shared/i18n/locales/en/settings.json @@ -224,6 +224,94 @@ }, "title": "Agent harnesses" }, + "custom": { + "actions": { + "addHeader": "Add header", + "addModel": "Add model", + "back": "Back", + "create": "Create provider", + "delete": "Delete", + "deleteProvider": "Delete {{name}}", + "deleting": "Deleting...", + "editProvider": "Edit {{name}}", + "hideApiKey": "Hide API key", + "removeHeader": "Remove header", + "removeModel": "Remove {{model}}", + "save": "Save changes", + "saving": "Saving...", + "showApiKey": "Show API key" + }, + "addButton": "Add custom provider", + "addTitle": "Add custom provider", + "confirmDelete": "This removes {{name}} and any API key owned by it.", + "confirmDeleteTitle": "Delete {{name}}?", + "description": "Connect an OpenAI-compatible, Anthropic-compatible, or local model endpoint.", + "editTitle": "Edit custom provider", + "emptyHeaders": "No custom headers.", + "emptyModels": "Add at least one model.", + "engines": { + "anthropic_compatible": "Anthropic-compatible", + "ollama_compatible": "Ollama-compatible", + "openai_compatible": "OpenAI-compatible" + }, + "errors": { + "deleteFailed": "Failed to delete custom provider.", + "loadFailed": "Failed to load custom provider.", + "saveFailed": "Failed to save custom provider.", + "templatesFailed": "Failed to load provider templates." + }, + "fields": { + "apiKey": "API key", + "apiKeyEditPlaceholder": "Leave blank to keep the saved key", + "apiKeyPlaceholder": "Paste your API key", + "apiUrl": "API URL", + "apiUrlPlaceholder": "https://api.example.com/v1", + "basePath": "Base path", + "basePathPlaceholder": "/v1/chat/completions", + "displayName": "Display name", + "displayNamePlaceholder": "My provider", + "engine": "Engine", + "headerKey": "Header", + "headerValue": "Value", + "headers": "Custom headers", + "models": "Models", + "modelsPlaceholder": "gpt-4.1, claude-3-5-sonnet", + "requiresAuth": "Requires authentication", + "requiresAuthDescription": "Store an API key for requests to this provider.", + "supportsStreaming": "Supports streaming", + "supportsStreamingDescription": "Stream responses as the model generates tokens." + }, + "modelCount_one": "{{count}} model", + "modelCount_other": "{{count}} models", + "notConfigured": "Not configured", + "validation": { + "apiKeyRequired": "API key is required.", + "apiUrlInvalid": "Enter a valid HTTP or HTTPS URL.", + "apiUrlRequired": "API URL is required.", + "displayNameRequired": "Display name is required.", + "engineRequired": "Choose a provider engine.", + "headerDuplicate": "Header names must be unique.", + "headerNameInvalid": "Header names can only contain valid HTTP token characters.", + "headerNameRequired": "Header name is required.", + "headerValueRequired": "Header value is required.", + "modelsRequired": "Add at least one model." + }, + "sections": { + "template": "Start from template" + }, + "templates": { + "compatibility": { + "all": "All compatibility types" + }, + "clear": "Clear template", + "empty": "No templates available yet.", + "manual": "Fully custom", + "manualDescription": "Start from a blank provider configuration.", + "searchPlaceholder": "Search templates", + "useTemplate": "Use a template", + "useTemplateDescription": "Start with endpoint and model defaults for a known provider." + } + }, "disconnect": "Disconnect", "models": { "description": "AI models power your agents. Goose requires one to work, but some agents bring their own.", diff --git a/ui/goose2/src/shared/i18n/locales/es/settings.json b/ui/goose2/src/shared/i18n/locales/es/settings.json index 3f2e87ae..2ef8fe56 100644 --- a/ui/goose2/src/shared/i18n/locales/es/settings.json +++ b/ui/goose2/src/shared/i18n/locales/es/settings.json @@ -224,6 +224,94 @@ }, "title": "Arneses de agentes" }, + "custom": { + "actions": { + "addHeader": "Agregar encabezado", + "addModel": "Agregar modelo", + "back": "Atrás", + "create": "Crear proveedor", + "delete": "Eliminar", + "deleteProvider": "Eliminar {{name}}", + "deleting": "Eliminando...", + "editProvider": "Editar {{name}}", + "hideApiKey": "Ocultar clave API", + "removeHeader": "Eliminar encabezado", + "removeModel": "Eliminar {{model}}", + "save": "Guardar cambios", + "saving": "Guardando...", + "showApiKey": "Mostrar clave API" + }, + "addButton": "Agregar proveedor personalizado", + "addTitle": "Agregar proveedor personalizado", + "confirmDelete": "Esto elimina {{name}} y cualquier clave API que le pertenezca.", + "confirmDeleteTitle": "¿Eliminar {{name}}?", + "description": "Conecta un endpoint de modelo compatible con OpenAI, Anthropic o local.", + "editTitle": "Editar proveedor personalizado", + "emptyHeaders": "No hay encabezados personalizados.", + "emptyModels": "Agrega al menos un modelo.", + "engines": { + "anthropic_compatible": "Compatible con Anthropic", + "ollama_compatible": "Compatible con Ollama", + "openai_compatible": "Compatible con OpenAI" + }, + "errors": { + "deleteFailed": "No se pudo eliminar el proveedor personalizado.", + "loadFailed": "No se pudo cargar el proveedor personalizado.", + "saveFailed": "No se pudo guardar el proveedor personalizado.", + "templatesFailed": "No se pudieron cargar las plantillas de proveedores." + }, + "fields": { + "apiKey": "Clave API", + "apiKeyEditPlaceholder": "Déjalo en blanco para conservar la clave guardada", + "apiKeyPlaceholder": "Pega tu clave API", + "apiUrl": "URL de API", + "apiUrlPlaceholder": "https://api.example.com/v1", + "basePath": "Ruta base", + "basePathPlaceholder": "/v1/chat/completions", + "displayName": "Nombre visible", + "displayNamePlaceholder": "Mi proveedor", + "engine": "Motor", + "headerKey": "Encabezado", + "headerValue": "Valor", + "headers": "Encabezados personalizados", + "models": "Modelos", + "modelsPlaceholder": "gpt-4.1, claude-3-5-sonnet", + "requiresAuth": "Requiere autenticación", + "requiresAuthDescription": "Guarda una clave API para las solicitudes a este proveedor.", + "supportsStreaming": "Admite streaming", + "supportsStreamingDescription": "Transmite respuestas mientras el modelo genera tokens." + }, + "modelCount_one": "{{count}} modelo", + "modelCount_other": "{{count}} modelos", + "notConfigured": "No configurado", + "validation": { + "apiKeyRequired": "La clave API es obligatoria.", + "apiUrlInvalid": "Introduce una URL HTTP o HTTPS válida.", + "apiUrlRequired": "La URL de API es obligatoria.", + "displayNameRequired": "El nombre visible es obligatorio.", + "engineRequired": "Elige un motor de proveedor.", + "headerDuplicate": "Los nombres de encabezado deben ser únicos.", + "headerNameInvalid": "Los nombres de encabezado solo pueden contener caracteres HTTP token válidos.", + "headerNameRequired": "El nombre del encabezado es obligatorio.", + "headerValueRequired": "El valor del encabezado es obligatorio.", + "modelsRequired": "Agrega al menos un modelo." + }, + "sections": { + "template": "Comenzar con plantilla" + }, + "templates": { + "compatibility": { + "all": "Todos los tipos de compatibilidad" + }, + "clear": "Borrar plantilla", + "empty": "Todavía no hay plantillas disponibles.", + "manual": "Totalmente personalizado", + "manualDescription": "Comienza con una configuración de proveedor en blanco.", + "searchPlaceholder": "Buscar plantillas", + "useTemplate": "Usar una plantilla", + "useTemplateDescription": "Comienza con valores predeterminados de endpoint y modelos para un proveedor conocido." + } + }, "disconnect": "Desconectar", "models": { "description": "Necesitas al menos un proveedor de modelos para usar el agente Goose. Algunos agentes pueden traer sus propias conexiones de modelo.", diff --git a/ui/goose2/src/shared/ui/alert-dialog.tsx b/ui/goose2/src/shared/ui/alert-dialog.tsx index b9dae5bc..e2bce14b 100644 --- a/ui/goose2/src/shared/ui/alert-dialog.tsx +++ b/ui/goose2/src/shared/ui/alert-dialog.tsx @@ -34,7 +34,7 @@ function AlertDialogOverlay({ - +
+ +
); } diff --git a/ui/goose2/src/shared/ui/dialog.tsx b/ui/goose2/src/shared/ui/dialog.tsx index db9613b4..b7c105d4 100644 --- a/ui/goose2/src/shared/ui/dialog.tsx +++ b/ui/goose2/src/shared/ui/dialog.tsx @@ -36,7 +36,7 @@ function DialogOverlay({ - - {children} - {showCloseButton && ( - - - Close - - )} - +
+ + {children} + {showCloseButton && ( + + + Close + + )} + +
); } diff --git a/ui/goose2/src/shared/ui/icons/ProviderIcons.tsx b/ui/goose2/src/shared/ui/icons/ProviderIcons.tsx index 5e2b88df..a9830c78 100644 --- a/ui/goose2/src/shared/ui/icons/ProviderIcons.tsx +++ b/ui/goose2/src/shared/ui/icons/ProviderIcons.tsx @@ -440,13 +440,16 @@ const PROVIDER_ICON_MAP: Record ReactNode> = { amp: (className) => , "amp-acp": (className) => , azure: (className) => , + azure_openai: (className) => , bedrock: (className) => , + aws_bedrock: (className) => , databricks: (className) => , gcp_vertex_ai: (className) => , ollama: (className) => , openrouter: (className) => , snowflake: (className) => , xai: (className) => , + lmstudio: (className) => , }; function normalizeProviderId(providerId: string) { @@ -482,10 +485,11 @@ export function getProviderIcon( return NORMALIZED_PROVIDER_ICON_MAP[normalizedId](className); } - for (const [key, render] of Object.entries(NORMALIZED_PROVIDER_ICON_MAP)) { - if (normalizedId.includes(key)) { - return render(className); - } + const fallback = Object.entries(NORMALIZED_PROVIDER_ICON_MAP).find( + ([providerFamily]) => normalizedId.includes(providerFamily), + )?.[1]; + if (fallback) { + return fallback(className); } return null; diff --git a/ui/goose2/src/shared/ui/select.tsx b/ui/goose2/src/shared/ui/select.tsx index 7af5ea6e..5dd1d46c 100644 --- a/ui/goose2/src/shared/ui/select.tsx +++ b/ui/goose2/src/shared/ui/select.tsx @@ -59,7 +59,7 @@ function SelectContent({ { + const raw = await this.conn.extMethod( + "_goose/providers/catalog/list", + params, + ); + return zProviderCatalogListResponse.parse( + raw, + ) as ProviderCatalogListResponse; + } + + async GooseProvidersCatalogTemplate( + params: ProviderCatalogTemplateRequest, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/providers/catalog/template", + params, + ); + return zProviderCatalogTemplateResponse.parse( + raw, + ) as ProviderCatalogTemplateResponse; + } + + async GooseProvidersCustomCreate( + params: CustomProviderCreateRequest, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/providers/custom/create", + params, + ); + return zCustomProviderCreateResponse.parse( + raw, + ) as CustomProviderCreateResponse; + } + + async GooseProvidersCustomRead( + params: CustomProviderReadRequest, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/providers/custom/read", + params, + ); + return zCustomProviderReadResponse.parse(raw) as CustomProviderReadResponse; + } + + async GooseProvidersCustomUpdate( + params: CustomProviderUpdateRequest, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/providers/custom/update", + params, + ); + return zCustomProviderUpdateResponse.parse( + raw, + ) as CustomProviderUpdateResponse; + } + + async GooseProvidersCustomDelete( + params: CustomProviderDeleteRequest, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/providers/custom/delete", + params, + ); + return zCustomProviderDeleteResponse.parse( + raw, + ) as CustomProviderDeleteResponse; + } + async GooseProvidersInventoryRefresh( params: RefreshProviderInventoryRequest, ): Promise { diff --git a/ui/sdk/src/generated/index.ts b/ui/sdk/src/generated/index.ts index 90469747..407e381d 100644 --- a/ui/sdk/src/generated/index.ts +++ b/ui/sdk/src/generated/index.ts @@ -1,6 +1,6 @@ // This file is auto-generated by @hey-api/openapi-ts -export type { AddConfigExtensionRequest, AddExtensionRequest, ArchiveSessionRequest, CheckSecretRequest, CheckSecretResponse, CreateSourceRequest, CreateSourceResponse, DeleteSessionRequest, DeleteSourceRequest, DictationConfigRequest, DictationConfigResponse, DictationDownloadProgress, DictationLocalModelStatus, DictationModelCancelRequest, DictationModelDeleteRequest, DictationModelDownloadProgressRequest, DictationModelDownloadProgressResponse, DictationModelDownloadRequest, DictationModelOption, DictationModelSelectRequest, DictationModelsListRequest, DictationModelsListResponse, DictationProviderStatusEntry, DictationTranscribeRequest, DictationTranscribeResponse, EmptyResponse, ExportSessionRequest, ExportSessionResponse, ExportSourceRequest, ExportSourceResponse, ExtRequest, ExtResponse, GetExtensionsRequest, GetExtensionsResponse, GetSessionExtensionsRequest, GetSessionExtensionsResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, ImportSessionResponse, ImportSourcesRequest, ImportSourcesResponse, ListProvidersRequest, ListProvidersResponse, ListSourcesRequest, ListSourcesResponse, ProviderConfigChangeResponse, ProviderConfigDeleteRequest, ProviderConfigFieldUpdate, ProviderConfigFieldValueDto, ProviderConfigKey, ProviderConfigReadRequest, ProviderConfigReadResponse, ProviderConfigSaveRequest, ProviderConfigStatusDto, ProviderConfigStatusRequest, ProviderConfigStatusResponse, ProviderInventoryEntryDto, ProviderInventoryModelDto, ReadConfigRequest, ReadConfigResponse, ReadResourceRequest, ReadResourceResponse, RefreshProviderInventoryRequest, RefreshProviderInventoryResponse, RefreshProviderInventorySkipDto, RefreshProviderInventorySkipReasonDto, RemoveConfigExtensionRequest, RemoveConfigRequest, RemoveExtensionRequest, RemoveSecretRequest, RenameSessionRequest, SourceEntry, SourceType, ToggleConfigExtensionRequest, UnarchiveSessionRequest, UpdateSessionProjectRequest, UpdateSourceRequest, UpdateSourceResponse, UpdateWorkingDirRequest, UpsertConfigRequest, UpsertSecretRequest } from './types.gen.js'; +export type { AddConfigExtensionRequest, AddExtensionRequest, ArchiveSessionRequest, CheckSecretRequest, CheckSecretResponse, CreateSourceRequest, CreateSourceResponse, CustomProviderConfigDto, CustomProviderCreateRequest, CustomProviderCreateResponse, CustomProviderDeleteRequest, CustomProviderDeleteResponse, CustomProviderReadRequest, CustomProviderReadResponse, CustomProviderUpdateRequest, CustomProviderUpdateResponse, DeleteSessionRequest, DeleteSourceRequest, DictationConfigRequest, DictationConfigResponse, DictationDownloadProgress, DictationLocalModelStatus, DictationModelCancelRequest, DictationModelDeleteRequest, DictationModelDownloadProgressRequest, DictationModelDownloadProgressResponse, DictationModelDownloadRequest, DictationModelOption, DictationModelSelectRequest, DictationModelsListRequest, DictationModelsListResponse, DictationProviderStatusEntry, DictationTranscribeRequest, DictationTranscribeResponse, EmptyResponse, ExportSessionRequest, ExportSessionResponse, ExportSourceRequest, ExportSourceResponse, ExtRequest, ExtResponse, GetExtensionsRequest, GetExtensionsResponse, GetSessionExtensionsRequest, GetSessionExtensionsResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, ImportSessionResponse, ImportSourcesRequest, ImportSourcesResponse, ListProvidersRequest, ListProvidersResponse, ListSourcesRequest, ListSourcesResponse, ProviderCatalogEntryDto, ProviderCatalogListRequest, ProviderCatalogListResponse, ProviderCatalogTemplateRequest, ProviderCatalogTemplateResponse, ProviderConfigChangeResponse, ProviderConfigDeleteRequest, ProviderConfigFieldUpdate, ProviderConfigFieldValueDto, ProviderConfigKey, ProviderConfigReadRequest, ProviderConfigReadResponse, ProviderConfigSaveRequest, ProviderConfigStatusDto, ProviderConfigStatusRequest, ProviderConfigStatusResponse, ProviderInventoryEntryDto, ProviderInventoryModelDto, ProviderTemplateCapabilitiesDto, ProviderTemplateDto, ProviderTemplateModelDto, ReadConfigRequest, ReadConfigResponse, ReadResourceRequest, ReadResourceResponse, RefreshProviderInventoryRequest, RefreshProviderInventoryResponse, RefreshProviderInventorySkipDto, RefreshProviderInventorySkipReasonDto, RemoveConfigExtensionRequest, RemoveConfigRequest, RemoveExtensionRequest, RemoveSecretRequest, RenameSessionRequest, SourceEntry, SourceType, ToggleConfigExtensionRequest, UnarchiveSessionRequest, UpdateSessionProjectRequest, UpdateSourceRequest, UpdateSourceResponse, UpdateWorkingDirRequest, UpsertConfigRequest, UpsertSecretRequest } from './types.gen.js'; export const GOOSE_EXT_METHODS = [ { @@ -63,6 +63,36 @@ export const GOOSE_EXT_METHODS = [ requestType: "ListProvidersRequest", responseType: "ListProvidersResponse", }, + { + method: "_goose/providers/catalog/list", + requestType: "ProviderCatalogListRequest", + responseType: "ProviderCatalogListResponse", + }, + { + method: "_goose/providers/catalog/template", + requestType: "ProviderCatalogTemplateRequest", + responseType: "ProviderCatalogTemplateResponse", + }, + { + method: "_goose/providers/custom/create", + requestType: "CustomProviderCreateRequest", + responseType: "CustomProviderCreateResponse", + }, + { + method: "_goose/providers/custom/read", + requestType: "CustomProviderReadRequest", + responseType: "CustomProviderReadResponse", + }, + { + method: "_goose/providers/custom/update", + requestType: "CustomProviderUpdateRequest", + responseType: "CustomProviderUpdateResponse", + }, + { + method: "_goose/providers/custom/delete", + requestType: "CustomProviderDeleteRequest", + responseType: "CustomProviderDeleteResponse", + }, { method: "_goose/providers/inventory/refresh", requestType: "RefreshProviderInventoryRequest", diff --git a/ui/sdk/src/generated/types.gen.ts b/ui/sdk/src/generated/types.gen.ts index 6e8a0557..1d9aa5fe 100644 --- a/ui/sdk/src/generated/types.gen.ts +++ b/ui/sdk/src/generated/types.gen.ts @@ -261,13 +261,90 @@ export type ProviderInventoryModelDto = { }; /** - * Trigger a background refresh of provider inventories. + * List custom-provider catalog entries. Omit `format` to list all formats. */ -export type RefreshProviderInventoryRequest = { - /** - * Which providers to refresh. Empty means all known providers. - */ - providerIds?: Array; +export type ProviderCatalogListRequest = { + format?: string | null; +}; + +export type ProviderCatalogListResponse = { + providers: Array; +}; + +export type ProviderCatalogEntryDto = { + providerId: string; + name: string; + format: string; + apiUrl: string; + modelCount: number; + docUrl: string; + envVar: string; +}; + +/** + * Return the editable template for one catalog provider. + */ +export type ProviderCatalogTemplateRequest = { + providerId: string; +}; + +export type ProviderCatalogTemplateResponse = { + template: ProviderTemplateDto; +}; + +export type ProviderTemplateDto = { + providerId: string; + name: string; + format: string; + apiUrl: string; + models: Array; + supportsStreaming: boolean; + envVar: string; + docUrl: string; +}; + +export type ProviderTemplateModelDto = { + id: string; + name: string; + contextLimit: number; + capabilities: ProviderTemplateCapabilitiesDto; + deprecated: boolean; +}; + +export type ProviderTemplateCapabilitiesDto = { + toolCall: boolean; + reasoning: boolean; + attachment: boolean; + temperature: boolean; +}; + +/** + * Create a custom provider backed by Goose's declarative provider store. + */ +export type CustomProviderCreateRequest = { + engine: string; + displayName: string; + apiUrl: string; + apiKey?: string | null; + models?: Array; + supportsStreaming?: boolean | null; + headers?: { + [key: string]: string; + }; + requiresAuth: boolean; + catalogProviderId?: string | null; + basePath?: string | null; +}; + +export type CustomProviderCreateResponse = { + providerId: string; + status: ProviderConfigStatusDto; + refresh: RefreshProviderInventoryResponse; +}; + +export type ProviderConfigStatusDto = { + providerId: string; + isConfigured: boolean; }; /** @@ -291,6 +368,83 @@ export type RefreshProviderInventorySkipDto = { export type RefreshProviderInventorySkipReasonDto = 'unknown_provider' | 'not_configured' | 'does_not_support_refresh' | 'already_refreshing'; +/** + * Read a declarative provider config. Custom configs are editable; bundled configs are read-only. + */ +export type CustomProviderReadRequest = { + providerId: string; +}; + +export type CustomProviderReadResponse = { + provider: CustomProviderConfigDto; + editable: boolean; + status: ProviderConfigStatusDto; +}; + +export type CustomProviderConfigDto = { + providerId: string; + engine: string; + displayName: string; + apiUrl: string; + models?: Array; + supportsStreaming?: boolean | null; + headers?: { + [key: string]: string; + }; + requiresAuth: boolean; + catalogProviderId?: string | null; + basePath?: string | null; + apiKeyEnv?: string | null; + apiKeySet: boolean; +}; + +/** + * Update a custom provider backed by Goose's declarative provider store. + */ +export type CustomProviderUpdateRequest = { + providerId: string; + engine: string; + displayName: string; + apiUrl: string; + apiKey?: string | null; + models?: Array; + supportsStreaming?: boolean | null; + headers?: { + [key: string]: string; + }; + requiresAuth: boolean; + catalogProviderId?: string | null; + basePath?: string | null; +}; + +export type CustomProviderUpdateResponse = { + providerId: string; + status: ProviderConfigStatusDto; + refresh: RefreshProviderInventoryResponse; +}; + +/** + * Delete a custom provider from Goose's declarative provider store. + */ +export type CustomProviderDeleteRequest = { + providerId: string; +}; + +export type CustomProviderDeleteResponse = { + providerId: string; + refresh: RefreshProviderInventoryResponse; +}; + +/** + * Trigger a background refresh of provider inventories. + */ +export type RefreshProviderInventoryRequest = { + /** + * Which providers to refresh. Empty means all known providers. + */ + providerIds?: Array; +}; + /** * Read saved configuration field values for one provider. */ @@ -321,11 +475,6 @@ export type ProviderConfigStatusResponse = { statuses: Array; }; -export type ProviderConfigStatusDto = { - providerId: string; - isConfigured: boolean; -}; - /** * Save provider configuration fields and start an inventory refresh when supported. */ @@ -727,14 +876,14 @@ export type DictationModelSelectRequest = { export type ExtRequest = { id: string; method: string; - params?: AddExtensionRequest | RemoveExtensionRequest | GetToolsRequest | ReadResourceRequest | UpdateWorkingDirRequest | DeleteSessionRequest | GetExtensionsRequest | AddConfigExtensionRequest | RemoveConfigExtensionRequest | ToggleConfigExtensionRequest | GetSessionExtensionsRequest | ListProvidersRequest | RefreshProviderInventoryRequest | ProviderConfigReadRequest | ProviderConfigStatusRequest | ProviderConfigSaveRequest | ProviderConfigDeleteRequest | ReadConfigRequest | UpsertConfigRequest | RemoveConfigRequest | CheckSecretRequest | UpsertSecretRequest | RemoveSecretRequest | ExportSessionRequest | ImportSessionRequest | UpdateSessionProjectRequest | RenameSessionRequest | ArchiveSessionRequest | UnarchiveSessionRequest | CreateSourceRequest | ListSourcesRequest | UpdateSourceRequest | DeleteSourceRequest | ExportSourceRequest | ImportSourcesRequest | DictationTranscribeRequest | DictationConfigRequest | DictationModelsListRequest | DictationModelDownloadRequest | DictationModelDownloadProgressRequest | DictationModelCancelRequest | DictationModelDeleteRequest | DictationModelSelectRequest | { + params?: AddExtensionRequest | RemoveExtensionRequest | GetToolsRequest | ReadResourceRequest | UpdateWorkingDirRequest | DeleteSessionRequest | GetExtensionsRequest | AddConfigExtensionRequest | RemoveConfigExtensionRequest | ToggleConfigExtensionRequest | GetSessionExtensionsRequest | ListProvidersRequest | ProviderCatalogListRequest | ProviderCatalogTemplateRequest | CustomProviderCreateRequest | CustomProviderReadRequest | CustomProviderUpdateRequest | CustomProviderDeleteRequest | RefreshProviderInventoryRequest | ProviderConfigReadRequest | ProviderConfigStatusRequest | ProviderConfigSaveRequest | ProviderConfigDeleteRequest | ReadConfigRequest | UpsertConfigRequest | RemoveConfigRequest | CheckSecretRequest | UpsertSecretRequest | RemoveSecretRequest | ExportSessionRequest | ImportSessionRequest | UpdateSessionProjectRequest | RenameSessionRequest | ArchiveSessionRequest | UnarchiveSessionRequest | CreateSourceRequest | ListSourcesRequest | UpdateSourceRequest | DeleteSourceRequest | ExportSourceRequest | ImportSourcesRequest | DictationTranscribeRequest | DictationConfigRequest | DictationModelsListRequest | DictationModelDownloadRequest | DictationModelDownloadProgressRequest | DictationModelCancelRequest | DictationModelDeleteRequest | DictationModelSelectRequest | { [key: string]: unknown; } | null; }; export type ExtResponse = { id: string; - result?: EmptyResponse | GetToolsResponse | ReadResourceResponse | GetExtensionsResponse | GetSessionExtensionsResponse | ListProvidersResponse | RefreshProviderInventoryResponse | ProviderConfigReadResponse | ProviderConfigStatusResponse | ProviderConfigChangeResponse | ReadConfigResponse | CheckSecretResponse | ExportSessionResponse | ImportSessionResponse | CreateSourceResponse | ListSourcesResponse | UpdateSourceResponse | ExportSourceResponse | ImportSourcesResponse | DictationTranscribeResponse | DictationConfigResponse | DictationModelsListResponse | DictationModelDownloadProgressResponse | unknown; + result?: EmptyResponse | GetToolsResponse | ReadResourceResponse | GetExtensionsResponse | GetSessionExtensionsResponse | ListProvidersResponse | ProviderCatalogListResponse | ProviderCatalogTemplateResponse | CustomProviderCreateResponse | CustomProviderReadResponse | CustomProviderUpdateResponse | CustomProviderDeleteResponse | RefreshProviderInventoryResponse | ProviderConfigReadResponse | ProviderConfigStatusResponse | ProviderConfigChangeResponse | ReadConfigResponse | CheckSecretResponse | ExportSessionResponse | ImportSessionResponse | CreateSourceResponse | ListSourcesResponse | UpdateSourceResponse | ExportSourceResponse | ImportSourcesResponse | DictationTranscribeResponse | DictationConfigResponse | DictationModelsListResponse | DictationModelDownloadProgressResponse | unknown; } | { error: { code: number; diff --git a/ui/sdk/src/generated/zod.gen.ts b/ui/sdk/src/generated/zod.gen.ts index f0f6e594..a562cfcf 100644 --- a/ui/sdk/src/generated/zod.gen.ts +++ b/ui/sdk/src/generated/zod.gen.ts @@ -196,10 +196,97 @@ export const zListProvidersResponse = z.object({ }); /** - * Trigger a background refresh of provider inventories. + * List custom-provider catalog entries. Omit `format` to list all formats. */ -export const zRefreshProviderInventoryRequest = z.object({ - providerIds: z.array(z.string()).optional().default([]) +export const zProviderCatalogListRequest = z.object({ + format: z.union([ + z.string(), + z.null() + ]).optional() +}); + +export const zProviderCatalogEntryDto = z.object({ + providerId: z.string(), + name: z.string(), + format: z.string(), + apiUrl: z.string(), + modelCount: z.number().int().gte(0), + docUrl: z.string(), + envVar: z.string() +}); + +export const zProviderCatalogListResponse = z.object({ + providers: z.array(zProviderCatalogEntryDto) +}); + +/** + * Return the editable template for one catalog provider. + */ +export const zProviderCatalogTemplateRequest = z.object({ + providerId: z.string() +}); + +export const zProviderTemplateCapabilitiesDto = z.object({ + toolCall: z.boolean(), + reasoning: z.boolean(), + attachment: z.boolean(), + temperature: z.boolean() +}); + +export const zProviderTemplateModelDto = z.object({ + id: z.string(), + name: z.string(), + contextLimit: z.number().int().gte(0), + capabilities: zProviderTemplateCapabilitiesDto, + deprecated: z.boolean() +}); + +export const zProviderTemplateDto = z.object({ + providerId: z.string(), + name: z.string(), + format: z.string(), + apiUrl: z.string(), + models: z.array(zProviderTemplateModelDto), + supportsStreaming: z.boolean(), + envVar: z.string(), + docUrl: z.string() +}); + +export const zProviderCatalogTemplateResponse = z.object({ + template: zProviderTemplateDto +}); + +/** + * Create a custom provider backed by Goose's declarative provider store. + */ +export const zCustomProviderCreateRequest = z.object({ + engine: z.string(), + displayName: z.string(), + apiUrl: z.string(), + apiKey: z.union([ + z.string(), + z.null() + ]).optional(), + models: z.array(z.string()).optional().default([]), + supportsStreaming: z.union([ + z.boolean(), + z.null() + ]).optional(), + headers: z.record(z.string()).optional().default({}), + requiresAuth: z.boolean(), + catalogProviderId: z.union([ + z.string(), + z.null() + ]).optional(), + basePath: z.union([ + z.string(), + z.null() + ]).optional() +}); + +export const zProviderConfigStatusDto = z.object({ + providerId: z.string(), + isConfigured: z.boolean() }); export const zRefreshProviderInventorySkipReasonDto = z.enum([ @@ -222,6 +309,106 @@ export const zRefreshProviderInventoryResponse = z.object({ skipped: z.array(zRefreshProviderInventorySkipDto).optional().default([]) }); +export const zCustomProviderCreateResponse = z.object({ + providerId: z.string(), + status: zProviderConfigStatusDto, + refresh: zRefreshProviderInventoryResponse +}); + +/** + * Read a declarative provider config. Custom configs are editable; bundled configs are read-only. + */ +export const zCustomProviderReadRequest = z.object({ + providerId: z.string() +}); + +export const zCustomProviderConfigDto = z.object({ + providerId: z.string(), + engine: z.string(), + displayName: z.string(), + apiUrl: z.string(), + models: z.array(z.string()).optional().default([]), + supportsStreaming: z.union([ + z.boolean(), + z.null() + ]).optional(), + headers: z.record(z.string()).optional().default({}), + requiresAuth: z.boolean(), + catalogProviderId: z.union([ + z.string(), + z.null() + ]).optional(), + basePath: z.union([ + z.string(), + z.null() + ]).optional(), + apiKeyEnv: z.union([ + z.string(), + z.null() + ]).optional(), + apiKeySet: z.boolean() +}); + +export const zCustomProviderReadResponse = z.object({ + provider: zCustomProviderConfigDto, + editable: z.boolean(), + status: zProviderConfigStatusDto +}); + +/** + * Update a custom provider backed by Goose's declarative provider store. + */ +export const zCustomProviderUpdateRequest = z.object({ + providerId: z.string(), + engine: z.string(), + displayName: z.string(), + apiUrl: z.string(), + apiKey: z.union([ + z.string(), + z.null() + ]).optional(), + models: z.array(z.string()).optional().default([]), + supportsStreaming: z.union([ + z.boolean(), + z.null() + ]).optional(), + headers: z.record(z.string()).optional().default({}), + requiresAuth: z.boolean(), + catalogProviderId: z.union([ + z.string(), + z.null() + ]).optional(), + basePath: z.union([ + z.string(), + z.null() + ]).optional() +}); + +export const zCustomProviderUpdateResponse = z.object({ + providerId: z.string(), + status: zProviderConfigStatusDto, + refresh: zRefreshProviderInventoryResponse +}); + +/** + * Delete a custom provider from Goose's declarative provider store. + */ +export const zCustomProviderDeleteRequest = z.object({ + providerId: z.string() +}); + +export const zCustomProviderDeleteResponse = z.object({ + providerId: z.string(), + refresh: zRefreshProviderInventoryResponse +}); + +/** + * Trigger a background refresh of provider inventories. + */ +export const zRefreshProviderInventoryRequest = z.object({ + providerIds: z.array(z.string()).optional().default([]) +}); + /** * Read saved configuration field values for one provider. */ @@ -251,11 +438,6 @@ export const zProviderConfigStatusRequest = z.object({ providerIds: z.array(z.string()).optional().default([]) }); -export const zProviderConfigStatusDto = z.object({ - providerId: z.string(), - isConfigured: z.boolean() -}); - export const zProviderConfigStatusResponse = z.object({ statuses: z.array(zProviderConfigStatusDto) }); @@ -690,6 +872,12 @@ export const zExtRequest = z.object({ zToggleConfigExtensionRequest, zGetSessionExtensionsRequest, zListProvidersRequest, + zProviderCatalogListRequest, + zProviderCatalogTemplateRequest, + zCustomProviderCreateRequest, + zCustomProviderReadRequest, + zCustomProviderUpdateRequest, + zCustomProviderDeleteRequest, zRefreshProviderInventoryRequest, zProviderConfigReadRequest, zProviderConfigStatusRequest, @@ -740,6 +928,12 @@ export const zExtResponse = z.union([ zGetExtensionsResponse, zGetSessionExtensionsResponse, zListProvidersResponse, + zProviderCatalogListResponse, + zProviderCatalogTemplateResponse, + zCustomProviderCreateResponse, + zCustomProviderReadResponse, + zCustomProviderUpdateResponse, + zCustomProviderDeleteResponse, zRefreshProviderInventoryResponse, zProviderConfigReadResponse, zProviderConfigStatusResponse,