fix: apply declarative provider context_limit via canonical create path (#8364)

Signed-off-by: Douwe Osinga <douwe@squareup.com>
Co-authored-by: juanxincai <zarttic@users.noreply.github.com>
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
juanxincai
2026-04-08 23:10:22 +08:00
committed by GitHub
parent 4506d00eff
commit 71b59805b3
2 changed files with 89 additions and 7 deletions
+61 -3
View File
@@ -144,8 +144,8 @@ pub async fn create(
model: ModelConfig,
extensions: Vec<ExtensionConfig>,
) -> Result<Arc<dyn Provider>> {
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<ExtensionConfig>,
) -> Result<Arc<dyn Provider>> {
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");
}
}
@@ -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<ExtensionConfig>,
) -> Result<Arc<dyn Provider>> {
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<ExtensionConfig>,
) -> Result<Arc<dyn Provider>> {
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)> {