From 78044d0df354f1820bc27f2ea63aa47e346b8752 Mon Sep 17 00:00:00 2001 From: dongjiang Date: Mon, 3 Aug 2026 22:36:09 +0800 Subject: [PATCH] fall back to static model list only for non-models payloads (#10189) Signed-off-by: dongjiang Co-authored-by: Douwe M Osinga --- crates/goose-providers/src/anthropic.rs | 303 ++++++++++++++++++++++-- crates/goose-providers/src/openai.rs | 198 +++++++++++++++- 2 files changed, 474 insertions(+), 27 deletions(-) diff --git a/crates/goose-providers/src/anthropic.rs b/crates/goose-providers/src/anthropic.rs index ee35ec273..09391acdf 100644 --- a/crates/goose-providers/src/anthropic.rs +++ b/crates/goose-providers/src/anthropic.rs @@ -20,7 +20,6 @@ use super::formats::anthropic::{ ANTHROPIC_PROVIDER_NAME, }; use super::openai_compatible::handle_status; -use super::openai_compatible::map_http_error_to_provider_error; use super::retry::ProviderRetry; use crate::conversation::message::Message; use crate::model::ModelConfig; @@ -213,33 +212,63 @@ impl AnthropicProvider { } async fn fetch_models_from_api(&self) -> Result, ProviderError> { - let response = self.api_client.request("v1/models").api_get().await?; + let response = self.api_client.request("v1/models").response_get().await?; - if response.status == StatusCode::NOT_FOUND { - let msg = response - .payload - .as_ref() - .and_then(|p| p.get("error").and_then(|e| e.get("message"))) - .and_then(|m| m.as_str()) - .unwrap_or("models endpoint not found") - .to_string(); + if response.status() == StatusCode::NOT_FOUND { + let body = response.text().await.unwrap_or_default(); + let msg = serde_json::from_str::(&body) + .ok() + .and_then(|p| { + p.get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + .map(String::from) + }) + .unwrap_or_else(|| "models endpoint not found".to_string()); return Err(ProviderError::EndpointNotFound(msg)); } - if response.status != StatusCode::OK { - return Err(map_http_error_to_provider_error( - response.status, - response.payload, - "v1/models", - )); + let response = handle_status(response).await?; + + let body = response.bytes().await.map_err(|e| { + ProviderError::NetworkError(format!("Failed to read response body: {}", e)) + })?; + let json: Value = serde_json::from_slice(&body).map_err(|e| { + ProviderError::EndpointNotFound(format!("Response body is not valid JSON: {}", e)) + })?; + + if let Some(err_obj) = json.get("error").filter(|error| !error.is_null()) { + let message = err_obj + .get("message") + .and_then(Value::as_str) + .unwrap_or("unknown error") + .to_string(); + let error_type = err_obj.get("type").and_then(Value::as_str); + return Err(match error_type { + Some("authentication_error" | "permission_error") => { + ProviderError::Authentication(message) + } + Some("rate_limit_error") => ProviderError::RateLimitExceeded { + details: message, + retry_delay: None, + }, + Some("billing_error") => ProviderError::CreditsExhausted { + details: message, + top_up_url: None, + }, + Some("api_error" | "overloaded_error") => ProviderError::ServerError(message), + _ => ProviderError::RequestFailed(message), + }); } - let json = response.payload.unwrap_or_default(); - let arr = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| { - ProviderError::RequestFailed( - "Missing 'data' array in Anthropic models response".to_string(), - ) - })?; + let arr = match json.get("data").and_then(|v| v.as_array()) { + Some(arr) => arr, + None => { + return Err(ProviderError::RequestFailed( + "response is not a models payload (missing 'data' array)".into(), + )); + } + }; let mut models: Vec = arr .iter() @@ -441,3 +470,233 @@ pub fn from_declarative_config( .skip_canonical_filtering(config.skip_canonical_filtering) .format_options(format_options)) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::api_client::AuthMethod; + use serde_json::json; + + fn make_provider_with_custom_models( + host: &str, + custom_models: Vec, + ) -> AnthropicProvider { + AnthropicProvider { + api_client: ApiClient::new_with_tls(host.to_string(), AuthMethod::NoAuth, None) + .unwrap(), + supports_streaming: true, + name: "test-provider".to_string(), + custom_models: Some(custom_models), + dynamic_models: Some(true), + skip_canonical_filtering: false, + format_options: AnthropicFormatOptions::default(), + } + } + + #[tokio::test] + async fn fetch_models_treats_invalid_json_as_endpoint_not_found() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with( + ResponseTemplate::new(200).set_body_string("not a models endpoint"), + ) + .expect(1) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + let err = provider.fetch_models_from_api().await.unwrap_err(); + assert!( + err.is_endpoint_not_found(), + "expected EndpointNotFound, got: {:?}", + err + ); + } + + #[tokio::test] + async fn fetch_models_treats_missing_data_field_as_request_failed() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"status": "ok"}))) + .expect(1) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + let err = provider.fetch_models_from_api().await.unwrap_err(); + assert!( + matches!(err, ProviderError::RequestFailed(_)), + "expected RequestFailed, got: {:?}", + err + ); + } + + #[tokio::test] + async fn fetch_supported_models_falls_back_on_invalid_payload() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_string("error page")) + .mount(&server) + .await; + + let predefined = vec![ + "claude-sonnet-4-5".to_string(), + "claude-haiku-4-5".to_string(), + ]; + let provider = make_provider_with_custom_models(&server.uri(), predefined.clone()); + + let models = provider + .fetch_supported_models() + .await + .expect("should fall back to predefined list on invalid payload"); + assert_eq!(models, predefined); + } + + #[tokio::test] + async fn fetch_supported_models_does_not_fall_back_on_missing_data() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"status": "ok"}))) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + let err = provider.fetch_supported_models().await.unwrap_err(); + assert!( + matches!(err, ProviderError::RequestFailed(_)), + "expected RequestFailed to propagate, got: {:?}", + err + ); + } + + #[tokio::test] + async fn fetch_supported_models_propagates_auth_error() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(401).set_body_json(json!({ + "type": "error", + "error": { + "type": "authentication_error", + "message": "invalid api key" + } + }))) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + let err = provider.fetch_supported_models().await.unwrap_err(); + assert!( + matches!(err, ProviderError::Authentication(_)), + "expected Authentication error, got: {:?}", + err + ); + } + + #[tokio::test] + async fn fetch_supported_models_accepts_null_error() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "data": [{"id": "model-a"}], + "error": null + }))) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + assert_eq!( + provider.fetch_supported_models().await.unwrap(), + vec!["model-a".to_string()] + ); + } + + #[tokio::test] + async fn fetch_supported_models_preserves_200_error_type() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "type": "error", + "error": { + "type": "rate_limit_error", + "message": "quota exceeded" + } + }))) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + assert!(matches!( + provider.fetch_supported_models().await.unwrap_err(), + ProviderError::RateLimitExceeded { .. } + )); + } + + #[tokio::test] + async fn fetch_supported_models_propagates_auth_error_from_200_payload() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "type": "error", + "error": { + "type": "authentication_error", + "message": "invalid api key" + } + }))) + .mount(&server) + .await; + + let provider = + make_provider_with_custom_models(&server.uri(), vec!["static-model".to_string()]); + + let err = provider.fetch_supported_models().await.unwrap_err(); + assert!( + matches!(err, ProviderError::Authentication(_)), + "expected Authentication error, got: {:?}", + err + ); + } +} diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index f36351cec..15a5c9832 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -543,8 +543,16 @@ impl OpenAiProvider { return Err(ProviderError::EndpointNotFound(body)); } - let json = handle_response_openai_compat(response).await?; - if let Some(err_obj) = json.get("error") { + let response = handle_status(response).await?; + + let body = response.bytes().await.map_err(|e| { + ProviderError::NetworkError(format!("Failed to read response body: {}", e)) + })?; + let json: serde_json::Value = serde_json::from_slice(&body).map_err(|e| { + ProviderError::EndpointNotFound(format!("Response body is not valid JSON: {}", e)) + })?; + + if let Some(err_obj) = json.get("error").filter(|error| !error.is_null()) { let msg = err_obj .get("message") .and_then(|v| v.as_str()) @@ -1415,9 +1423,6 @@ mod tests { #[test] fn derive_base_path_preserves_non_v1_version_prefix() { - // Zhipu's default base_url is https://open.bigmodel.cn/api/paas/v4 and - // from_custom_config passes url.path() ("/api/paas/v4") here. The - // existing /api/paas/v4 version must not gain an extra /v1 segment. let r = derive_base_path("/api/paas/v4"); assert_eq!(r, "api/paas/v4/chat/completions"); } @@ -1428,6 +1433,189 @@ mod tests { assert_eq!(r, "api/voice/v1/chat/completions"); } + fn make_provider_with_custom_models( + host: &str, + base_path: &str, + custom_models: Vec, + ) -> OpenAiProvider { + OpenAiProvider { + api_client: ApiClient::new_with_tls(host.to_string(), AuthMethod::NoAuth, None) + .unwrap(), + base_path: base_path.to_string(), + organization: None, + project: None, + custom_headers: None, + supports_streaming: true, + name: "test-provider".to_string(), + custom_models: Some( + custom_models + .into_iter() + .map(|model| ModelInfo::new(model, 4096)) + .collect(), + ), + dynamic_models: Some(true), + skip_canonical_filtering: false, + preserve_thinking_context: false, + n_ctx_cache: Arc::new(Mutex::new(HashMap::new())), + } + } + + #[tokio::test] + async fn fetch_models_treats_invalid_json_as_endpoint_not_found() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with( + ResponseTemplate::new(200).set_body_string("not a models endpoint"), + ) + .expect(1) + .mount(&server) + .await; + + let provider = make_provider_with_custom_models( + &server.uri(), + "v1/chat/completions", + vec!["static-model".to_string()], + ); + + let err = provider.fetch_models_from_api().await.unwrap_err(); + assert!(err.is_endpoint_not_found(), "got: {:?}", err); + } + + #[tokio::test] + async fn fetch_models_returns_request_failed_for_missing_data_field() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"status": "ok"}))) + .expect(1) + .mount(&server) + .await; + + let provider = make_provider_with_custom_models( + &server.uri(), + "v1/chat/completions", + vec!["static-model".to_string()], + ); + + let err = provider.fetch_models_from_api().await.unwrap_err(); + assert!( + matches!(err, ProviderError::RequestFailed(_)), + "expected RequestFailed, got: {:?}", + err + ); + } + + #[tokio::test] + async fn fetch_supported_models_falls_back_on_invalid_payload() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_string("error page")) + .mount(&server) + .await; + + let predefined = vec!["glm-4.5".to_string(), "glm-5".to_string()]; + let provider = make_provider_with_custom_models( + &server.uri(), + "v1/chat/completions", + predefined.clone(), + ); + + let models = provider + .fetch_supported_models() + .await + .expect("should fall back"); + assert_eq!(models, predefined); + } + + #[tokio::test] + async fn fetch_supported_models_propagates_auth_error() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "error": {"message": "invalid api key"} + }))) + .mount(&server) + .await; + + let provider = make_provider_with_custom_models( + &server.uri(), + "v1/chat/completions", + vec!["static-model".to_string()], + ); + + let err = provider.fetch_supported_models().await.unwrap_err(); + assert!( + matches!(err, ProviderError::Authentication(_)), + "got: {:?}", + err + ); + } + + #[tokio::test] + async fn fetch_supported_models_does_not_reclassify_400_as_endpoint_not_found() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": {"message": "request is not valid JSON"} + }))) + .mount(&server) + .await; + + let provider = make_provider_with_custom_models( + &server.uri(), + "v1/chat/completions", + vec!["static-model".to_string()], + ); + + let err = provider.fetch_supported_models().await.unwrap_err(); + assert!(!err.is_endpoint_not_found(), "got: {:?}", err); + } + + #[tokio::test] + async fn fetch_supported_models_accepts_payload_with_extra_fields() { + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/models")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "data": [{"id": "model-a"}, {"id": "model-b"}], + "message": "ok", + "error": null + }))) + .mount(&server) + .await; + + let provider = make_provider_with_custom_models( + &server.uri(), + "v1/chat/completions", + vec!["static-model".to_string()], + ); + + let models = provider.fetch_supported_models().await.unwrap(); + assert_eq!(models, vec!["model-a".to_string(), "model-b".to_string()]); + } + use crate::base::ThinkingPreservationFormat; fn cerebras_config() -> DeclarativeProviderConfig {