Skip to main content

llm/providers/
generic.rs

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