1use crate::config::{ProviderApi, ProviderConfig};
2use crate::layers::transform::{ProviderRequest, TransformLayer};
3use crate::types::{ChatMessage, ChatRequest, ChatResponse, Choice, StreamChunk, Usage};
4use anyhow::{Context, Result};
5use async_trait::async_trait;
6use bytes::Bytes;
7use serde::{Deserialize, Serialize};
8
9fn normalize_openai_usage_fields(response: &mut serde_json::Value) {
10 let Some(usage) = response.get_mut("usage") else {
11 return;
12 };
13 let Some(usage_object) = usage.as_object_mut() else {
14 return;
15 };
16 let prompt_tokens = usage_object
17 .get("prompt_tokens")
18 .and_then(serde_json::Value::as_u64)
19 .unwrap_or(0);
20 let completion_tokens = usage_object
21 .get("completion_tokens")
22 .and_then(serde_json::Value::as_u64)
23 .unwrap_or(0);
24 let total_tokens = usage_object
25 .get("total_tokens")
26 .and_then(serde_json::Value::as_u64)
27 .unwrap_or(prompt_tokens.saturating_add(completion_tokens));
28 let prompt_cache_hit_tokens = usage_object
29 .get("prompt_cache_hit_tokens")
30 .and_then(serde_json::Value::as_u64);
31 let prompt_cache_miss_tokens = usage_object
32 .get("prompt_cache_miss_tokens")
33 .and_then(serde_json::Value::as_u64);
34 let cache_read_input_tokens = usage_object
38 .get("cache_read_input_tokens")
39 .and_then(serde_json::Value::as_u64);
40 let cache_creation_input_tokens = usage_object
41 .get("cache_creation_input_tokens")
42 .and_then(serde_json::Value::as_u64);
43
44 *usage = serde_json::to_value(Usage {
45 prompt_tokens,
46 completion_tokens,
47 total_tokens,
48 prompt_cache_hit_tokens,
49 prompt_cache_miss_tokens,
50 cache_read_input_tokens,
51 cache_creation_input_tokens,
52 })
53 .unwrap_or_else(|_| usage.clone());
54}
55
56pub struct OpenAITransform;
60
61fn normalize_openai_base_url(base_url: &str) -> String {
62 let trimmed = base_url.trim().trim_end_matches('/');
63 let lower = trimmed.to_lowercase();
64
65 if lower.ends_with("/chat/completions") || lower.ends_with("/responses") {
66 return trimmed.to_string();
67 }
68
69 if lower.ends_with("/v1") || lower.ends_with("/openai/v1") {
70 return trimmed.to_string();
71 }
72
73 if lower.contains("api.openai.com") || lower.contains("api.deepseek.com") {
74 return format!("{trimmed}/v1");
75 }
76
77 trimmed.to_string()
78}
79
80fn openai_endpoint(provider: &ProviderConfig) -> &'static str {
81 match provider.api {
82 ProviderApi::ChatCompletions => "chat/completions",
83 ProviderApi::Responses => "responses",
84 }
85}
86
87#[derive(Debug, Serialize)]
88struct ResponsesRequest {
89 model: String,
90 input: Vec<ResponsesInputMessage>,
91 #[serde(skip_serializing_if = "Option::is_none")]
92 temperature: Option<f64>,
93 #[serde(skip_serializing_if = "Option::is_none")]
94 max_output_tokens: Option<u64>,
95 #[serde(skip_serializing_if = "Option::is_none")]
96 top_p: Option<f64>,
97 #[serde(skip_serializing_if = "Option::is_none")]
98 tools: Option<Vec<serde_json::Value>>,
99 #[serde(skip_serializing_if = "Option::is_none")]
100 tool_choice: Option<serde_json::Value>,
101 stream: bool,
102}
103
104#[derive(Debug, Serialize)]
105struct ResponsesInputMessage {
106 role: String,
107 content: String,
108}
109
110#[derive(Debug, Deserialize)]
111struct ResponsesResponse {
112 id: String,
113 model: String,
114 #[serde(default)]
115 output_text: Option<String>,
116 #[serde(default)]
117 output: Vec<ResponsesOutputItem>,
118 #[serde(default)]
119 usage: Option<ResponsesUsage>,
120}
121
122#[derive(Debug, Deserialize)]
123struct ResponsesOutputItem {
124 #[serde(default)]
125 content: Vec<ResponsesContentItem>,
126}
127
128#[derive(Debug, Deserialize)]
129struct ResponsesContentItem {
130 #[serde(default)]
131 text: Option<String>,
132}
133
134#[derive(Debug, Deserialize)]
135struct ResponsesUsage {
136 #[serde(default)]
137 input_tokens: u64,
138 #[serde(default)]
139 output_tokens: u64,
140 #[serde(default)]
141 total_tokens: u64,
142}
143
144fn responses_body(request: &ChatRequest) -> ResponsesRequest {
145 ResponsesRequest {
146 model: request.model.clone(),
147 input: request
148 .messages
149 .iter()
150 .map(|message| ResponsesInputMessage {
151 role: message.role.clone(),
152 content: message.content_text(),
157 })
158 .collect(),
159 temperature: request.temperature,
160 max_output_tokens: request.max_tokens,
161 top_p: request.top_p,
162 tools: request.tools.clone(),
163 tool_choice: request.tool_choice.clone(),
164 stream: request.stream,
165 }
166}
167
168fn response_text(resp: &ResponsesResponse) -> String {
169 if let Some(text) = &resp.output_text {
170 if !text.is_empty() {
171 return text.clone();
172 }
173 }
174
175 resp.output
176 .iter()
177 .flat_map(|item| item.content.iter())
178 .filter_map(|content| content.text.as_deref())
179 .collect::<Vec<_>>()
180 .join("")
181}
182
183#[async_trait]
184impl TransformLayer for OpenAITransform {
185 async fn transform_request(
186 &self,
187 request: &ChatRequest,
188 api_key: &str,
189 provider: &ProviderConfig,
190 ) -> Result<ProviderRequest> {
191 let normalized_base_url = normalize_openai_base_url(&provider.base_url);
192 let url = format!(
193 "{}/{}",
194 normalized_base_url.trim_end_matches('/'),
195 openai_endpoint(provider)
196 );
197 let body = match provider.api {
198 ProviderApi::ChatCompletions => {
199 serde_json::to_vec(request).context("Failed to serialize request")?
200 }
201 ProviderApi::Responses => serde_json::to_vec(&responses_body(request))
202 .context("Failed to serialize Responses request")?,
203 };
204
205 Ok(ProviderRequest {
206 url,
207 method: "POST".into(),
208 headers: vec![
209 ("Content-Type".into(), "application/json".into()),
210 ("Authorization".into(), format!("Bearer {}", api_key)),
211 ],
212 body: Bytes::from(body),
213 })
214 }
215
216 async fn transform_response(
217 &self,
218 status: u16,
219 body: Bytes,
220 provider: &ProviderConfig,
221 ) -> Result<ChatResponse> {
222 if status != 200 {
223 let text = String::from_utf8_lossy(&body);
224 anyhow::bail!("Provider returned status {}: {}", status, text);
225 }
226 match provider.api {
227 ProviderApi::ChatCompletions => {
228 let mut response_value: serde_json::Value =
229 serde_json::from_slice(&body).context("Failed to parse provider response")?;
230 normalize_openai_usage_fields(&mut response_value);
231 let resp: ChatResponse = serde_json::from_value(response_value)
232 .context("Failed to normalize provider response")?;
233 Ok(resp)
234 }
235 ProviderApi::Responses => {
236 let resp: ResponsesResponse =
237 serde_json::from_slice(&body).context("Failed to parse Responses response")?;
238 let usage = resp.usage.as_ref().map(|usage| Usage {
239 prompt_tokens: usage.input_tokens,
240 completion_tokens: usage.output_tokens,
241 total_tokens: if usage.total_tokens == 0 {
242 usage.input_tokens + usage.output_tokens
243 } else {
244 usage.total_tokens
245 },
246 prompt_cache_hit_tokens: None,
247 prompt_cache_miss_tokens: None,
248 cache_read_input_tokens: None,
249 cache_creation_input_tokens: None,
250 });
251 let content = response_text(&resp);
252 Ok(ChatResponse {
253 id: resp.id,
254 object: "chat.completion".into(),
255 created: 0,
256 model: resp.model,
257 choices: vec![Choice {
258 index: 0,
259 message: ChatMessage {
260 role: "assistant".into(),
261 content: serde_json::Value::String(content),
262 tool_calls: None,
263 tool_call_id: None,
264 },
265 finish_reason: Some("stop".into()),
266 }],
267 usage,
268 })
269 }
270 }
271 }
272
273 async fn transform_stream_chunk(
274 &self,
275 chunk: &str,
276 _provider: &ProviderConfig,
277 ) -> Result<Option<StreamChunk>> {
278 let line = chunk.trim();
279
280 if line.is_empty() || line.starts_with(':') {
282 return Ok(None); }
284
285 let data = line.strip_prefix("data: ").unwrap_or(line);
286
287 if data == "[DONE]" {
288 return Ok(None);
289 }
290
291 let chunk: StreamChunk =
292 serde_json::from_str(data).context("Failed to parse stream chunk")?;
293 Ok(Some(chunk))
294 }
295}
296
297#[cfg(test)]
298mod tests {
299 use super::*;
300 use crate::config::{ApiKeyConfig, ProviderApi, ProviderConfig};
301 use crate::types::{ChatMessage, ChatRequest};
302
303 fn provider() -> ProviderConfig {
304 ProviderConfig {
305 name: "openrouter".into(),
306 base_url: "https://openrouter.ai/api/v1".into(),
307 format: "openai".into(),
308 api: ProviderApi::ChatCompletions,
309 keys: vec![ApiKeyConfig {
310 value: "sk-test".into(),
311 label: None,
312 enabled: true,
313 }],
314 models: vec!["gpt-4o".into()],
315 }
316 }
317
318 fn provider_with_base_url(base_url: &str) -> ProviderConfig {
319 ProviderConfig {
320 base_url: base_url.into(),
321 ..provider()
322 }
323 }
324
325 fn request() -> ChatRequest {
326 ChatRequest {
327 model: "gpt-4o".into(),
328 messages: vec![ChatMessage {
329 role: "user".into(),
330 content: "hello".into(),
331 tool_calls: None,
332 tool_call_id: None,
333 }],
334 stream: false,
335 temperature: None,
336 max_tokens: None,
337 top_p: None,
338 tools: None,
339 tool_choice: None,
340 response_format: None,
341 x_provider: None,
342 }
343 }
344
345 #[tokio::test]
346 async fn test_transform_request_url() {
347 let t = OpenAITransform;
348 let req = t
349 .transform_request(&request(), "sk-test", &provider())
350 .await
351 .unwrap();
352 assert_eq!(req.url, "https://openrouter.ai/api/v1/chat/completions");
353 assert!(req
354 .headers
355 .iter()
356 .any(|(k, v)| k == "Authorization" && v == "Bearer sk-test"));
357 }
358
359 #[tokio::test]
360 async fn test_transform_request_url_normalizes_deepseek_v1() {
361 let t = OpenAITransform;
362 let provider = provider_with_base_url("https://api.deepseek.com");
363 let req = t
364 .transform_request(&request(), "sk-test", &provider)
365 .await
366 .unwrap();
367
368 assert_eq!(req.url, "https://api.deepseek.com/v1/chat/completions");
369 }
370
371 #[tokio::test]
372 async fn test_transform_request_url_normalizes_deepseek_v1_with_trailing_slash() {
373 let t = OpenAITransform;
374 let provider = provider_with_base_url(" https://api.deepseek.com/ ");
375 let req = t
376 .transform_request(&request(), "sk-test", &provider)
377 .await
378 .unwrap();
379
380 assert_eq!(req.url, "https://api.deepseek.com/v1/chat/completions");
381 }
382
383 #[tokio::test]
384 async fn test_transform_request_url_preserves_deepseek_explicit_v1() {
385 let t = OpenAITransform;
386 let provider = provider_with_base_url("https://api.deepseek.com/v1");
387 let req = t
388 .transform_request(&request(), "sk-test", &provider)
389 .await
390 .unwrap();
391
392 assert_eq!(req.url, "https://api.deepseek.com/v1/chat/completions");
393 }
394
395 #[tokio::test]
396 async fn test_transform_request_url_preserves_explicit_v1() {
397 let t = OpenAITransform;
398 let provider = provider_with_base_url("https://api.openai.com/v1");
399 let req = t
400 .transform_request(&request(), "sk-test", &provider)
401 .await
402 .unwrap();
403
404 assert_eq!(req.url, "https://api.openai.com/v1/chat/completions");
405 }
406
407 #[tokio::test]
408 async fn test_transform_stream_done() {
409 let t = OpenAITransform;
410 let result = t
411 .transform_stream_chunk("data: [DONE]", &provider())
412 .await
413 .unwrap();
414 assert!(result.is_none());
415 }
416
417 #[tokio::test]
418 async fn test_transform_stream_keepalive() {
419 let t = OpenAITransform;
420 let result = t.transform_stream_chunk("", &provider()).await.unwrap();
421 assert!(result.is_none());
422 let result = t
423 .transform_stream_chunk(": ping", &provider())
424 .await
425 .unwrap();
426 assert!(result.is_none());
427 }
428
429 #[tokio::test]
430 async fn test_transform_response_preserves_prompt_cache_usage_fields() {
431 let t = OpenAITransform;
432 let body = Bytes::from_static(
433 br#"{
434 "id": "chatcmpl_1",
435 "object": "chat.completion",
436 "created": 1,
437 "model": "deepseek-v4-flash",
438 "choices": [{
439 "index": 0,
440 "message": {"role": "assistant", "content": "ok"},
441 "finish_reason": "stop"
442 }],
443 "usage": {
444 "prompt_tokens": 100,
445 "completion_tokens": 20,
446 "total_tokens": 120,
447 "prompt_cache_hit_tokens": 80,
448 "prompt_cache_miss_tokens": 20
449 }
450 }"#,
451 );
452
453 let response = t.transform_response(200, body, &provider()).await.unwrap();
454 let usage = response.usage.expect("usage should be present");
455
456 assert_eq!(usage.prompt_tokens, 100);
457 assert_eq!(usage.completion_tokens, 20);
458 assert_eq!(usage.total_tokens, 120);
459 assert_eq!(usage.prompt_cache_hit_tokens, Some(80));
460 assert_eq!(usage.prompt_cache_miss_tokens, Some(20));
461 }
462}