create goose-providers crate with canonical models, conversation and other types (#9588)
This commit is contained in:
Generated
+21
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 }
|
||||
+3
-3
@@ -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.
|
||||
+34
-164
@@ -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<ProviderSetupConfigKey>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ProviderSetupConfigKey {
|
||||
pub name: String,
|
||||
pub required: bool,
|
||||
pub secret: bool,
|
||||
pub default: Option<String>,
|
||||
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<ProviderCatalogEntry> {
|
||||
let native_provider_ids = super::init::providers()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|(metadata, _)| metadata.name)
|
||||
.collect::<std::collections::HashSet<_>>();
|
||||
|
||||
pub fn get_providers_by_format(
|
||||
format: ProviderFormat,
|
||||
native_provider_ids: &HashSet<String>,
|
||||
) -> Vec<ProviderCatalogEntry> {
|
||||
let mut entries: Vec<ProviderCatalogEntry> = PROVIDER_METADATA
|
||||
.values()
|
||||
.filter_map(|metadata| {
|
||||
@@ -1038,13 +1053,9 @@ pub async fn get_providers_by_format(format: ProviderFormat) -> Vec<ProviderCata
|
||||
entries
|
||||
}
|
||||
|
||||
pub async fn get_setup_catalog_entries() -> Vec<ProviderSetupCatalogEntry> {
|
||||
let registry_metadata = super::providers()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|(metadata, _)| (metadata.name.clone(), metadata))
|
||||
.collect::<HashMap<_, _>>();
|
||||
|
||||
pub fn get_setup_catalog_entries(
|
||||
registry_metadata: &HashMap<String, ProviderSetupMetadata>,
|
||||
) -> Vec<ProviderSetupCatalogEntry> {
|
||||
SETUP_METADATA
|
||||
.iter()
|
||||
.filter_map(|curated| {
|
||||
@@ -1122,144 +1133,3 @@ pub fn get_provider_template(provider_id: &str) -> Option<ProviderTemplate> {
|
||||
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::<Vec<_>>(),
|
||||
["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::<Vec<_>>(),
|
||||
["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::<std::collections::HashSet<_>>();
|
||||
|
||||
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"));
|
||||
}
|
||||
}
|
||||
+1
@@ -1,3 +1,4 @@
|
||||
pub mod catalog;
|
||||
mod model;
|
||||
mod name_builder;
|
||||
mod registry;
|
||||
+2
-1
@@ -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<String, serde_json::Value>;
|
||||
pub type ToolResult<T> = Result<T, rmcp::model::ErrorData>;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
+1
-1
@@ -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};
|
||||
@@ -0,0 +1,4 @@
|
||||
pub mod canonical;
|
||||
pub mod conversation;
|
||||
mod mcp_utils;
|
||||
mod utils;
|
||||
@@ -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<u8> = 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<u8> = 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!!!");
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
+5
-4
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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<ProviderCatalogEntry> {
|
||||
let native_provider_ids = super::init::providers()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|(metadata, _)| metadata.name)
|
||||
.collect::<HashSet<_>>();
|
||||
|
||||
goose_providers::canonical::catalog::get_providers_by_format(format, &native_provider_ids)
|
||||
}
|
||||
|
||||
pub async fn get_setup_catalog_entries() -> Vec<ProviderSetupCatalogEntry> {
|
||||
let registry_metadata = super::providers()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|(metadata, _)| {
|
||||
let name = metadata.name.clone();
|
||||
(name, setup_metadata(metadata))
|
||||
})
|
||||
.collect::<HashMap<_, _>>();
|
||||
|
||||
goose_providers::canonical::catalog::get_setup_catalog_entries(®istry_metadata)
|
||||
}
|
||||
|
||||
pub fn get_provider_setup_category(provider_id: &str) -> Option<ProviderSetupCategory> {
|
||||
goose_providers::canonical::catalog::get_provider_setup_category(provider_id)
|
||||
}
|
||||
|
||||
pub fn get_provider_template(provider_id: &str) -> Option<ProviderTemplate> {
|
||||
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::<Vec<_>>(),
|
||||
["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::<Vec<_>>(),
|
||||
["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::<std::collections::HashSet<_>>();
|
||||
|
||||
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"));
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user