Skip to main content

weft_core/defaults/transforms/
openai.rs

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    // Some OpenAI-compatible upstreams (Anthropic via OpenAI-compat shim, xAI)
35    // expose Anthropic-style cache fields too. Capture them so per-provider
36    // cache attribution survives the normalize step.
37    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
56/// Transforms for OpenAI-compatible APIs (OpenRouter, OpenAI, etc.)
57/// These APIs accept the same JSON format we use internally,
58/// so the transform is mostly pass-through + auth header.
59pub 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                // Responses API only accepts string content. Structured content
153                // (e.g. Anthropic-style `cache_control` blocks) gets flattened
154                // to plain text here — cache_control is meaningless on this
155                // endpoint anyway.
156                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        // SSE format: "data: {...}" or "data: [DONE]"
281        if line.is_empty() || line.starts_with(':') {
282            return Ok(None); // comment or keep-alive
283        }
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}