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