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
92pub 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}