llm/providers/openai_compatible/
generic.rs1use async_openai::{Client, config::OpenAIConfig};
2use schemars::Schema;
3
4use crate::catalog::Provider;
5use crate::provider::{error_stream, get_context_window};
6use crate::providers::http::openai_client;
7use crate::tool_schema::normalize_for_moonshot;
8use crate::{
9 Context, LlmError, LlmModel, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, Result,
10 StreamingModelProvider,
11};
12
13use super::{AetherOpenAiConfig, PromptCacheKeySource, build_chat_request, create_custom_stream_generic};
14
15pub struct ProviderConfig {
20 pub provider: Provider,
21 pub api_base: Option<&'static str>,
22 pub default_model: &'static str,
23 pub tool_schema_transform: Option<fn(&mut Schema)>,
24 pub prompt_cache_key: PromptCacheKeySource,
25}
26
27pub const DEEPSEEK: ProviderConfig = ProviderConfig {
28 provider: Provider::DeepSeek,
29 api_base: Some("https://api.deepseek.com"),
30 default_model: "deepseek-v4-flash",
31 tool_schema_transform: None,
32 prompt_cache_key: PromptCacheKeySource::Omit,
33};
34
35pub const MOONSHOT: ProviderConfig = ProviderConfig {
36 provider: Provider::Moonshot,
37 api_base: Some("https://api.moonshot.ai/v1"),
38 default_model: "moonshot-v1-8k",
39 tool_schema_transform: Some(normalize_for_moonshot),
40 prompt_cache_key: PromptCacheKeySource::Omit,
41};
42
43pub const ZAI: ProviderConfig = ProviderConfig {
44 provider: Provider::ZAi,
45 api_base: Some("https://api.z.ai/api/coding/paas/v4"),
46 default_model: "GLM-4.6",
47 tool_schema_transform: None,
48 prompt_cache_key: PromptCacheKeySource::Omit,
49};
50
51pub const AZURE_FOUNDRY: ProviderConfig = ProviderConfig {
52 provider: Provider::AzureFoundry,
53 api_base: None,
54 default_model: "gpt-5.5",
55 tool_schema_transform: None,
56 prompt_cache_key: PromptCacheKeySource::Prefix,
57};
58
59pub const FIREWORKS: ProviderConfig = ProviderConfig {
60 provider: Provider::Fireworks,
61 api_base: Some("https://api.fireworks.ai/inference/v1"),
62 default_model: "accounts/fireworks/models/glm-5p1",
63 tool_schema_transform: None,
64 prompt_cache_key: PromptCacheKeySource::SessionAffinity,
65};
66
67pub(crate) const BUILT_INS: &[&ProviderConfig] = &[&DEEPSEEK, &MOONSHOT, &ZAI, &AZURE_FOUNDRY, &FIREWORKS];
68
69pub struct GenericOpenAiProvider {
71 client: Client<AetherOpenAiConfig>,
72 model: String,
73 request_model: Option<String>,
74 config: &'static ProviderConfig,
75}
76
77impl GenericOpenAiProvider {
78 pub fn from_env(config: &'static ProviderConfig) -> Result<Self> {
79 Self::from_env_with_connection(config, ProviderConnectionConfig::default())
80 }
81
82 pub fn from_env_with_connection(
83 config: &'static ProviderConfig,
84 connection: ProviderConnectionConfig,
85 ) -> Result<Self> {
86 let api_key = match connection.auth_mode {
87 ProviderAuthMode::Default => {
88 let env_var = config.provider.required_env_var().expect("generic providers require an API key");
89 std::env::var(env_var).map_err(|_| LlmError::MissingApiKey(env_var.to_string()))?
90 }
91 ProviderAuthMode::None => String::new(),
92 };
93 Self::new_with_connection(api_key, config, connection)
94 }
95
96 pub fn new(api_key: String, config: &'static ProviderConfig) -> Result<Self> {
97 Self::new_with_connection(api_key, config, ProviderConnectionConfig::default())
98 }
99
100 pub fn new_with_connection(
101 api_key: String,
102 config: &'static ProviderConfig,
103 connection: ProviderConnectionConfig,
104 ) -> Result<Self> {
105 let api_base = connection
106 .base_url
107 .or_else(|| config.api_base.map(str::to_string))
108 .ok_or_else(|| LlmError::MissingProviderUrl { provider: config.provider.parser_name().to_string() })?
109 .trim_end_matches('/')
110 .to_string();
111 let openai_config = OpenAIConfig::new().with_api_key(api_key).with_api_base(api_base);
112 let openai_config = AetherOpenAiConfig::new(openai_config, connection.auth_mode);
113
114 Ok(Self {
115 client: openai_client(openai_config, reqwest::Client::new()),
116 model: config.default_model.to_string(),
117 request_model: connection.request_model,
118 config,
119 })
120 }
121
122 pub fn with_model(mut self, model: &str) -> Self {
123 self.model = model.to_string();
124 self
125 }
126}
127
128impl StreamingModelProvider for GenericOpenAiProvider {
129 fn model(&self) -> Option<LlmModel> {
130 format!("{}:{}", self.config.provider.parser_name(), self.model).parse().ok()
131 }
132
133 fn context_window(&self) -> Option<u32> {
134 get_context_window(self.config.provider.parser_name(), &self.model)
135 }
136
137 fn stream_response(&self, context: &Context) -> LlmResponseStream {
138 let mut request = match build_chat_request(
139 self.request_model.as_deref().unwrap_or(&self.model),
140 context,
141 self.config.tool_schema_transform,
142 ) {
143 Ok(req) => req,
144 Err(e) => return error_stream(e),
145 };
146 request.prompt_cache_key = self.config.prompt_cache_key.resolve(context).map(String::from);
147 create_custom_stream_generic(&self.client, request)
148 }
149
150 fn display_name(&self) -> String {
151 format!("{} ({})", self.config.provider.display_name(), self.model)
152 }
153}
154
155#[cfg(test)]
156mod tests {
157 use futures::StreamExt;
158
159 use super::*;
160 use crate::providers::test_capture_server::CaptureServer;
161
162 use crate::ChatMessage;
163
164 #[test]
165 fn azure_foundry_requires_a_configured_url() {
166 let Err(error) = GenericOpenAiProvider::new("key".to_string(), &AZURE_FOUNDRY) else {
167 panic!("Azure Foundry must require a URL");
168 };
169 assert!(matches!(error, LlmError::MissingProviderUrl { provider } if provider == "azure-foundry"));
170 }
171
172 #[tokio::test]
173 async fn request_model_routes_the_request_without_changing_catalog_identity() {
174 let mut server = CaptureServer::start_chat_completions().await;
175 let provider = GenericOpenAiProvider::new_with_connection(
176 "key".to_string(),
177 &AZURE_FOUNDRY,
178 ProviderConnectionConfig {
179 base_url: Some(format!("{}/", server.base_url)),
180 auth_mode: ProviderAuthMode::None,
181 request_model: Some("production-coding".to_string()),
182 ..Default::default()
183 },
184 )
185 .unwrap()
186 .with_model("gpt-5.5");
187 let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
188
189 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
190 let captured = server.captured().await;
191
192 assert_successful_stream(&responses);
193 assert_eq!(captured.path, "/chat/completions");
194 assert_eq!(captured.body["model"], "production-coding");
195 assert_eq!(captured.body["stream"], true);
196 assert_eq!(captured.body["stream_options"]["include_usage"], true);
197 assert!(captured.headers.get("authorization").is_none());
198 assert_eq!(provider.model().unwrap().to_string(), "azure-foundry:gpt-5.5");
199 assert_eq!(provider.display_name(), "Microsoft Foundry (gpt-5.5)");
200 }
201
202 #[tokio::test]
203 async fn providers_apply_their_declared_prompt_cache_policy() {
204 for (config, expected_key) in [
205 (&AZURE_FOUNDRY, Some("prefix-abc")),
206 (&FIREWORKS, Some("conversation-abc")),
207 (&DEEPSEEK, None),
208 (&MOONSHOT, None),
209 (&ZAI, None),
210 ] {
211 let mut server = CaptureServer::start_chat_completions().await;
212 let provider = capture_backed_provider(&server, config);
213 let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
214 context.set_prompt_cache_key(Some("prefix-abc".to_string()));
215 context.set_session_affinity_key(Some("conversation-abc".to_string()));
216
217 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
218 let captured = server.captured().await;
219
220 assert_successful_stream(&responses);
221 assert_eq!(captured.body.get("prompt_cache_key").and_then(serde_json::Value::as_str), expected_key);
222 assert!(captured.body.get("user").is_none());
223 assert!(captured.body.get("session_id").is_none());
224 }
225 }
226
227 #[tokio::test]
228 async fn providers_omit_unset_context_keys() {
229 for config in [&AZURE_FOUNDRY, &FIREWORKS] {
230 let mut server = CaptureServer::start_chat_completions().await;
231 let provider = capture_backed_provider(&server, config);
232 let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
233
234 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
235 let captured = server.captured().await;
236
237 assert_successful_stream(&responses);
238 assert!(captured.body.get("prompt_cache_key").is_none());
239 assert!(captured.body.get("session_id").is_none());
240 }
241 }
242
243 fn assert_successful_stream(responses: &[Result<crate::LlmResponse>]) {
244 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
245 assert!(responses.iter().any(|response| matches!(response, Ok(crate::LlmResponse::Done { .. }))));
246 }
247
248 fn capture_backed_provider(server: &CaptureServer, config: &'static ProviderConfig) -> GenericOpenAiProvider {
249 GenericOpenAiProvider::new_with_connection(
250 "key".to_string(),
251 config,
252 ProviderConnectionConfig {
253 base_url: Some(server.base_url.clone()),
254 auth_mode: ProviderAuthMode::None,
255 ..Default::default()
256 },
257 )
258 .unwrap()
259 }
260}