From 71b59805b34ca24e30613284782bc4e8ff0bcb89 Mon Sep 17 00:00:00 2001 From: juanxincai <332209078@qq.com> Date: Wed, 8 Apr 2026 23:10:22 +0800 Subject: [PATCH] fix: apply declarative provider context_limit via canonical create path (#8364) Signed-off-by: Douwe Osinga Co-authored-by: juanxincai Co-authored-by: Douwe Osinga --- crates/goose/src/providers/init.rs | 64 ++++++++++++++++++- .../goose/src/providers/provider_registry.rs | 32 ++++++++-- 2 files changed, 89 insertions(+), 7 deletions(-) diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index 34febc8f..afb9db27 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -144,8 +144,8 @@ pub async fn create( model: ModelConfig, extensions: Vec, ) -> Result> { - let constructor = get_from_registry(name).await?.constructor.clone(); - constructor(model, extensions).await + let entry = get_from_registry(name).await?; + entry.create(model, extensions).await } pub async fn create_with_default_model( @@ -177,13 +177,15 @@ pub async fn create_with_named_model( model_name: &str, extensions: Vec, ) -> Result> { - let config = ModelConfig::new(model_name)?.with_canonical_limits(provider_name); + let config = ModelConfig::new(model_name)?; create(provider_name, config, extensions).await } #[cfg(test)] mod tests { use super::*; + use crate::config::paths::Paths; + use std::fs; #[tokio::test] async fn test_tanzu_declarative_provider_registry_wiring() { @@ -274,4 +276,60 @@ mod tests { ); } } + + #[tokio::test] + async fn test_custom_provider_context_limit_is_applied_from_file() { + let _guard = env_lock::lock_env([("GOOSE_PATH_ROOT", None::<&str>)]); + let temp_dir = tempfile::tempdir().expect("tempdir should be created"); + std::env::set_var("GOOSE_PATH_ROOT", temp_dir.path()); + + let custom_dir = Paths::config_dir().join("custom_providers"); + fs::create_dir_all(&custom_dir).expect("custom providers dir should be created"); + + let custom_inf = r#"{ + "name": "custom_inf", + "engine": "openai", + "display_name": "Custom Inf", + "description": "test provider", + "api_key_env": "", + "base_url": "https://example.invalid/v1/chat/completions", + "models": [ + {"name": "kimi-k2.5", "context_limit": 256000} + ], + "requires_auth": false +}"#; + fs::write(custom_dir.join("custom_inf.json"), custom_inf) + .expect("custom_inf.json should be written"); + + let custom_zero = r#"{ + "name": "custom_zero", + "engine": "openai", + "display_name": "Custom Zero", + "description": "test provider", + "api_key_env": "", + "base_url": "https://example.invalid/v1/chat/completions", + "models": [ + {"name": "zero-model", "context_limit": 0} + ], + "requires_auth": false +}"#; + fs::write(custom_dir.join("custom_zero.json"), custom_zero) + .expect("custom_zero.json should be written"); + + refresh_custom_providers() + .await + .expect("custom providers should refresh"); + + let provider = create_with_named_model("custom_inf", "kimi-k2.5", Vec::new()) + .await + .expect("custom_inf provider should be creatable"); + assert_eq!(provider.get_model_config().context_limit, Some(256_000)); + + let zero_provider = create_with_named_model("custom_zero", "zero-model", Vec::new()) + .await + .expect("custom_zero provider should be creatable"); + assert_eq!(zero_provider.get_model_config().context_limit, None); + + std::env::remove_var("GOOSE_PATH_ROOT"); + } } diff --git a/crates/goose/src/providers/provider_registry.rs b/crates/goose/src/providers/provider_registry.rs index 860421c4..be2c5a23 100644 --- a/crates/goose/src/providers/provider_registry.rs +++ b/crates/goose/src/providers/provider_registry.rs @@ -27,16 +27,40 @@ impl ProviderEntry { &self.metadata } + fn normalize_model_config(&self, mut model: ModelConfig) -> ModelConfig { + model = model.with_canonical_limits(&self.metadata.name); + + if model.context_limit.is_none() { + if let Some(info) = self + .metadata + .known_models + .iter() + .find(|m| m.name == model.model_name && m.context_limit > 0) + { + model.context_limit = Some(info.context_limit); + } + } + + model + } + pub async fn create_with_default_model( &self, extensions: Vec, ) -> Result> { let default_model = &self.metadata.default_model; - let provider_name = &self.metadata.name; - let model_config = - ModelConfig::new(default_model.as_str())?.with_canonical_limits(provider_name); + let model_config = self.normalize_model_config(ModelConfig::new(default_model.as_str())?); (self.constructor)(model_config, extensions).await } + + pub async fn create( + &self, + model: ModelConfig, + extensions: Vec, + ) -> Result> { + let model = self.normalize_model_config(model); + (self.constructor)(model, extensions).await + } } #[derive(Default)] @@ -194,7 +218,7 @@ impl ProviderRegistry { .get(name) .ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", name))?; - (entry.constructor)(model, extensions).await + entry.create(model, extensions).await } pub fn all_metadata_with_types(&self) -> Vec<(ProviderMetadata, ProviderType)> {