Skip to main content

llm/providers/openai_compatible/
generic.rs

1use 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
15/// Configuration for an OpenAI-compatible provider.
16///
17/// Each provider that uses the standard `build_chat_request → create_custom_stream_generic`
18/// flow differs only in these constants.
19pub 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
69/// A generic provider for APIs that are fully OpenAI-compatible.
70pub 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        if let Err(error) = crate::provider::validate_reasoning(context, self.model().as_ref()) {
139            return crate::provider::error_stream(error);
140        }
141        let mut request = match build_chat_request(
142            self.request_model.as_deref().unwrap_or(&self.model),
143            context,
144            self.config.tool_schema_transform,
145        ) {
146            Ok(req) => req,
147            Err(e) => return error_stream(e),
148        };
149        request.prompt_cache_key = self.config.prompt_cache_key.resolve(context).map(String::from);
150        create_custom_stream_generic(&self.client, request)
151    }
152
153    fn display_name(&self) -> String {
154        format!("{} ({})", self.config.provider.display_name(), self.model)
155    }
156}
157
158#[cfg(test)]
159mod tests {
160    use futures::StreamExt;
161
162    use super::*;
163    use crate::providers::test_capture_server::CaptureServer;
164
165    use crate::ChatMessage;
166
167    #[tokio::test]
168    async fn disabled_toggle_and_unknown_models_never_send_requests() {
169        use crate::testing::FakeHttpService;
170        let service = FakeHttpService::default();
171        let mut provider = GenericOpenAiProvider::new("key".to_string(), &DEEPSEEK).unwrap();
172        provider.client =
173            openai_client(AetherOpenAiConfig::new(OpenAIConfig::new(), ProviderAuthMode::None), service.clone());
174        for model in ["deepseek-v4-flash", "unknown"] {
175            provider = provider.with_model(model);
176            let mut context = Context::new(vec![], vec![]);
177            context.set_reasoning_effort(crate::ReasoningEffort::Disabled);
178            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
179            assert_eq!(responses.len(), 1);
180            let error = responses[0].as_ref().unwrap_err();
181            if model == "unknown" {
182                assert!(matches!(error, LlmError::ReasoningValidation(_)));
183            } else {
184                assert!(matches!(error, LlmError::UnsupportedDisableTransport { .. }));
185            }
186            assert!(!error.is_retryable());
187            assert!(service.take_requests().is_empty());
188        }
189    }
190
191    #[test]
192    fn azure_foundry_requires_a_configured_url() {
193        let Err(error) = GenericOpenAiProvider::new("key".to_string(), &AZURE_FOUNDRY) else {
194            panic!("Azure Foundry must require a URL");
195        };
196        assert!(matches!(error, LlmError::MissingProviderUrl { provider } if provider == "azure-foundry"));
197    }
198
199    #[tokio::test]
200    async fn request_model_routes_the_request_without_changing_catalog_identity() {
201        let mut server = CaptureServer::start_chat_completions().await;
202        let provider = GenericOpenAiProvider::new_with_connection(
203            "key".to_string(),
204            &AZURE_FOUNDRY,
205            ProviderConnectionConfig {
206                base_url: Some(format!("{}/", server.base_url)),
207                auth_mode: ProviderAuthMode::None,
208                request_model: Some("production-coding".to_string()),
209                ..Default::default()
210            },
211        )
212        .unwrap()
213        .with_model("gpt-5.5");
214        let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
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.path, "/chat/completions");
221        assert_eq!(captured.body["model"], "production-coding");
222        assert_eq!(captured.body["stream"], true);
223        assert_eq!(captured.body["stream_options"]["include_usage"], true);
224        assert!(captured.headers.get("authorization").is_none());
225        assert_eq!(provider.model().unwrap().to_string(), "azure-foundry:gpt-5.5");
226        assert_eq!(provider.display_name(), "Microsoft Foundry (gpt-5.5)");
227    }
228
229    #[tokio::test]
230    async fn providers_apply_their_declared_prompt_cache_policy() {
231        for (config, expected_key) in [
232            (&AZURE_FOUNDRY, Some("prefix-abc")),
233            (&FIREWORKS, Some("conversation-abc")),
234            (&DEEPSEEK, None),
235            (&MOONSHOT, None),
236            (&ZAI, None),
237        ] {
238            let mut server = CaptureServer::start_chat_completions().await;
239            let provider = capture_backed_provider(&server, config);
240            let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
241            context.set_prompt_cache_key(Some("prefix-abc".to_string()));
242            context.set_session_affinity_key(Some("conversation-abc".to_string()));
243
244            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
245            let captured = server.captured().await;
246
247            assert_successful_stream(&responses);
248            assert_eq!(captured.body.get("prompt_cache_key").and_then(serde_json::Value::as_str), expected_key);
249            assert!(captured.body.get("user").is_none());
250            assert!(captured.body.get("session_id").is_none());
251        }
252    }
253
254    #[tokio::test]
255    async fn providers_omit_unset_context_keys() {
256        for config in [&AZURE_FOUNDRY, &FIREWORKS] {
257            let mut server = CaptureServer::start_chat_completions().await;
258            let provider = capture_backed_provider(&server, config);
259            let context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
260
261            let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
262            let captured = server.captured().await;
263
264            assert_successful_stream(&responses);
265            assert!(captured.body.get("prompt_cache_key").is_none());
266            assert!(captured.body.get("session_id").is_none());
267        }
268    }
269
270    fn assert_successful_stream(responses: &[Result<crate::LlmResponse>]) {
271        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
272        assert!(responses.iter().any(|response| matches!(response, Ok(crate::LlmResponse::Done { .. }))));
273    }
274
275    fn capture_backed_provider(server: &CaptureServer, config: &'static ProviderConfig) -> GenericOpenAiProvider {
276        GenericOpenAiProvider::new_with_connection(
277            "key".to_string(),
278            config,
279            ProviderConnectionConfig {
280                base_url: Some(server.base_url.clone()),
281                auth_mode: ProviderAuthMode::None,
282                ..Default::default()
283            },
284        )
285        .unwrap()
286    }
287}