Skip to main content

gateway_core/
openai.rs

1use serde_json::{Value, json};
2
3use crate::{
4    Capabilities, ModelUsage, ProviderAdapter, ProviderError, ProviderRequest, ProviderResponse,
5    ProviderStreamDecoder, ProviderStreamEvent, SseEvent, Surface, provider::chat_usage,
6};
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum OpenAiFlavor {
10    OpenAi,
11    Foundry,
12    Compatible,
13}
14
15pub struct OpenAiCompatibleAdapter {
16    flavor: OpenAiFlavor,
17}
18
19impl OpenAiCompatibleAdapter {
20    pub fn new(flavor: OpenAiFlavor) -> Self {
21        Self { flavor }
22    }
23
24    pub fn openai() -> Self {
25        Self::new(OpenAiFlavor::OpenAi)
26    }
27
28    pub fn foundry() -> Self {
29        Self::new(OpenAiFlavor::Foundry)
30    }
31}
32
33pub fn normalize_foundry_endpoint(endpoint: &str) -> String {
34    let endpoint = endpoint.trim().trim_end_matches('/');
35    let has_path = endpoint
36        .split_once("://")
37        .is_some_and(|(_, authority)| authority.contains('/'));
38    if has_path {
39        endpoint.to_owned()
40    } else {
41        format!("{endpoint}/openai/v1")
42    }
43}
44
45/// Usage from an OpenAI-compatible embeddings response. Embeddings generate no
46/// completion, so only the prompt is billed: whatever the provider reports as
47/// output is ignored rather than priced.
48pub fn embeddings_usage(response: &Value) -> ModelUsage {
49    ModelUsage {
50        output_tokens: 0,
51        ..chat_usage(response)
52    }
53}
54
55pub fn responses_usage(response: &Value) -> ModelUsage {
56    response.get("usage").map(chat_usage).unwrap_or_default()
57}
58
59impl ProviderAdapter for OpenAiCompatibleAdapter {
60    fn name(&self) -> &'static str {
61        match self.flavor {
62            OpenAiFlavor::OpenAi => "openai",
63            OpenAiFlavor::Foundry => "foundry",
64            OpenAiFlavor::Compatible => "openai_compatible",
65        }
66    }
67
68    fn capabilities(&self) -> Capabilities {
69        Capabilities {
70            chat: true,
71            responses: true,
72            vision: true,
73            reasoning: true,
74            embeddings: false,
75        }
76    }
77
78    fn encode_request(
79        &self,
80        surface: Surface,
81        request: ProviderRequest,
82    ) -> Result<Value, ProviderError> {
83        let mut body = request.body;
84        let object = body.as_object_mut().ok_or_else(|| {
85            ProviderError::InvalidRequest("request body must be an object".into())
86        })?;
87        object.insert("model".into(), json!(request.model));
88        if surface == Surface::ChatCompletions && object.get("stream") == Some(&Value::Bool(true)) {
89            if let Some(options) = object
90                .get_mut("stream_options")
91                .and_then(Value::as_object_mut)
92            {
93                options.insert("include_usage".into(), Value::Bool(true));
94            } else {
95                object.insert("stream_options".into(), json!({ "include_usage": true }));
96            }
97        }
98        Ok(body)
99    }
100
101    fn decode_response(
102        &self,
103        surface: Surface,
104        response: Value,
105    ) -> Result<ProviderResponse, ProviderError> {
106        let usage = match surface {
107            Surface::ChatCompletions => chat_usage(&response),
108            Surface::Responses => responses_usage(&response),
109        };
110        Ok(ProviderResponse {
111            body: response,
112            usage,
113        })
114    }
115
116    fn stream_decoder(
117        &self,
118        surface: Surface,
119    ) -> Result<Box<dyn ProviderStreamDecoder>, ProviderError> {
120        Ok(Box::new(OpenAiStreamDecoder {
121            surface,
122            usage: ModelUsage::default(),
123            done: false,
124        }))
125    }
126}
127
128struct OpenAiStreamDecoder {
129    surface: Surface,
130    usage: ModelUsage,
131    done: bool,
132}
133
134impl ProviderStreamDecoder for OpenAiStreamDecoder {
135    fn decode(&mut self, event: SseEvent) -> Result<Vec<ProviderStreamEvent>, ProviderError> {
136        if event.data.trim() == "[DONE]" {
137            self.done = true;
138            return Ok(vec![ProviderStreamEvent::Done(self.usage)]);
139        }
140        let data: Value = serde_json::from_str(&event.data)
141            .map_err(|error| ProviderError::InvalidStream(error.to_string()))?;
142        if crate::is_rate_limit_payload(&data) {
143            let message = data
144                .pointer("/error/message")
145                .and_then(Value::as_str)
146                .unwrap_or("OpenAI stream rate limited")
147                .to_owned();
148            return Err(ProviderError::RateLimitedStream(message));
149        }
150        let usage = match self.surface {
151            Surface::ChatCompletions => data.get("usage"),
152            Surface::Responses => data
153                .pointer("/response/usage")
154                .or_else(|| data.get("usage")),
155        };
156        if let Some(usage) = usage.filter(|usage| usage.is_object()) {
157            self.usage = chat_usage(usage);
158        }
159        let event_name = match self.surface {
160            Surface::ChatCompletions => event.event,
161            Surface::Responses => event
162                .event
163                .or_else(|| data.get("type").and_then(Value::as_str).map(str::to_owned)),
164        };
165        Ok(vec![ProviderStreamEvent::Data {
166            event: event_name,
167            data,
168        }])
169    }
170
171    fn finish(&mut self) -> Result<Vec<ProviderStreamEvent>, ProviderError> {
172        if self.done {
173            Ok(Vec::new())
174        } else {
175            self.done = true;
176            Ok(vec![ProviderStreamEvent::Done(self.usage)])
177        }
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184
185    #[test]
186    fn foundry_endpoint_normalization_preserves_explicit_paths() {
187        assert_eq!(
188            normalize_foundry_endpoint("https://example.openai.azure.com/"),
189            "https://example.openai.azure.com/openai/v1"
190        );
191        assert_eq!(
192            normalize_foundry_endpoint("https://example.test/custom/v1/"),
193            "https://example.test/custom/v1"
194        );
195    }
196
197    #[test]
198    fn rewrites_model_and_forces_stream_usage() {
199        let body = OpenAiCompatibleAdapter::foundry()
200            .encode_request(
201                Surface::ChatCompletions,
202                ProviderRequest {
203                    model: "deployment".into(),
204                    body: json!({ "model": "foundry/deployment", "stream": true }),
205                },
206            )
207            .unwrap();
208        assert_eq!(body["model"], "deployment");
209        assert_eq!(body["stream_options"]["include_usage"], true);
210    }
211
212    #[test]
213    fn embeddings_usage_is_prompt_only() {
214        assert_eq!(
215            embeddings_usage(&json!({
216                "object": "list",
217                "data": [{ "embedding": [0.1, 0.2] }],
218                "usage": { "prompt_tokens": 8, "total_tokens": 8, "completion_tokens": 3 }
219            })),
220            ModelUsage {
221                input_tokens: 8,
222                ..ModelUsage::default()
223            }
224        );
225    }
226
227    #[test]
228    fn responses_usage_reads_the_responses_usage_block() {
229        assert_eq!(
230            responses_usage(&json!({
231                "usage": {
232                    "input_tokens": 20,
233                    "output_tokens": 8,
234                    "output_tokens_details": { "reasoning_tokens": 6 }
235                }
236            })),
237            ModelUsage {
238                input_tokens: 20,
239                output_tokens: 8,
240                reasoning_tokens: 6,
241                cache_read_tokens: 0,
242                cache_write_tokens: 0,
243            }
244        );
245    }
246
247    #[test]
248    fn chat_and_responses_preserve_unknown_fields_verbatim() {
249        for surface in [Surface::ChatCompletions, Surface::Responses] {
250            let original = json!({
251                "model": "qualified/model",
252                "stream": false,
253                "future_field": { "nested": [1, 2, 3] },
254                "tools": [{ "future_tool_field": true }],
255                "reasoning": { "effort": "high" }
256            });
257            let encoded = OpenAiCompatibleAdapter::openai()
258                .encode_request(
259                    surface,
260                    ProviderRequest {
261                        model: "bare-model".into(),
262                        body: original.clone(),
263                    },
264                )
265                .unwrap();
266            let mut expected = original;
267            expected["model"] = json!("bare-model");
268            assert_eq!(encoded, expected);
269
270            let response = json!({
271                "id": "response_1",
272                "future_response_field": { "opaque": true },
273                "usage": { "input_tokens": 3, "output_tokens": 4 }
274            });
275            assert_eq!(
276                OpenAiCompatibleAdapter::openai()
277                    .decode_response(surface, response.clone())
278                    .unwrap()
279                    .body,
280                response
281            );
282        }
283    }
284
285    #[test]
286    fn stream_usage_rewrite_preserves_other_stream_options() {
287        let body = OpenAiCompatibleAdapter::openai()
288            .encode_request(
289                Surface::ChatCompletions,
290                ProviderRequest {
291                    model: "model".into(),
292                    body: json!({
293                        "stream": true,
294                        "stream_options": { "future_option": "keep", "include_usage": false }
295                    }),
296                },
297            )
298            .unwrap();
299        assert_eq!(body["stream_options"]["future_option"], "keep");
300        assert_eq!(body["stream_options"]["include_usage"], true);
301    }
302
303    #[test]
304    fn stream_decoder_forwards_verbatim_and_finishes_with_usage() {
305        let mut decoder = OpenAiCompatibleAdapter::openai()
306            .stream_decoder(Surface::ChatCompletions)
307            .unwrap();
308        let chunk = json!({
309            "id": "chunk_1",
310            "choices": [],
311            "opaque": { "keep": true },
312            "usage": {
313                "prompt_tokens": 10,
314                "completion_tokens": 5,
315                "completion_tokens_details": { "reasoning_tokens": 2 },
316                "prompt_tokens_details": { "cached_tokens": 3 }
317            }
318        });
319        let forwarded = decoder
320            .decode(SseEvent {
321                event: Some("custom".into()),
322                data: chunk.to_string(),
323            })
324            .unwrap();
325        assert_eq!(
326            forwarded,
327            vec![ProviderStreamEvent::Data {
328                event: Some("custom".into()),
329                data: chunk
330            }]
331        );
332        assert_eq!(
333            decoder
334                .decode(SseEvent {
335                    event: None,
336                    data: "[DONE]".into(),
337                })
338                .unwrap(),
339            vec![ProviderStreamEvent::Done(ModelUsage {
340                input_tokens: 7,
341                output_tokens: 5,
342                reasoning_tokens: 2,
343                cache_read_tokens: 3,
344                cache_write_tokens: 0,
345            })]
346        );
347
348        let mut responses = OpenAiCompatibleAdapter::openai()
349            .stream_decoder(Surface::Responses)
350            .unwrap();
351        responses
352            .decode(SseEvent {
353                event: None,
354                data: json!({
355                    "type": "response.completed",
356                    "response": { "usage": {
357                        "input_tokens": 20,
358                        "output_tokens": 8,
359                        "output_tokens_details": { "reasoning_tokens": 6 },
360                        "input_tokens_details": { "cached_tokens": 4 }
361                    }}
362                })
363                .to_string(),
364            })
365            .unwrap();
366        assert_eq!(
367            responses.finish().unwrap(),
368            vec![ProviderStreamEvent::Done(ModelUsage {
369                input_tokens: 16,
370                output_tokens: 8,
371                reasoning_tokens: 6,
372                cache_read_tokens: 4,
373                cache_write_tokens: 0,
374            })]
375        );
376    }
377
378    #[test]
379    fn informational_rate_limits_updated_event_is_not_a_stream_error() {
380        let mut decoder = OpenAiCompatibleAdapter::openai()
381            .stream_decoder(Surface::Responses)
382            .unwrap();
383        let events = decoder
384            .decode(SseEvent {
385                event: None,
386                data: json!({
387                    "type": "rate_limits.updated",
388                    "rate_limits": { "requests": 10 }
389                })
390                .to_string(),
391            })
392            .unwrap();
393        assert!(matches!(
394            events.as_slice(),
395            [ProviderStreamEvent::Data { .. }]
396        ));
397    }
398
399    #[test]
400    fn rate_limit_stream_error_uses_provider_message() {
401        let mut decoder = OpenAiCompatibleAdapter::openai()
402            .stream_decoder(Surface::ChatCompletions)
403            .unwrap();
404        let error = decoder
405            .decode(SseEvent {
406                event: None,
407                data: json!({
408                    "error": {
409                        "type": "rate_limit_exceeded",
410                        "message": "slow down"
411                    }
412                })
413                .to_string(),
414            })
415            .unwrap_err();
416        assert_eq!(
417            error,
418            ProviderError::RateLimitedStream("slow down".to_owned())
419        );
420    }
421}