fix(providers): honor dynamic_models: false in declarative provider configs (#8795)
Signed-off-by: Douwe Osinga <douwe@squareup.com> Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -86,6 +86,12 @@ pub struct DeclarativeProviderConfig {
|
|||||||
pub base_path: Option<String>,
|
pub base_path: Option<String>,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub env_vars: Option<Vec<EnvVarConfig>>,
|
pub env_vars: Option<Vec<EnvVarConfig>>,
|
||||||
|
/// Controls whether `fetch_supported_models` calls the provider's `/v1/models`
|
||||||
|
/// endpoint or returns the static `models` list directly.
|
||||||
|
///
|
||||||
|
/// - `Some(false)` + non-empty `models`: return the static list; no API call.
|
||||||
|
/// Construction fails if `models` is empty.
|
||||||
|
/// - `Some(true)` or `None`: try the API; fall back to `models` on 404.
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub dynamic_models: Option<bool>,
|
pub dynamic_models: Option<bool>,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ pub struct AnthropicProvider {
|
|||||||
supports_streaming: bool,
|
supports_streaming: bool,
|
||||||
name: String,
|
name: String,
|
||||||
custom_models: Option<Vec<String>>,
|
custom_models: Option<Vec<String>>,
|
||||||
|
dynamic_models: Option<bool>,
|
||||||
skip_canonical_filtering: bool,
|
skip_canonical_filtering: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -84,6 +85,7 @@ impl AnthropicProvider {
|
|||||||
supports_streaming: true,
|
supports_streaming: true,
|
||||||
name: ANTHROPIC_PROVIDER_NAME.to_string(),
|
name: ANTHROPIC_PROVIDER_NAME.to_string(),
|
||||||
custom_models: None,
|
custom_models: None,
|
||||||
|
dynamic_models: None,
|
||||||
skip_canonical_filtering: false,
|
skip_canonical_filtering: false,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -92,6 +94,26 @@ impl AnthropicProvider {
|
|||||||
model: ModelConfig,
|
model: ModelConfig,
|
||||||
config: DeclarativeProviderConfig,
|
config: DeclarativeProviderConfig,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
|
let custom_models = if !config.models.is_empty() {
|
||||||
|
Some(
|
||||||
|
config
|
||||||
|
.models
|
||||||
|
.iter()
|
||||||
|
.map(|m| m.name.clone())
|
||||||
|
.collect::<Vec<String>>(),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
|
||||||
|
if config.dynamic_models == Some(false) && custom_models.is_none() {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"Provider '{}' has dynamic_models: false but no static models listed; \
|
||||||
|
at least one entry in `models` is required.",
|
||||||
|
config.name
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
let global_config = crate::config::Config::global();
|
let global_config = crate::config::Config::global();
|
||||||
let api_key: String = global_config
|
let api_key: String = global_config
|
||||||
.get_secret(&config.api_key_env)
|
.get_secret(&config.api_key_env)
|
||||||
@@ -124,12 +146,6 @@ impl AnthropicProvider {
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
let custom_models = if !config.models.is_empty() {
|
|
||||||
Some(config.models.iter().map(|m| m.name.clone()).collect())
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
};
|
|
||||||
|
|
||||||
let model = if let Some(ref fast_model_name) = config.fast_model {
|
let model = if let Some(ref fast_model_name) = config.fast_model {
|
||||||
model.with_fast(fast_model_name, &config.name)?
|
model.with_fast(fast_model_name, &config.name)?
|
||||||
} else {
|
} else {
|
||||||
@@ -142,6 +158,7 @@ impl AnthropicProvider {
|
|||||||
supports_streaming,
|
supports_streaming,
|
||||||
name: config.name.clone(),
|
name: config.name.clone(),
|
||||||
custom_models,
|
custom_models,
|
||||||
|
dynamic_models: config.dynamic_models,
|
||||||
skip_canonical_filtering: config.skip_canonical_filtering,
|
skip_canonical_filtering: config.skip_canonical_filtering,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -281,6 +298,9 @@ impl Provider for AnthropicProvider {
|
|||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
if let Some(custom_models) = &self.custom_models {
|
if let Some(custom_models) = &self.custom_models {
|
||||||
|
if self.dynamic_models == Some(false) {
|
||||||
|
return Ok(custom_models.clone());
|
||||||
|
}
|
||||||
match self.fetch_models_from_api().await {
|
match self.fetch_models_from_api().await {
|
||||||
Ok(models) => return Ok(models),
|
Ok(models) => return Ok(models),
|
||||||
Err(e) if e.is_endpoint_not_found() => {
|
Err(e) if e.is_endpoint_not_found() => {
|
||||||
@@ -345,3 +365,94 @@ impl Provider for AnthropicProvider {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::config::declarative_providers::{DeclarativeProviderConfig, ProviderEngine};
|
||||||
|
use wiremock::matchers::method;
|
||||||
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
|
fn make_provider_with_server(
|
||||||
|
server_uri: &str,
|
||||||
|
custom_models: Option<Vec<String>>,
|
||||||
|
dynamic_models: Option<bool>,
|
||||||
|
) -> AnthropicProvider {
|
||||||
|
let auth = AuthMethod::ApiKey {
|
||||||
|
header_name: "x-api-key".to_string(),
|
||||||
|
key: "test-key".to_string(),
|
||||||
|
};
|
||||||
|
let api_client = ApiClient::new(server_uri.to_string(), auth)
|
||||||
|
.unwrap()
|
||||||
|
.with_header("anthropic-version", ANTHROPIC_API_VERSION)
|
||||||
|
.unwrap();
|
||||||
|
AnthropicProvider {
|
||||||
|
api_client,
|
||||||
|
model: ModelConfig::new_or_fail("claude-test"),
|
||||||
|
supports_streaming: true,
|
||||||
|
name: "custom_anthropic".to_string(),
|
||||||
|
custom_models,
|
||||||
|
dynamic_models,
|
||||||
|
skip_canonical_filtering: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn base_declarative_config(
|
||||||
|
models: Vec<ModelInfo>,
|
||||||
|
dynamic_models: Option<bool>,
|
||||||
|
) -> DeclarativeProviderConfig {
|
||||||
|
DeclarativeProviderConfig {
|
||||||
|
name: "custom_anthropic".to_string(),
|
||||||
|
engine: ProviderEngine::Anthropic,
|
||||||
|
display_name: "Custom Anthropic".to_string(),
|
||||||
|
description: None,
|
||||||
|
api_key_env: String::new(),
|
||||||
|
base_url: "http://localhost:1".to_string(),
|
||||||
|
models,
|
||||||
|
headers: None,
|
||||||
|
timeout_seconds: None,
|
||||||
|
supports_streaming: Some(true),
|
||||||
|
requires_auth: false,
|
||||||
|
catalog_provider_id: None,
|
||||||
|
base_path: None,
|
||||||
|
env_vars: None,
|
||||||
|
dynamic_models,
|
||||||
|
skip_canonical_filtering: false,
|
||||||
|
model_doc_link: None,
|
||||||
|
setup_steps: vec![],
|
||||||
|
fast_model: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn fetch_supported_models_static_only_skips_api() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("GET"))
|
||||||
|
.respond_with(ResponseTemplate::new(500))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = make_provider_with_server(
|
||||||
|
&server.uri(),
|
||||||
|
Some(vec!["m1".to_string(), "m2".to_string()]),
|
||||||
|
Some(false),
|
||||||
|
);
|
||||||
|
|
||||||
|
let models = provider.fetch_supported_models().await.unwrap();
|
||||||
|
assert_eq!(models, vec!["m1".to_string(), "m2".to_string()]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn from_custom_config_rejects_static_only_without_models() {
|
||||||
|
let config = base_declarative_config(vec![], Some(false));
|
||||||
|
let err =
|
||||||
|
AnthropicProvider::from_custom_config(ModelConfig::new_or_fail("claude-test"), config)
|
||||||
|
.err()
|
||||||
|
.expect("expected construction error for dynamic_models: false with empty models");
|
||||||
|
let msg = err.to_string();
|
||||||
|
assert!(
|
||||||
|
msg.contains("dynamic_models: false"),
|
||||||
|
"error message should mention dynamic_models: false; got: {msg}"
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -125,6 +125,7 @@ pub struct OpenAiProvider {
|
|||||||
supports_streaming: bool,
|
supports_streaming: bool,
|
||||||
name: String,
|
name: String,
|
||||||
custom_models: Option<Vec<String>>,
|
custom_models: Option<Vec<String>>,
|
||||||
|
dynamic_models: Option<bool>,
|
||||||
skip_canonical_filtering: bool,
|
skip_canonical_filtering: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -277,6 +278,7 @@ impl OpenAiProvider {
|
|||||||
supports_streaming: true,
|
supports_streaming: true,
|
||||||
name: OPEN_AI_PROVIDER_NAME.to_string(),
|
name: OPEN_AI_PROVIDER_NAME.to_string(),
|
||||||
custom_models: None,
|
custom_models: None,
|
||||||
|
dynamic_models: None,
|
||||||
skip_canonical_filtering: false,
|
skip_canonical_filtering: false,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -293,6 +295,7 @@ impl OpenAiProvider {
|
|||||||
supports_streaming: true,
|
supports_streaming: true,
|
||||||
name: OPEN_AI_PROVIDER_NAME.to_string(),
|
name: OPEN_AI_PROVIDER_NAME.to_string(),
|
||||||
custom_models: None,
|
custom_models: None,
|
||||||
|
dynamic_models: None,
|
||||||
skip_canonical_filtering: false,
|
skip_canonical_filtering: false,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -301,6 +304,26 @@ impl OpenAiProvider {
|
|||||||
model: ModelConfig,
|
model: ModelConfig,
|
||||||
config: DeclarativeProviderConfig,
|
config: DeclarativeProviderConfig,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
|
let custom_models = if !config.models.is_empty() {
|
||||||
|
Some(
|
||||||
|
config
|
||||||
|
.models
|
||||||
|
.iter()
|
||||||
|
.map(|m| m.name.clone())
|
||||||
|
.collect::<Vec<String>>(),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
|
||||||
|
if config.dynamic_models == Some(false) && custom_models.is_none() {
|
||||||
|
return Err(anyhow::anyhow!(
|
||||||
|
"Provider '{}' has dynamic_models: false but no static models listed; \
|
||||||
|
at least one entry in `models` is required.",
|
||||||
|
config.name
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
let global_config = crate::config::Config::global();
|
let global_config = crate::config::Config::global();
|
||||||
|
|
||||||
let api_key: Option<String> = if config.requires_auth && !config.api_key_env.is_empty() {
|
let api_key: Option<String> = if config.requires_auth && !config.api_key_env.is_empty() {
|
||||||
@@ -360,12 +383,6 @@ impl OpenAiProvider {
|
|||||||
api_client = api_client.with_headers(header_map)?;
|
api_client = api_client.with_headers(header_map)?;
|
||||||
}
|
}
|
||||||
|
|
||||||
let custom_models = if !config.models.is_empty() {
|
|
||||||
Some(config.models.iter().map(|m| m.name.clone()).collect())
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
};
|
|
||||||
|
|
||||||
let model = if let Some(ref fast_model_name) = config.fast_model {
|
let model = if let Some(ref fast_model_name) = config.fast_model {
|
||||||
model.with_fast(fast_model_name, &config.name)?
|
model.with_fast(fast_model_name, &config.name)?
|
||||||
} else {
|
} else {
|
||||||
@@ -382,6 +399,7 @@ impl OpenAiProvider {
|
|||||||
supports_streaming: config.supports_streaming.unwrap_or(true),
|
supports_streaming: config.supports_streaming.unwrap_or(true),
|
||||||
name: config.name.clone(),
|
name: config.name.clone(),
|
||||||
custom_models,
|
custom_models,
|
||||||
|
dynamic_models: config.dynamic_models,
|
||||||
skip_canonical_filtering: config.skip_canonical_filtering,
|
skip_canonical_filtering: config.skip_canonical_filtering,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -657,6 +675,9 @@ impl Provider for OpenAiProvider {
|
|||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
if let Some(custom_models) = &self.custom_models {
|
if let Some(custom_models) = &self.custom_models {
|
||||||
|
if self.dynamic_models == Some(false) {
|
||||||
|
return Ok(custom_models.clone());
|
||||||
|
}
|
||||||
match self.fetch_models_from_api().await {
|
match self.fetch_models_from_api().await {
|
||||||
Ok(models) => return Ok(models),
|
Ok(models) => return Ok(models),
|
||||||
Err(e) if e.is_endpoint_not_found() => {
|
Err(e) if e.is_endpoint_not_found() => {
|
||||||
@@ -904,6 +925,7 @@ mod tests {
|
|||||||
supports_streaming: true,
|
supports_streaming: true,
|
||||||
name: name.to_string(),
|
name: name.to_string(),
|
||||||
custom_models: None,
|
custom_models: None,
|
||||||
|
dynamic_models: None,
|
||||||
skip_canonical_filtering: false,
|
skip_canonical_filtering: false,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1124,4 +1146,91 @@ mod tests {
|
|||||||
"chat/completions"
|
"chat/completions"
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ── dynamic_models behavior ─────────────────────────────────────────────
|
||||||
|
|
||||||
|
use crate::config::declarative_providers::{DeclarativeProviderConfig, ProviderEngine};
|
||||||
|
use wiremock::matchers::method;
|
||||||
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
|
fn make_provider_with_server(
|
||||||
|
server_uri: &str,
|
||||||
|
custom_models: Option<Vec<String>>,
|
||||||
|
dynamic_models: Option<bool>,
|
||||||
|
) -> OpenAiProvider {
|
||||||
|
OpenAiProvider {
|
||||||
|
api_client: ApiClient::new(server_uri.to_string(), AuthMethod::NoAuth).unwrap(),
|
||||||
|
base_path: "v1/chat/completions".to_string(),
|
||||||
|
organization: None,
|
||||||
|
project: None,
|
||||||
|
model: ModelConfig::new_or_fail("test-model"),
|
||||||
|
custom_headers: None,
|
||||||
|
supports_streaming: true,
|
||||||
|
name: "custom_test".to_string(),
|
||||||
|
custom_models,
|
||||||
|
dynamic_models,
|
||||||
|
skip_canonical_filtering: false,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn base_declarative_config(
|
||||||
|
models: Vec<ModelInfo>,
|
||||||
|
dynamic_models: Option<bool>,
|
||||||
|
) -> DeclarativeProviderConfig {
|
||||||
|
DeclarativeProviderConfig {
|
||||||
|
name: "custom_test".to_string(),
|
||||||
|
engine: ProviderEngine::OpenAI,
|
||||||
|
display_name: "Custom Test".to_string(),
|
||||||
|
description: None,
|
||||||
|
api_key_env: String::new(),
|
||||||
|
base_url: "http://localhost:1".to_string(),
|
||||||
|
models,
|
||||||
|
headers: None,
|
||||||
|
timeout_seconds: None,
|
||||||
|
supports_streaming: Some(true),
|
||||||
|
requires_auth: false,
|
||||||
|
catalog_provider_id: None,
|
||||||
|
base_path: None,
|
||||||
|
env_vars: None,
|
||||||
|
dynamic_models,
|
||||||
|
skip_canonical_filtering: false,
|
||||||
|
model_doc_link: None,
|
||||||
|
setup_steps: vec![],
|
||||||
|
fast_model: None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn fetch_supported_models_static_only_skips_api() {
|
||||||
|
// Any request to the mock returns 500 — if the fix calls the API, the test fails.
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("GET"))
|
||||||
|
.respond_with(ResponseTemplate::new(500))
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
let provider = make_provider_with_server(
|
||||||
|
&server.uri(),
|
||||||
|
Some(vec!["m1".to_string(), "m2".to_string()]),
|
||||||
|
Some(false),
|
||||||
|
);
|
||||||
|
|
||||||
|
let models = provider.fetch_supported_models().await.unwrap();
|
||||||
|
assert_eq!(models, vec!["m1".to_string(), "m2".to_string()]);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn from_custom_config_rejects_static_only_without_models() {
|
||||||
|
let config = base_declarative_config(vec![], Some(false));
|
||||||
|
let err =
|
||||||
|
OpenAiProvider::from_custom_config(ModelConfig::new_or_fail("test-model"), config)
|
||||||
|
.expect_err(
|
||||||
|
"expected construction error for dynamic_models: false with empty models",
|
||||||
|
);
|
||||||
|
let msg = err.to_string();
|
||||||
|
assert!(
|
||||||
|
msg.contains("dynamic_models: false"),
|
||||||
|
"error message should mention dynamic_models: false; got: {msg}"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4582,6 +4582,7 @@
|
|||||||
},
|
},
|
||||||
"dynamic_models": {
|
"dynamic_models": {
|
||||||
"type": "boolean",
|
"type": "boolean",
|
||||||
|
"description": "Controls whether `fetch_supported_models` calls the provider's `/v1/models`\nendpoint or returns the static `models` list directly.\n\n- `Some(false)` + non-empty `models`: return the static list; no API call.\nConstruction fails if `models` is empty.\n- `Some(true)` or `None`: try the API; fall back to `models` on 404.",
|
||||||
"nullable": true
|
"nullable": true
|
||||||
},
|
},
|
||||||
"engine": {
|
"engine": {
|
||||||
|
|||||||
@@ -213,6 +213,14 @@ export type DeclarativeProviderConfig = {
|
|||||||
catalog_provider_id?: string | null;
|
catalog_provider_id?: string | null;
|
||||||
description?: string | null;
|
description?: string | null;
|
||||||
display_name: string;
|
display_name: string;
|
||||||
|
/**
|
||||||
|
* Controls whether `fetch_supported_models` calls the provider's `/v1/models`
|
||||||
|
* endpoint or returns the static `models` list directly.
|
||||||
|
*
|
||||||
|
* - `Some(false)` + non-empty `models`: return the static list; no API call.
|
||||||
|
* Construction fails if `models` is empty.
|
||||||
|
* - `Some(true)` or `None`: try the API; fall back to `models` on 404.
|
||||||
|
*/
|
||||||
dynamic_models?: boolean | null;
|
dynamic_models?: boolean | null;
|
||||||
engine: ProviderEngine;
|
engine: ProviderEngine;
|
||||||
env_vars?: Array<EnvVarConfig> | null;
|
env_vars?: Array<EnvVarConfig> | null;
|
||||||
|
|||||||
Reference in New Issue
Block a user