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