Skip to main content

llm/providers/
generic.rs

1use async_openai::Client;
2use async_openai::config::{Config, OpenAIConfig};
3use reqwest::Url;
4use schemars::Schema;
5
6use crate::catalog::Provider;
7use crate::provider::{error_stream, get_context_window, stream_from, validate_reasoning};
8use crate::providers::http::openai_client;
9use crate::providers::openai_compatible::{
10    AetherOpenAiConfig, PromptCacheKeySource, build_chat_request, create_custom_stream_generic,
11};
12use crate::providers::openai_responses::mappers::build_wire_request;
13use crate::providers::openai_responses::transport::{process_connection, send};
14use crate::tool_schema::normalize_for_moonshot;
15use crate::{
16    Context, LlmError, LlmModel, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, Result,
17    StreamingModelProvider,
18};
19
20pub use crate::providers::openai_responses::mappers::ResponsesRequestPolicy;
21
22pub struct ProviderConfig {
23    pub provider: Provider,
24    pub api_base: Option<&'static str>,
25    pub default_model: &'static str,
26    pub api: Api,
27}
28
29pub enum Api {
30    ChatCompletions { tool_schema_transform: Option<fn(&mut Schema)>, prompt_cache_key: PromptCacheKeySource },
31    Responses(ResponsesRequestPolicy),
32}
33
34pub const OPENAI: ProviderConfig = ProviderConfig {
35    provider: Provider::Openai,
36    api_base: Some("https://api.openai.com/v1"),
37    default_model: "gpt-4.1",
38    api: Api::Responses(ResponsesRequestPolicy::OPENAI),
39};
40
41pub const XIAOMI: ProviderConfig = ProviderConfig {
42    provider: Provider::Xiaomi,
43    api_base: Some("https://api.xiaomimimo.com/v1"),
44    default_model: "mimo-v2.6-pro",
45    api: Api::Responses(ResponsesRequestPolicy::XIAOMI),
46};
47
48pub const DEEPSEEK: ProviderConfig = ProviderConfig {
49    provider: Provider::DeepSeek,
50    api_base: Some("https://api.deepseek.com"),
51    default_model: "deepseek-v4-flash",
52    api: Api::ChatCompletions { tool_schema_transform: None, prompt_cache_key: PromptCacheKeySource::Omit },
53};
54
55pub const MOONSHOT: ProviderConfig = ProviderConfig {
56    provider: Provider::Moonshot,
57    api_base: Some("https://api.moonshot.ai/v1"),
58    default_model: "moonshot-v1-8k",
59    api: Api::ChatCompletions {
60        tool_schema_transform: Some(normalize_for_moonshot),
61        prompt_cache_key: PromptCacheKeySource::Omit,
62    },
63};
64
65pub const ZAI: ProviderConfig = ProviderConfig {
66    provider: Provider::ZAi,
67    api_base: Some("https://api.z.ai/api/coding/paas/v4"),
68    default_model: "GLM-4.6",
69    api: Api::ChatCompletions { tool_schema_transform: None, prompt_cache_key: PromptCacheKeySource::Omit },
70};
71
72pub const AZURE_FOUNDRY: ProviderConfig = ProviderConfig {
73    provider: Provider::AzureFoundry,
74    api_base: None,
75    default_model: "gpt-5.5",
76    api: Api::ChatCompletions { tool_schema_transform: None, prompt_cache_key: PromptCacheKeySource::Prefix },
77};
78
79pub const FIREWORKS: ProviderConfig = ProviderConfig {
80    provider: Provider::Fireworks,
81    api_base: Some("https://api.fireworks.ai/inference/v1"),
82    default_model: "accounts/fireworks/models/glm-5p1",
83    api: Api::ChatCompletions { tool_schema_transform: None, prompt_cache_key: PromptCacheKeySource::SessionAffinity },
84};
85
86pub(crate) const BUILT_INS: &[&ProviderConfig] =
87    &[&OPENAI, &XIAOMI, &DEEPSEEK, &MOONSHOT, &ZAI, &AZURE_FOUNDRY, &FIREWORKS];
88
89/// A provider whose behavior is fully described by a [`ProviderConfig`].
90pub struct GenericProvider {
91    config: &'static ProviderConfig,
92    openai_config: AetherOpenAiConfig,
93    http: reqwest::Client,
94    chat_client: Client<AetherOpenAiConfig>,
95    model: String,
96    request_model: Option<String>,
97}
98
99impl GenericProvider {
100    pub fn from_env(config: &'static ProviderConfig) -> Result<Self> {
101        Self::from_env_with_connection(config, ProviderConnectionConfig::default())
102    }
103
104    pub fn from_env_with_connection(
105        config: &'static ProviderConfig,
106        connection: ProviderConnectionConfig,
107    ) -> Result<Self> {
108        let api_key = match connection.auth_mode {
109            ProviderAuthMode::Default => {
110                let env_var = config.provider.required_env_var().expect("generic providers require an API key");
111                std::env::var(env_var).map_err(|_| LlmError::MissingApiKey(env_var.to_string()))?
112            }
113            ProviderAuthMode::None => String::new(),
114        };
115        Self::new_with_connection(api_key, config, connection)
116    }
117
118    pub fn new(api_key: String, config: &'static ProviderConfig) -> Result<Self> {
119        Self::new_with_connection(api_key, config, ProviderConnectionConfig::default())
120    }
121
122    pub fn new_with_connection(
123        api_key: String,
124        config: &'static ProviderConfig,
125        connection: ProviderConnectionConfig,
126    ) -> Result<Self> {
127        let api_base = connection
128            .base_url
129            .or_else(|| config.api_base.map(str::to_string))
130            .ok_or_else(|| LlmError::MissingProviderUrl { provider: config.provider.parser_name().to_string() })?;
131
132        let openai_config = AetherOpenAiConfig::new(
133            OpenAIConfig::new().with_api_key(api_key).with_api_base(api_base.trim_end_matches('/')),
134            connection.auth_mode,
135        );
136
137        let http = reqwest::Client::new();
138        Ok(Self {
139            config,
140            chat_client: openai_client(openai_config.clone(), http.clone()),
141            openai_config,
142            http,
143            model: config.default_model.to_string(),
144            request_model: connection.request_model,
145        })
146    }
147
148    pub fn with_model(mut self, model: &str) -> Self {
149        if !model.is_empty() {
150            self.model = model.to_string();
151        }
152        self
153    }
154}
155
156impl StreamingModelProvider for GenericProvider {
157    fn stream_response(&self, context: &Context) -> LlmResponseStream {
158        match &self.config.api {
159            Api::ChatCompletions { tool_schema_transform, prompt_cache_key } => {
160                self.stream_chat_completions(context, *tool_schema_transform, *prompt_cache_key)
161            }
162            Api::Responses(policy) => self.stream_responses(context, policy),
163        }
164    }
165
166    fn display_name(&self) -> String {
167        format!("{} ({})", self.config.provider.display_name(), self.model)
168    }
169
170    fn context_window(&self) -> Option<u32> {
171        get_context_window(self.config.provider.parser_name(), &self.model)
172    }
173
174    fn model(&self) -> Option<LlmModel> {
175        format!("{}:{}", self.config.provider.parser_name(), self.model).parse().ok()
176    }
177}
178
179impl GenericProvider {
180    fn stream_chat_completions(
181        &self,
182        context: &Context,
183        tool_schema_transform: Option<fn(&mut Schema)>,
184        prompt_cache_key: PromptCacheKeySource,
185    ) -> LlmResponseStream {
186        if let Err(error) = validate_reasoning(context, self.model().as_ref()) {
187            return error_stream(error);
188        }
189        let mut request = match build_chat_request(
190            self.request_model.as_deref().unwrap_or(&self.model),
191            context,
192            tool_schema_transform,
193        ) {
194            Ok(request) => request,
195            Err(error) => return error_stream(error),
196        };
197        request.prompt_cache_key = prompt_cache_key.resolve(context).map(String::from);
198        create_custom_stream_generic(&self.chat_client, request)
199    }
200
201    fn stream_responses(&self, context: &Context, policy: &ResponsesRequestPolicy) -> LlmResponseStream {
202        let mut url = match Url::parse(&self.openai_config.url("/responses")) {
203            Ok(url) => url,
204            Err(error) => return error_stream(LlmError::ProviderRequest(error.to_string())),
205        };
206
207        url.query_pairs_mut().extend_pairs(self.openai_config.query());
208
209        let mut request = match build_wire_request(&self.model, context, policy) {
210            Ok(request) => request,
211            Err(error) => return error_stream(error),
212        };
213        if let Some(model) = &self.request_model {
214            request["model"] = model.clone().into();
215        }
216
217        let http = self.http.clone();
218        let headers = self.openai_config.headers();
219        stream_from(async move { send(&http, url.as_str(), headers, request).await }, process_connection)
220    }
221}
222
223#[cfg(test)]
224mod tests {
225    use futures::StreamExt;
226    use serde_json::json;
227
228    use super::*;
229    use crate::providers::test_capture_server::{CaptureServer, ResponseSpec};
230    use crate::testing::FakeHttpService;
231    use crate::types::IsoString;
232    use crate::{
233        AssistantReasoning, ChatMessage, ContentBlock, LlmResponse, MessageId, ProviderErrorKind, ReasoningEffort,
234        ToolDefinition,
235    };
236
237    #[tokio::test]
238    async fn disabled_toggle_and_unknown_models_never_send_requests() {
239        let service = FakeHttpService::default();
240        let mut provider = GenericProvider::new("key".to_string(), &DEEPSEEK).unwrap();
241        provider.chat_client =
242            openai_client(AetherOpenAiConfig::new(OpenAIConfig::new(), ProviderAuthMode::None), service.clone());
243        for model in ["deepseek-v4-flash", "unknown"] {
244            provider = provider.with_model(model);
245            let mut context = Context::new(vec![], vec![]);
246            context.set_reasoning_effort(ReasoningEffort::Disabled);
247            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
248            assert_eq!(responses.len(), 1);
249            let error = responses[0].as_ref().unwrap_err();
250            if model == "unknown" {
251                assert!(matches!(error, LlmError::ReasoningValidation(_)));
252            } else {
253                assert!(matches!(error, LlmError::UnsupportedDisableTransport { .. }));
254            }
255            assert!(!error.is_retryable());
256            assert!(service.take_requests().is_empty());
257        }
258    }
259
260    #[test]
261    fn azure_foundry_requires_a_configured_url() {
262        let Err(error) = GenericProvider::new("key".to_string(), &AZURE_FOUNDRY) else {
263            panic!("Azure Foundry must require a URL");
264        };
265        assert!(matches!(error, LlmError::MissingProviderUrl { provider } if provider == "azure-foundry"));
266    }
267
268    #[tokio::test]
269    async fn chat_request_model_routes_the_request_without_changing_catalog_identity() {
270        let mut server = CaptureServer::start_chat_completions().await;
271        let provider = GenericProvider::new_with_connection(
272            "key".to_string(),
273            &AZURE_FOUNDRY,
274            ProviderConnectionConfig {
275                base_url: Some(format!("{}/", server.base_url)),
276                auth_mode: ProviderAuthMode::None,
277                request_model: Some("production-coding".to_string()),
278                ..Default::default()
279            },
280        )
281        .unwrap()
282        .with_model("gpt-5.5");
283        let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
284
285        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
286        let captured = server.captured().await;
287
288        assert_successful_stream(&responses);
289        assert_eq!(captured.path, "/chat/completions");
290        assert_eq!(captured.body["model"], "production-coding");
291        assert_eq!(captured.body["stream"], true);
292        assert_eq!(captured.body["stream_options"]["include_usage"], true);
293        assert!(captured.headers.get("authorization").is_none());
294        assert_eq!(provider.model().unwrap().to_string(), "azure-foundry:gpt-5.5");
295        assert_eq!(provider.display_name(), "Microsoft Foundry (gpt-5.5)");
296    }
297
298    #[tokio::test]
299    async fn chat_providers_apply_their_declared_prompt_cache_policy() {
300        for (config, expected_key) in [
301            (&AZURE_FOUNDRY, Some("prefix-abc")),
302            (&FIREWORKS, Some("conversation-abc")),
303            (&DEEPSEEK, None),
304            (&MOONSHOT, None),
305            (&ZAI, None),
306        ] {
307            let mut server = CaptureServer::start_chat_completions().await;
308            let provider = capture_backed_provider(&server, config);
309            let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
310            context.set_prompt_cache_key(Some("prefix-abc".to_string()));
311            context.set_session_affinity_key(Some("conversation-abc".to_string()));
312
313            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
314            let captured = server.captured().await;
315
316            assert_successful_stream(&responses);
317            assert_eq!(captured.body.get("prompt_cache_key").and_then(serde_json::Value::as_str), expected_key);
318            assert!(captured.body.get("user").is_none());
319            assert!(captured.body.get("session_id").is_none());
320        }
321    }
322
323    #[tokio::test]
324    async fn chat_providers_omit_unset_context_keys() {
325        for config in [&AZURE_FOUNDRY, &FIREWORKS] {
326            let mut server = CaptureServer::start_chat_completions().await;
327            let provider = capture_backed_provider(&server, config);
328            let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
329
330            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
331            let captured = server.captured().await;
332
333            assert_successful_stream(&responses);
334            assert!(captured.body.get("prompt_cache_key").is_none());
335            assert!(captured.body.get("session_id").is_none());
336        }
337    }
338
339    #[tokio::test]
340    async fn openai_distinguishes_default_disabled_and_low_effort() {
341        for (effort, expected) in [
342            (ReasoningEffort::Default, None),
343            (ReasoningEffort::Disabled, Some("none")),
344            (ReasoningEffort::Low, Some("low")),
345        ] {
346            let mut server = CaptureServer::start_responses().await;
347            let provider = capture_backed_provider(&server, &OPENAI).with_model("gpt-5.4");
348            let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
349            context.set_reasoning_effort(effort);
350
351            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
352            let body = server.captured().await.body;
353
354            assert!(responses.iter().all(Result::is_ok), "{responses:?}");
355            assert_eq!(body["reasoning"]["effort"].as_str(), expected);
356            if effort == ReasoningEffort::Disabled {
357                assert!(body["reasoning"]["summary"].is_null());
358            }
359        }
360    }
361
362    #[tokio::test]
363    async fn openai_sends_max_effort_and_prompt_cache_key() {
364        let mut server = CaptureServer::start_responses().await;
365        let provider = capture_backed_provider(&server, &OPENAI).with_model("gpt-5.6");
366        let mut context = Context::new(vec![ChatMessage::user("Think harder")], vec![]);
367        context.set_reasoning_effort(ReasoningEffort::Max);
368        context.set_prompt_cache_key(Some("cache-key".to_string()));
369
370        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
371        let captured = server.captured().await;
372
373        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
374        assert_eq!(captured.body["reasoning"]["effort"], "max");
375        assert_eq!(captured.body["model"], "gpt-5.6");
376        assert_eq!(captured.body["prompt_cache_key"], "cache-key");
377        assert_eq!(captured.body["include"], json!(["reasoning.encrypted_content"]));
378        assert_eq!(captured.body["stream"], true);
379        assert_eq!(provider.display_name(), "OpenAI (gpt-5.6)");
380    }
381
382    #[tokio::test]
383    async fn responses_http_200_failed_server_error_is_retryable_with_request_id() {
384        let spec = ResponseSpec::sse(include_str!("../../tests/fixtures/openai_responses/04_failed_server.sse"))
385            .with_header("x-request-id", "req-openai-1");
386        let mut server = CaptureServer::start_with_spec(spec).await;
387        let provider = capture_backed_provider(&server, &OPENAI);
388        let context = Context::new(vec![ChatMessage::user("hi")], vec![]);
389
390        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
391        let _ = server.captured().await;
392
393        assert!(!responses.iter().any(|r| matches!(r, Ok(LlmResponse::Done { .. }))));
394        let err = responses.iter().find_map(|r| r.as_ref().err()).expect("expected a failure");
395        assert!(err.is_retryable(), "server_error must be retryable: {err:?}");
396        let provider_error = err.provider().expect("expected provider error");
397        assert_eq!(provider_error.kind, ProviderErrorKind::Server);
398        assert_eq!(provider_error.http_status, Some(200));
399        assert_eq!(provider_error.request_id.as_deref(), Some("req-openai-1"));
400        assert_eq!(provider_error.code.as_deref(), Some("server_error"));
401    }
402
403    #[tokio::test]
404    async fn responses_surface_a_mapping_failure_as_the_only_item() {
405        let connection = ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() };
406        let provider = GenericProvider::from_env_with_connection(&OPENAI, connection).unwrap();
407        let context = Context::new(
408            vec![ChatMessage::User {
409                message_id: MessageId::new(),
410                content: vec![ContentBlock::Audio { data: "YXVkaW8=".to_string(), mime_type: "audio/wav".to_string() }],
411                timestamp: IsoString::now(),
412            }],
413            vec![],
414        );
415
416        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
417
418        assert_eq!(responses.len(), 1);
419        assert!(matches!(responses[0], Err(LlmError::UnsupportedContent(_))), "{responses:?}");
420    }
421
422    #[tokio::test]
423    async fn xiaomi_omits_encrypted_reasoning_summaries_and_prompt_cache_key() {
424        let mut server = CaptureServer::start_responses().await;
425        let provider = capture_backed_provider(&server, &XIAOMI);
426        let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
427        context.set_reasoning_effort(ReasoningEffort::High);
428        context.set_prompt_cache_key(Some("cache-key".to_string()));
429
430        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
431        let captured = server.captured().await;
432
433        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
434        assert_eq!(captured.path, "/responses");
435        assert_eq!(captured.body["model"], "mimo-v2.6-pro");
436        assert_eq!(captured.body["reasoning"], json!({ "effort": "high" }));
437        assert!(captured.body.get("include").is_none());
438        assert!(captured.body.get("prompt_cache_key").is_none());
439        assert_eq!(provider.display_name(), "Xiaomi (mimo-v2.6-pro)");
440    }
441
442    #[tokio::test]
443    async fn xiaomi_replays_prior_reasoning_as_plain_text() {
444        let mut server = CaptureServer::start_responses().await;
445        let provider = capture_backed_provider(&server, &XIAOMI);
446        let context = Context::new(
447            vec![
448                ChatMessage::user("Hello"),
449                ChatMessage::Assistant {
450                    message_id: MessageId::new(),
451                    content: "Hi".to_string(),
452                    reasoning: AssistantReasoning::from_parts("greeting the user".to_string(), None),
453                    timestamp: IsoString::now(),
454                    tool_calls: vec![],
455                },
456                ChatMessage::user("Again"),
457            ],
458            vec![],
459        );
460
461        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
462        let captured = server.captured().await;
463
464        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
465        let reasoning = captured.body["input"]
466            .as_array()
467            .unwrap()
468            .iter()
469            .find(|item| item["type"] == "reasoning")
470            .expect("reasoning item should be replayed");
471        assert_eq!(reasoning["content"], json!([{ "type": "reasoning_text", "text": "greeting the user" }]));
472        assert!(reasoning.get("encrypted_content").is_none_or(serde_json::Value::is_null));
473    }
474
475    #[tokio::test]
476    async fn xiaomi_drops_null_from_optional_tool_parameters() {
477        let mut server = CaptureServer::start_responses().await;
478        let provider = capture_backed_provider(&server, &XIAOMI);
479        let tool = ToolDefinition::new(
480            "bash",
481            "Run a command",
482            json!({
483                "type": "object",
484                "properties": {
485                    "command": { "type": "string" },
486                    "description": { "type": ["string", "null"] }
487                },
488                "required": ["command"]
489            }),
490        );
491        let context = Context::new(vec![ChatMessage::user("Hello")], vec![tool]);
492
493        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
494        let captured = server.captured().await;
495
496        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
497        assert_eq!(captured.body["tools"][0]["parameters"]["properties"]["description"], json!({ "type": "string" }));
498    }
499
500    #[tokio::test]
501    async fn responses_request_model_routes_the_request_without_changing_catalog_identity() {
502        let mut server = CaptureServer::start_responses().await;
503        let provider = GenericProvider::new_with_connection(
504            "key".to_string(),
505            &XIAOMI,
506            ProviderConnectionConfig {
507                base_url: Some(server.base_url.clone()),
508                auth_mode: ProviderAuthMode::None,
509                request_model: Some("mimo-deployment".to_string()),
510                ..Default::default()
511            },
512        )
513        .unwrap();
514        let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
515
516        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
517        let captured = server.captured().await;
518
519        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
520        assert_eq!(captured.body["model"], "mimo-deployment");
521        assert_eq!(provider.model().unwrap().to_string(), "xiaomi:mimo-v2.6-pro");
522    }
523
524    fn assert_successful_stream(responses: &[Result<LlmResponse>]) {
525        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
526        assert!(responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))));
527    }
528
529    fn capture_backed_provider(server: &CaptureServer, config: &'static ProviderConfig) -> GenericProvider {
530        GenericProvider::new_with_connection(
531            "key".to_string(),
532            config,
533            ProviderConnectionConfig {
534                base_url: Some(server.base_url.clone()),
535                auth_mode: ProviderAuthMode::None,
536                ..Default::default()
537            },
538        )
539        .unwrap()
540    }
541}