From 826f13f343606c287165ffb373304b262736b0e8 Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Fri, 5 Jun 2026 15:28:50 -0400 Subject: [PATCH] create goose-providers crate with canonical models, conversation and other types (#9588) --- Cargo.lock | 21 ++ Justfile | 4 +- crates/goose-providers/Cargo.toml | 30 +++ .../src}/canonical/README.md | 6 +- .../src/canonical}/catalog.rs | 198 +++-------------- .../data/canonical_mapping_report.json | 0 .../src}/canonical/data/canonical_models.json | 0 .../canonical/data/provider_metadata.json | 0 .../src}/canonical/mod.rs | 1 + .../src}/canonical/model.rs | 0 .../src}/canonical/name_builder.rs | 0 .../src}/canonical/registry.rs | 0 .../src/conversation/message.rs | 3 +- .../src/conversation/mod.rs | 0 .../src/conversation/tool_result_serde.rs | 2 +- crates/goose-providers/src/lib.rs | 4 + crates/goose-providers/src/mcp_utils.rs | 105 +++++++++ crates/goose-providers/src/utils.rs | 14 ++ crates/goose/Cargo.toml | 3 +- .../build_canonical_models.rs | 9 +- crates/goose/src/lib.rs | 4 +- crates/goose/src/providers/catalog_util.rs | 205 ++++++++++++++++++ crates/goose/src/providers/mod.rs | 9 +- 23 files changed, 439 insertions(+), 179 deletions(-) create mode 100644 crates/goose-providers/Cargo.toml rename crates/{goose/src/providers => goose-providers/src}/canonical/README.md (74%) rename crates/{goose/src/providers => goose-providers/src/canonical}/catalog.rs (86%) rename crates/{goose/src/providers => goose-providers/src}/canonical/data/canonical_mapping_report.json (100%) rename crates/{goose/src/providers => goose-providers/src}/canonical/data/canonical_models.json (100%) rename crates/{goose/src/providers => goose-providers/src}/canonical/data/provider_metadata.json (100%) rename crates/{goose/src/providers => goose-providers/src}/canonical/mod.rs (99%) rename crates/{goose/src/providers => goose-providers/src}/canonical/model.rs (100%) rename crates/{goose/src/providers => goose-providers/src}/canonical/name_builder.rs (100%) rename crates/{goose/src/providers => goose-providers/src}/canonical/registry.rs (100%) rename crates/{goose => goose-providers}/src/conversation/message.rs (99%) rename crates/{goose => goose-providers}/src/conversation/mod.rs (100%) rename crates/{goose => goose-providers}/src/conversation/tool_result_serde.rs (99%) create mode 100644 crates/goose-providers/src/lib.rs create mode 100644 crates/goose-providers/src/mcp_utils.rs create mode 100644 crates/goose-providers/src/utils.rs rename crates/goose/src/{providers/canonical => bin}/build_canonical_models.rs (99%) create mode 100644 crates/goose/src/providers/catalog_util.rs diff --git a/Cargo.lock b/Cargo.lock index fd7f5d23..db116e32 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4574,6 +4574,7 @@ dependencies = [ "futures", "goose-acp-macros", "goose-mcp", + "goose-providers", "goose-sdk-types", "goose-test-support", "http 1.4.1", @@ -4762,6 +4763,26 @@ dependencies = [ "url", ] +[[package]] +name = "goose-providers" +version = "1.37.0" +dependencies = [ + "anyhow", + "base64 0.22.1", + "chrono", + "once_cell", + "regex", + "rmcp", + "serde", + "serde_json", + "test-case", + "thiserror 1.0.69", + "tracing", + "unicode-normalization", + "utoipa 4.2.3", + "uuid", +] + [[package]] name = "goose-sdk" version = "1.37.0" diff --git a/Justfile b/Justfile index 8e8f6e49..017f0c50 100644 --- a/Justfile +++ b/Justfile @@ -330,8 +330,8 @@ prepare-release version: ui/desktop/package.json \ ui/pnpm-lock.yaml \ ui/desktop/openapi.json \ - crates/goose/src/providers/canonical/data/canonical_models.json \ - crates/goose/src/providers/canonical/data/provider_metadata.json + crates/goose-providers/src/canonical/data/canonical_models.json \ + crates/goose-providers/src/canonical/data/provider_metadata.json @git commit --message "chore(release): release version {{ version }}" set-openapi-version version: diff --git a/crates/goose-providers/Cargo.toml b/crates/goose-providers/Cargo.toml new file mode 100644 index 00000000..e51dea54 --- /dev/null +++ b/crates/goose-providers/Cargo.toml @@ -0,0 +1,30 @@ +[package] +name = "goose-providers" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +authors.workspace = true +license.workspace = true +repository.workspace = true +description.workspace = true + +[lints] +workspace = true + +[dependencies] +anyhow = { workspace = true } +base64 = { workspace = true } +chrono = { workspace = true } +once_cell = { workspace = true } +regex = { workspace = true } +rmcp = { workspace = true, features = ["server"] } +serde = { workspace = true } +serde_json = { workspace = true } +thiserror = { workspace = true } +tracing = { workspace = true } +unicode-normalization = { version = "0.1.22", default-features = false, features = ["std"] } +utoipa = { workspace = true, features = ["chrono"] } +uuid = { workspace = true, features = ["v4", "std"] } + +[dev-dependencies] +test-case = { workspace = true } diff --git a/crates/goose/src/providers/canonical/README.md b/crates/goose-providers/src/canonical/README.md similarity index 74% rename from crates/goose/src/providers/canonical/README.md rename to crates/goose-providers/src/canonical/README.md index 410d102c..f0c682f2 100644 --- a/crates/goose/src/providers/canonical/README.md +++ b/crates/goose-providers/src/canonical/README.md @@ -13,11 +13,11 @@ cargo run --bin build_canonical_models --no-check # Build only, skip checker This script performs two operations by default: 1. **Builds canonical models** - Fetches from OpenRouter API and updates the registry - - Writes to: `src/providers/canonical/data/canonical_models.json` + - Writes to: `crates/goose-providers/src/canonical/data/canonical_models.json` 2. **Checks model mappings** (unless `--no-check` is passed) - Tests provider mappings and tracks changes over time - Reports unmapped models - Compares with previous runs (like a lock file) - Shows changed/added/removed mappings - - Writes to: `src/providers/canonical/data/canonical_mapping_report.json` + - Writes to: `crates/goose-providers/src/canonical/data/canonical_mapping_report.json` -The script is located in this directory: `build_canonical_models.rs` +The script is currently built from `crates/goose/src/bin/build_canonical_models.rs` and writes into this crate's `src/canonical/data` directory. diff --git a/crates/goose/src/providers/catalog.rs b/crates/goose-providers/src/canonical/catalog.rs similarity index 86% rename from crates/goose/src/providers/catalog.rs rename to crates/goose-providers/src/canonical/catalog.rs index 05b032cc..d7e3617d 100644 --- a/crates/goose/src/providers/catalog.rs +++ b/crates/goose-providers/src/canonical/catalog.rs @@ -1,13 +1,10 @@ use once_cell::sync::Lazy; use serde::{Deserialize, Serialize}; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; -use super::{ - base::{ConfigKey, ProviderMetadata}, - canonical::CanonicalModelRegistry, -}; +use super::CanonicalModelRegistry; -const PROVIDER_METADATA_JSON: &str = include_str!("canonical/data/provider_metadata.json"); +const PROVIDER_METADATA_JSON: &str = include_str!("data/provider_metadata.json"); #[derive(Debug, Clone, Serialize, Deserialize)] struct ProviderMetadataEntry { @@ -175,6 +172,24 @@ pub struct ProviderSetupCatalogEntry { pub show_only_when_installed: bool, } +#[derive(Debug, Clone)] +pub struct ProviderSetupMetadata { + pub name: String, + pub display_name: String, + pub description: String, + pub model_doc_link: String, + pub config_keys: Vec, +} + +#[derive(Debug, Clone)] +pub struct ProviderSetupConfigKey { + pub name: String, + pub required: bool, + pub secret: bool, + pub default: Option, + pub primary: bool, +} + #[derive(Debug, Clone, Copy)] struct CuratedSetupMetadata { provider_id: &'static str, @@ -900,7 +915,7 @@ fn field_label(key: &str) -> String { fn field_override<'a>( key: &str, - config_key: &ConfigKey, + config_key: &ProviderSetupConfigKey, curated: &'a CuratedSetupMetadata, ) -> Option<&'a CuratedFieldMetadata> { if let Some(field) = curated @@ -918,7 +933,10 @@ fn field_override<'a>( None } -fn setup_field(config_key: &ConfigKey, curated: &CuratedSetupMetadata) -> ProviderSetupField { +fn setup_field( + config_key: &ProviderSetupConfigKey, + curated: &CuratedSetupMetadata, +) -> ProviderSetupField { let field_override = field_override(&config_key.name, config_key, curated); ProviderSetupField { key: config_key.name.clone(), @@ -936,7 +954,7 @@ fn setup_field(config_key: &ConfigKey, curated: &CuratedSetupMetadata) -> Provid fn setup_entry_from_metadata( curated: &CuratedSetupMetadata, - metadata: &ProviderMetadata, + metadata: &ProviderSetupMetadata, ) -> ProviderSetupCatalogEntry { ProviderSetupCatalogEntry { provider_id: curated.provider_id.to_string(), @@ -994,13 +1012,10 @@ fn synthetic_goose_setup_entry(curated: &CuratedSetupMetadata) -> ProviderSetupC } } -pub async fn get_providers_by_format(format: ProviderFormat) -> Vec { - let native_provider_ids = super::init::providers() - .await - .into_iter() - .map(|(metadata, _)| metadata.name) - .collect::>(); - +pub fn get_providers_by_format( + format: ProviderFormat, + native_provider_ids: &HashSet, +) -> Vec { let mut entries: Vec = PROVIDER_METADATA .values() .filter_map(|metadata| { @@ -1038,13 +1053,9 @@ pub async fn get_providers_by_format(format: ProviderFormat) -> Vec Vec { - let registry_metadata = super::providers() - .await - .into_iter() - .map(|(metadata, _)| (metadata.name.clone(), metadata)) - .collect::>(); - +pub fn get_setup_catalog_entries( + registry_metadata: &HashMap, +) -> Vec { SETUP_METADATA .iter() .filter_map(|curated| { @@ -1122,144 +1133,3 @@ pub fn get_provider_template(provider_id: &str) -> Option { doc_url: metadata.doc.clone().unwrap_or_default(), }) } - -#[cfg(test)] -mod tests { - use super::*; - use crate::providers::base::ProviderType; - - #[tokio::test] - async fn test_zai_provider() { - let zai = crate::providers::get_from_registry("zai") - .await - .expect("z.ai should be registered as a declarative provider"); - assert_eq!(zai.provider_type(), ProviderType::Declarative); - - let metadata = zai.metadata(); - assert_eq!(metadata.display_name, "Z.AI"); - assert!( - !metadata.known_models.is_empty(), - "z.ai should have known models" - ); - assert!( - metadata - .config_keys - .iter() - .any(|key| key.name == "ZHIPU_API_KEY"), - "z.ai should expose its API key config" - ); - - let setup_entries = get_setup_catalog_entries().await; - let setup_entry = setup_entries - .iter() - .find(|entry| entry.provider_id == "zai") - .expect("z.ai should be in the setup catalog"); - assert_eq!(setup_entry.setup_method, ProviderSetupMethod::SingleApiKey); - - let template = get_provider_template("zai"); - assert!(template.is_some(), "z.ai should have a template"); - - let template = template.unwrap(); - println!("Z.AI template: {} models", template.models.len()); - for model in template.models.iter().take(3) { - println!( - " - {} ({}K context)", - model.name, - model.context_limit / 1000 - ); - } - assert!( - !template.models.is_empty(), - "z.ai template should have models" - ); - } - - #[tokio::test] - async fn setup_catalog_includes_goose_and_curated_fields() { - let entries = get_setup_catalog_entries().await; - - let goose = entries - .iter() - .find(|entry| entry.provider_id == "goose") - .expect("setup catalog should include synthetic goose"); - assert_eq!(goose.category, ProviderSetupCategory::Agent); - assert_eq!(goose.setup_method, ProviderSetupMethod::None); - assert!(goose.fields.is_empty()); - - let ollama = entries - .iter() - .find(|entry| entry.provider_id == "ollama") - .expect("setup catalog should include ollama"); - assert_eq!(ollama.setup_method, ProviderSetupMethod::ConfigFields); - assert_eq!(ollama.fields.len(), 1); - assert_eq!(ollama.fields[0].key, "OLLAMA_HOST"); - assert_eq!(ollama.fields[0].label, "Host"); - assert_eq!( - ollama.fields[0].default_value.as_deref(), - Some("http://localhost:11434") - ); - - let databricks = entries - .iter() - .find(|entry| entry.provider_id == "databricks") - .expect("setup catalog should include databricks"); - assert_eq!( - databricks.setup_method, - ProviderSetupMethod::HostWithOauthFallback - ); - assert_eq!( - databricks - .fields - .iter() - .map(|field| field.key.as_str()) - .collect::>(), - ["DATABRICKS_HOST", "DATABRICKS_TOKEN"] - ); - - let huggingface = entries - .iter() - .find(|entry| entry.provider_id == "huggingface") - .expect("setup catalog should include huggingface"); - assert_eq!(huggingface.setup_method, ProviderSetupMethod::SingleApiKey); - assert_eq!( - huggingface - .fields - .iter() - .map(|field| field.key.as_str()) - .collect::>(), - ["HF_TOKEN"] - ); - - let atomic_chat = entries - .iter() - .find(|entry| entry.provider_id == "atomic_chat") - .expect("setup catalog should include atomic_chat declarative provider"); - assert_eq!(atomic_chat.setup_method, ProviderSetupMethod::ConfigFields); - let host_field = atomic_chat - .fields - .iter() - .find(|field| field.key == "ATOMIC_CHAT_HOST") - .expect("atomic_chat should expose ATOMIC_CHAT_HOST"); - assert_eq!(host_field.label, "Host URL"); - assert_eq!( - host_field.default_value.as_deref(), - Some("http://localhost:1337") - ); - } - - #[tokio::test] - async fn setup_catalog_excludes_uncurated_deprecated_providers() { - let provider_ids = get_setup_catalog_entries() - .await - .into_iter() - .map(|entry| entry.provider_id) - .collect::>(); - - assert!(provider_ids.contains("claude-acp")); - assert!(provider_ids.contains("codex-acp")); - assert!(provider_ids.contains("atomic_chat")); - assert!(!provider_ids.contains("claude_code")); - assert!(!provider_ids.contains("codex")); - assert!(!provider_ids.contains("gemini_cli")); - } -} diff --git a/crates/goose/src/providers/canonical/data/canonical_mapping_report.json b/crates/goose-providers/src/canonical/data/canonical_mapping_report.json similarity index 100% rename from crates/goose/src/providers/canonical/data/canonical_mapping_report.json rename to crates/goose-providers/src/canonical/data/canonical_mapping_report.json diff --git a/crates/goose/src/providers/canonical/data/canonical_models.json b/crates/goose-providers/src/canonical/data/canonical_models.json similarity index 100% rename from crates/goose/src/providers/canonical/data/canonical_models.json rename to crates/goose-providers/src/canonical/data/canonical_models.json diff --git a/crates/goose/src/providers/canonical/data/provider_metadata.json b/crates/goose-providers/src/canonical/data/provider_metadata.json similarity index 100% rename from crates/goose/src/providers/canonical/data/provider_metadata.json rename to crates/goose-providers/src/canonical/data/provider_metadata.json diff --git a/crates/goose/src/providers/canonical/mod.rs b/crates/goose-providers/src/canonical/mod.rs similarity index 99% rename from crates/goose/src/providers/canonical/mod.rs rename to crates/goose-providers/src/canonical/mod.rs index ab7bc81a..cc0fceec 100644 --- a/crates/goose/src/providers/canonical/mod.rs +++ b/crates/goose-providers/src/canonical/mod.rs @@ -1,3 +1,4 @@ +pub mod catalog; mod model; mod name_builder; mod registry; diff --git a/crates/goose/src/providers/canonical/model.rs b/crates/goose-providers/src/canonical/model.rs similarity index 100% rename from crates/goose/src/providers/canonical/model.rs rename to crates/goose-providers/src/canonical/model.rs diff --git a/crates/goose/src/providers/canonical/name_builder.rs b/crates/goose-providers/src/canonical/name_builder.rs similarity index 100% rename from crates/goose/src/providers/canonical/name_builder.rs rename to crates/goose-providers/src/canonical/name_builder.rs diff --git a/crates/goose/src/providers/canonical/registry.rs b/crates/goose-providers/src/canonical/registry.rs similarity index 100% rename from crates/goose/src/providers/canonical/registry.rs rename to crates/goose-providers/src/canonical/registry.rs diff --git a/crates/goose/src/conversation/message.rs b/crates/goose-providers/src/conversation/message.rs similarity index 99% rename from crates/goose/src/conversation/message.rs rename to crates/goose-providers/src/conversation/message.rs index 83ab9321..a89d4760 100644 --- a/crates/goose/src/conversation/message.rs +++ b/crates/goose-providers/src/conversation/message.rs @@ -1,5 +1,5 @@ use crate::conversation::tool_result_serde; -use crate::mcp_utils::{extract_text_from_resource, ToolResult}; +use crate::mcp_utils::extract_text_from_resource; use crate::utils::sanitize_unicode_tags; use chrono::Utc; use rmcp::model::{ @@ -74,6 +74,7 @@ where /// Provider-specific metadata for tool requests/responses. /// Allows providers to store custom data without polluting the core model. pub type ProviderMetadata = serde_json::Map; +pub type ToolResult = Result; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] diff --git a/crates/goose/src/conversation/mod.rs b/crates/goose-providers/src/conversation/mod.rs similarity index 100% rename from crates/goose/src/conversation/mod.rs rename to crates/goose-providers/src/conversation/mod.rs diff --git a/crates/goose/src/conversation/tool_result_serde.rs b/crates/goose-providers/src/conversation/tool_result_serde.rs similarity index 99% rename from crates/goose/src/conversation/tool_result_serde.rs rename to crates/goose-providers/src/conversation/tool_result_serde.rs index 54578bcc..b0d6c737 100644 --- a/crates/goose/src/conversation/tool_result_serde.rs +++ b/crates/goose-providers/src/conversation/tool_result_serde.rs @@ -1,4 +1,4 @@ -use crate::mcp_utils::ToolResult; +use crate::conversation::message::ToolResult; use rmcp::model::{CallToolRequestParams, ErrorCode, ErrorData, JsonObject}; use serde::ser::SerializeStruct; use serde::{Deserialize, Deserializer, Serialize, Serializer}; diff --git a/crates/goose-providers/src/lib.rs b/crates/goose-providers/src/lib.rs new file mode 100644 index 00000000..7b77b3a9 --- /dev/null +++ b/crates/goose-providers/src/lib.rs @@ -0,0 +1,4 @@ +pub mod canonical; +pub mod conversation; +mod mcp_utils; +mod utils; diff --git a/crates/goose-providers/src/mcp_utils.rs b/crates/goose-providers/src/mcp_utils.rs new file mode 100644 index 00000000..ad707642 --- /dev/null +++ b/crates/goose-providers/src/mcp_utils.rs @@ -0,0 +1,105 @@ +use base64::Engine; +use rmcp::model::ResourceContents; + +pub fn extract_text_from_resource(resource: &ResourceContents) -> String { + match resource { + ResourceContents::TextResourceContents { text, .. } => text.clone(), + ResourceContents::BlobResourceContents { + blob, mime_type, .. + } => match base64::engine::general_purpose::STANDARD.decode(blob) { + Ok(bytes) => { + let byte_len = bytes.len(); + match String::from_utf8(bytes) { + Ok(text) => text, + Err(_) => { + let mime = mime_type + .as_ref() + .map(|m| m.as_str()) + .unwrap_or("application/octet-stream"); + format!("[Binary content ({}) - {} bytes]", mime, byte_len) + } + } + } + Err(_) => blob.clone(), + }, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use test_case::test_case; + + #[test_case("Hello, World!", "Hello, World!" ; "simple text")] + #[test_case("Hello from GitHub!", "Hello from GitHub!" ; "github content")] + #[test_case("", "" ; "empty text")] + fn test_extract_text_from_text_resource(input: &str, expected: &str) { + let resource = ResourceContents::TextResourceContents { + uri: "file:///test.txt".to_string(), + mime_type: Some("text/plain".to_string()), + text: input.to_string(), + meta: None, + }; + assert_eq!(extract_text_from_resource(&resource), expected); + } + + #[test_case("Hello from GitHub!", "Hello from GitHub!" ; "utf8 markdown")] + #[test_case("Simple text", "Simple text" ; "utf8 plain")] + fn test_extract_text_from_blob_utf8(input: &str, expected: &str) { + let blob = base64::engine::general_purpose::STANDARD.encode(input.as_bytes()); + let resource = ResourceContents::BlobResourceContents { + uri: "github://repo/file.md".to_string(), + mime_type: Some("text/markdown".to_string()), + blob, + meta: None, + }; + assert_eq!(extract_text_from_resource(&resource), expected); + } + + #[test] + fn test_extract_text_from_blob_binary() { + let binary_data: Vec = vec![0xFF, 0xFE, 0x00, 0x01, 0x89, 0x50, 0x4E, 0x47]; + let blob = base64::engine::general_purpose::STANDARD.encode(&binary_data); + + let resource = ResourceContents::BlobResourceContents { + uri: "file:///image.png".to_string(), + mime_type: Some("image/png".to_string()), + blob, + meta: None, + }; + + assert_eq!( + extract_text_from_resource(&resource), + "[Binary content (image/png) - 8 bytes]" + ); + } + + #[test] + fn test_extract_text_from_blob_binary_no_mime_type() { + let binary_data: Vec = vec![0xFF, 0xFE]; + let blob = base64::engine::general_purpose::STANDARD.encode(&binary_data); + + let resource = ResourceContents::BlobResourceContents { + uri: "file:///unknown".to_string(), + mime_type: None, + blob, + meta: None, + }; + + assert_eq!( + extract_text_from_resource(&resource), + "[Binary content (application/octet-stream) - 2 bytes]" + ); + } + + #[test] + fn test_extract_text_from_blob_invalid_base64() { + let resource = ResourceContents::BlobResourceContents { + uri: "file:///test.txt".to_string(), + mime_type: Some("text/plain".to_string()), + blob: "not valid base64!!!".to_string(), + meta: None, + }; + assert_eq!(extract_text_from_resource(&resource), "not valid base64!!!"); + } +} diff --git a/crates/goose-providers/src/utils.rs b/crates/goose-providers/src/utils.rs new file mode 100644 index 00000000..3768ad6a --- /dev/null +++ b/crates/goose-providers/src/utils.rs @@ -0,0 +1,14 @@ +use unicode_normalization::UnicodeNormalization; + +fn is_in_unicode_tag_range(c: char) -> bool { + matches!(c, '\u{E0000}'..='\u{E007F}') +} + +pub fn sanitize_unicode_tags(text: &str) -> String { + let normalized: String = text.nfc().collect(); + + normalized + .chars() + .filter(|&c| !is_in_unicode_tag_range(c)) + .collect() +} diff --git a/crates/goose/Cargo.toml b/crates/goose/Cargo.toml index 90d12ae0..4f3b6842 100644 --- a/crates/goose/Cargo.toml +++ b/crates/goose/Cargo.toml @@ -122,6 +122,7 @@ strum = { workspace = true } once_cell = { workspace = true } etcetera = { workspace = true } fs-err = { version = "3.1", default-features = false } +goose-providers = { path = "../goose-providers", default-features = false } goose-sdk-types = { path = "../goose-sdk-types" } rand = { workspace = true } utoipa = { workspace = true, features = ["chrono"] } @@ -267,7 +268,7 @@ path = "src/bin/analyze_cli.rs" [[bin]] name = "build_canonical_models" -path = "src/providers/canonical/build_canonical_models.rs" +path = "src/bin/build_canonical_models.rs" [[bin]] name = "generate-acp-schema" diff --git a/crates/goose/src/providers/canonical/build_canonical_models.rs b/crates/goose/src/bin/build_canonical_models.rs similarity index 99% rename from crates/goose/src/providers/canonical/build_canonical_models.rs rename to crates/goose/src/bin/build_canonical_models.rs index cea82f16..042917f2 100644 --- a/crates/goose/src/providers/canonical/build_canonical_models.rs +++ b/crates/goose/src/bin/build_canonical_models.rs @@ -9,10 +9,11 @@ /// use anyhow::{Context, Result}; use clap::Parser; -use goose::providers::canonical::{ - canonical_name, CanonicalModel, CanonicalModelRegistry, Limit, Modalities, Modality, Pricing, +use goose::providers::create_with_named_model; +use goose_providers::canonical::{ + canonical_name, CanonicalModel, CanonicalModelRegistry, Limit, Modalities, Modality, + ModelMapping, Pricing, }; -use goose::providers::{canonical::ModelMapping, create_with_named_model}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::{BTreeMap, BTreeSet, HashMap}; @@ -320,7 +321,7 @@ impl MappingReport { fn data_file_path(filename: &str) -> PathBuf { PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .join("src/providers/canonical/data") + .join("../goose-providers/src/canonical/data") .join(filename) } diff --git a/crates/goose/src/lib.rs b/crates/goose/src/lib.rs index 4b19247d..c610d4a7 100644 --- a/crates/goose/src/lib.rs +++ b/crates/goose/src/lib.rs @@ -9,7 +9,9 @@ pub mod builtin_extension; pub mod checks; pub mod config; pub mod context_mgmt; -pub mod conversation; +pub mod conversation { + pub use goose_providers::conversation::*; +} pub mod dictation; pub mod doctor; pub mod download_manager; diff --git a/crates/goose/src/providers/catalog_util.rs b/crates/goose/src/providers/catalog_util.rs new file mode 100644 index 00000000..8e306e0c --- /dev/null +++ b/crates/goose/src/providers/catalog_util.rs @@ -0,0 +1,205 @@ +pub use goose_providers::canonical::catalog::{ + ModelCapabilities, ModelTemplate, ProviderCatalogEntry, ProviderFormat, + ProviderSetupCapabilities, ProviderSetupCatalogEntry, ProviderSetupCategory, + ProviderSetupConfigKey, ProviderSetupField, ProviderSetupGroup, ProviderSetupMetadata, + ProviderSetupMethod, ProviderTemplate, +}; +use std::collections::{HashMap, HashSet}; + +use super::base::{ConfigKey, ProviderMetadata}; + +fn setup_config_key(config_key: ConfigKey) -> ProviderSetupConfigKey { + ProviderSetupConfigKey { + name: config_key.name, + required: config_key.required, + secret: config_key.secret, + default: config_key.default, + primary: config_key.primary, + } +} + +fn setup_metadata(metadata: ProviderMetadata) -> ProviderSetupMetadata { + ProviderSetupMetadata { + name: metadata.name, + display_name: metadata.display_name, + description: metadata.description, + model_doc_link: metadata.model_doc_link, + config_keys: metadata + .config_keys + .into_iter() + .map(setup_config_key) + .collect(), + } +} + +pub async fn get_providers_by_format(format: ProviderFormat) -> Vec { + let native_provider_ids = super::init::providers() + .await + .into_iter() + .map(|(metadata, _)| metadata.name) + .collect::>(); + + goose_providers::canonical::catalog::get_providers_by_format(format, &native_provider_ids) +} + +pub async fn get_setup_catalog_entries() -> Vec { + let registry_metadata = super::providers() + .await + .into_iter() + .map(|(metadata, _)| { + let name = metadata.name.clone(); + (name, setup_metadata(metadata)) + }) + .collect::>(); + + goose_providers::canonical::catalog::get_setup_catalog_entries(®istry_metadata) +} + +pub fn get_provider_setup_category(provider_id: &str) -> Option { + goose_providers::canonical::catalog::get_provider_setup_category(provider_id) +} + +pub fn get_provider_template(provider_id: &str) -> Option { + goose_providers::canonical::catalog::get_provider_template(provider_id) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::providers::base::ProviderType; + + #[tokio::test] + async fn test_zai_provider() { + let zai = crate::providers::get_from_registry("zai") + .await + .expect("z.ai should be registered as a declarative provider"); + assert_eq!(zai.provider_type(), ProviderType::Declarative); + + let metadata = zai.metadata(); + assert_eq!(metadata.display_name, "Z.AI"); + assert!( + !metadata.known_models.is_empty(), + "z.ai should have known models" + ); + assert!( + metadata + .config_keys + .iter() + .any(|key| key.name == "ZHIPU_API_KEY"), + "z.ai should expose its API key config" + ); + + let setup_entries = get_setup_catalog_entries().await; + let setup_entry = setup_entries + .iter() + .find(|entry| entry.provider_id == "zai") + .expect("z.ai should be in the setup catalog"); + assert_eq!(setup_entry.setup_method, ProviderSetupMethod::SingleApiKey); + + let template = get_provider_template("zai"); + assert!(template.is_some(), "z.ai should have a template"); + + let template = template.unwrap(); + println!("Z.AI template: {} models", template.models.len()); + for model in template.models.iter().take(3) { + println!( + " - {} ({}K context)", + model.name, + model.context_limit / 1000 + ); + } + assert!( + !template.models.is_empty(), + "z.ai template should have models" + ); + } + + #[tokio::test] + async fn setup_catalog_includes_goose_and_curated_fields() { + let entries = get_setup_catalog_entries().await; + + let goose = entries + .iter() + .find(|entry| entry.provider_id == "goose") + .expect("setup catalog should include synthetic goose"); + assert_eq!(goose.category, ProviderSetupCategory::Agent); + assert_eq!(goose.setup_method, ProviderSetupMethod::None); + assert!(goose.fields.is_empty()); + + let ollama = entries + .iter() + .find(|entry| entry.provider_id == "ollama") + .expect("setup catalog should include ollama"); + assert_eq!(ollama.setup_method, ProviderSetupMethod::ConfigFields); + assert_eq!(ollama.fields.len(), 1); + assert_eq!(ollama.fields[0].key, "OLLAMA_HOST"); + assert_eq!(ollama.fields[0].label, "Host"); + assert_eq!( + ollama.fields[0].default_value.as_deref(), + Some("http://localhost:11434") + ); + + let databricks = entries + .iter() + .find(|entry| entry.provider_id == "databricks") + .expect("setup catalog should include databricks"); + assert_eq!( + databricks.setup_method, + ProviderSetupMethod::HostWithOauthFallback + ); + assert_eq!( + databricks + .fields + .iter() + .map(|field| field.key.as_str()) + .collect::>(), + ["DATABRICKS_HOST", "DATABRICKS_TOKEN"] + ); + + let huggingface = entries + .iter() + .find(|entry| entry.provider_id == "huggingface") + .expect("setup catalog should include huggingface"); + assert_eq!(huggingface.setup_method, ProviderSetupMethod::SingleApiKey); + assert_eq!( + huggingface + .fields + .iter() + .map(|field| field.key.as_str()) + .collect::>(), + ["HF_TOKEN"] + ); + + let atomic_chat = entries + .iter() + .find(|entry| entry.provider_id == "atomic_chat") + .expect("setup catalog should include atomic_chat declarative provider"); + assert_eq!(atomic_chat.setup_method, ProviderSetupMethod::ConfigFields); + let host_field = atomic_chat + .fields + .iter() + .find(|field| field.key == "ATOMIC_CHAT_HOST") + .expect("atomic_chat should expose ATOMIC_CHAT_HOST"); + assert_eq!(host_field.label, "Host URL"); + assert_eq!( + host_field.default_value.as_deref(), + Some("http://localhost:1337") + ); + } + + #[tokio::test] + async fn setup_catalog_excludes_uncurated_deprecated_providers() { + let provider_ids = get_setup_catalog_entries() + .await + .into_iter() + .map(|entry| entry.provider_id) + .collect::>(); + + assert!(provider_ids.contains("claude-acp")); + assert!(provider_ids.contains("codex-acp")); + assert!(provider_ids.contains("atomic_chat")); + assert!(!provider_ids.contains("claude_code")); + assert!(!provider_ids.contains("codex")); + assert!(!provider_ids.contains("gemini_cli")); + } +} diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 1ca4d793..359d0465 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -8,8 +8,13 @@ pub mod azureauth; pub mod base; #[cfg(feature = "aws-providers")] pub mod bedrock; -pub mod canonical; -pub mod catalog; +pub mod canonical { + pub use goose_providers::canonical::*; +} +mod catalog_util; +pub mod catalog { + pub use super::catalog_util::*; +} pub mod chatgpt_codex; pub mod claude_acp; pub mod claude_code;