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