Skip to main content

claude_codex/providers/codex/chat_completions/
mod.rs

1pub mod request;
2pub mod response;
3pub mod stream;
4
5use std::{sync::Arc, time::Duration};
6
7use axum::{
8    Json,
9    response::{IntoResponse, Response},
10};
11use futures_util::StreamExt;
12use http::StatusCode;
13use serde_json::{Value, json};
14
15use crate::provider::RequestContext;
16
17use super::client::{CodexError, CodexHttpClient};
18use request::TranslatedRequest;
19
20pub struct ChatCompletionsBackend {
21    client: Arc<CodexHttpClient>,
22}
23
24impl Default for ChatCompletionsBackend {
25    fn default() -> Self {
26        Self::new()
27    }
28}
29
30impl ChatCompletionsBackend {
31    pub fn new() -> Self {
32        Self {
33            client: Arc::new(CodexHttpClient::new()),
34        }
35    }
36
37    #[cfg(test)]
38    fn with_client(client: CodexHttpClient) -> Self {
39        Self {
40            client: Arc::new(client),
41        }
42    }
43
44    pub async fn handle(&self, request: TranslatedRequest, ctx: RequestContext) -> Response {
45        if let Some(monitor) = ctx.monitor.as_ref() {
46            monitor.model_resolved(&ctx.req_id, &request.model);
47            monitor.upstream_started(&ctx.req_id);
48        }
49        let upstream = match self
50            .client
51            .post_native_responses(&request.upstream, &ctx, request.use_responses_lite, true)
52            .await
53        {
54            Ok(upstream) => upstream,
55            Err(error) => return codex_error_response(error),
56        };
57
58        if !upstream.status().is_success() {
59            return upstream_error_response(upstream, self.client.body_idle_timeout_ms()).await;
60        }
61        if request.stream {
62            return stream::streaming_response(
63                upstream,
64                ctx,
65                request.model,
66                request.include_usage,
67                self.client.body_idle_timeout_ms(),
68            );
69        }
70
71        let headers = stream::response_headers(upstream.headers());
72        let bytes =
73            match collect_body(upstream, self.client.body_idle_timeout_ms(), Some(&ctx)).await {
74                Ok(bytes) => bytes,
75                Err(error) => return error.response(),
76            };
77        if let Some(traffic) = ctx.traffic.as_deref() {
78            traffic.write_bytes("032-upstream-response-body.sse", &bytes);
79        }
80        let completion = match response::aggregate_sse(&bytes, &request.model) {
81            Ok(completion) => completion,
82            Err(error) => return error.response(),
83        };
84        if let Some(usage) = completion.get("usage")
85            && let Some(monitor) = ctx.monitor.as_ref()
86        {
87            monitor.usage_updated(
88                &ctx.req_id,
89                usage.get("prompt_tokens").and_then(Value::as_u64),
90                usage.get("completion_tokens").and_then(Value::as_u64),
91            );
92        }
93        if let Some(traffic) = ctx.traffic.as_deref() {
94            traffic.write_json("050-openai-chat-completion-response", &completion);
95        }
96        let mut downstream = Json(completion).into_response();
97        *downstream.headers_mut() = headers;
98        downstream.headers_mut().insert(
99            http::header::CONTENT_TYPE,
100            http::HeaderValue::from_static("application/json"),
101        );
102        downstream
103    }
104}
105
106async fn collect_body(
107    upstream: reqwest::Response,
108    idle_timeout_ms: u64,
109    ctx: Option<&RequestContext>,
110) -> Result<Vec<u8>, ChatError> {
111    let mut stream = upstream.bytes_stream();
112    let mut bytes = Vec::new();
113    let mut started = false;
114    loop {
115        match tokio::time::timeout(Duration::from_millis(idle_timeout_ms), stream.next()).await {
116            Ok(Some(Ok(chunk))) => {
117                if !started {
118                    if let Some(ctx) = ctx
119                        && let Some(monitor) = ctx.monitor.as_ref()
120                    {
121                        monitor.generation_started(&ctx.req_id);
122                    }
123                    started = true;
124                }
125                bytes.extend_from_slice(&chunk);
126                if let Some(ctx) = ctx
127                    && let Some(monitor) = ctx.monitor.as_ref()
128                {
129                    monitor.stream_progress(&ctx.req_id, chunk.len() as u64, 0, None, None);
130                }
131            }
132            Ok(Some(Err(error))) => {
133                return Err(ChatError::upstream(format!(
134                    "Codex response body read failed: {error}"
135                )));
136            }
137            Ok(None) => return Ok(bytes),
138            Err(_) => {
139                return Err(ChatError::timeout(format!(
140                    "Timed out waiting {idle_timeout_ms}ms for the next Codex response body chunk"
141                )));
142            }
143        }
144    }
145}
146
147fn codex_error_response(error: CodexError) -> Response {
148    let retry_after = error.retry_after.clone();
149    let response = ChatError::from_codex(error).response();
150    if let Some(retry_after) = retry_after
151        && let Ok(value) = http::HeaderValue::from_str(&retry_after)
152    {
153        let (mut parts, body) = response.into_parts();
154        parts.headers.insert(http::header::RETRY_AFTER, value);
155        Response::from_parts(parts, body)
156    } else {
157        response
158    }
159}
160
161async fn upstream_error_response(upstream: reqwest::Response, idle_timeout_ms: u64) -> Response {
162    let status = upstream.status();
163    let retry_after = upstream.headers().get(http::header::RETRY_AFTER).cloned();
164    let bytes = match collect_body(upstream, idle_timeout_ms, None).await {
165        Ok(bytes) => bytes,
166        Err(error) => return error.response(),
167    };
168    let message = serde_json::from_slice::<Value>(&bytes)
169        .ok()
170        .and_then(|value| {
171            value
172                .pointer("/error/message")
173                .or_else(|| value.get("message"))
174                .or_else(|| value.get("detail"))
175                .and_then(Value::as_str)
176                .map(str::to_string)
177        })
178        .filter(|message| !message.is_empty())
179        .unwrap_or_else(|| format!("Codex request failed with status {}", status.as_u16()));
180    let kind = match status {
181        StatusCode::UNAUTHORIZED => "authentication_error",
182        StatusCode::FORBIDDEN => "permission_error",
183        StatusCode::TOO_MANY_REQUESTS => "rate_limit_error",
184        _ => "api_error",
185    };
186    let response = ChatError::new(status, kind, message, None, None).response();
187    if let Some(retry_after) = retry_after {
188        let (mut parts, body) = response.into_parts();
189        parts.headers.insert(http::header::RETRY_AFTER, retry_after);
190        Response::from_parts(parts, body)
191    } else {
192        response
193    }
194}
195
196#[derive(Debug, Clone)]
197pub struct ChatError {
198    pub status: StatusCode,
199    pub kind: &'static str,
200    pub message: String,
201    pub param: Option<String>,
202    pub code: Option<String>,
203}
204
205impl ChatError {
206    pub fn new(
207        status: StatusCode,
208        kind: &'static str,
209        message: impl Into<String>,
210        param: Option<&str>,
211        code: Option<&str>,
212    ) -> Self {
213        Self {
214            status,
215            kind,
216            message: message.into(),
217            param: param.map(str::to_string),
218            code: code.map(str::to_string),
219        }
220    }
221
222    pub fn invalid(message: impl Into<String>, param: Option<&str>, code: Option<&str>) -> Self {
223        Self::new(
224            StatusCode::BAD_REQUEST,
225            "invalid_request_error",
226            message,
227            param,
228            code,
229        )
230    }
231
232    pub fn unsupported(param: impl Into<String>) -> Self {
233        let param = param.into();
234        Self::invalid(
235            format!("Unsupported parameter: {param}"),
236            Some(&param),
237            Some("unsupported_parameter"),
238        )
239    }
240
241    pub fn upstream(message: impl Into<String>) -> Self {
242        Self::new(StatusCode::BAD_GATEWAY, "api_error", message, None, None)
243    }
244
245    pub fn timeout(message: impl Into<String>) -> Self {
246        Self::new(
247            StatusCode::GATEWAY_TIMEOUT,
248            "api_error",
249            message,
250            None,
251            None,
252        )
253    }
254
255    fn from_codex(error: CodexError) -> Self {
256        let status = match error.status {
257            401 => StatusCode::UNAUTHORIZED,
258            403 => StatusCode::FORBIDDEN,
259            429 => StatusCode::TOO_MANY_REQUESTS,
260            _ if error.message.contains("Timed out waiting") => StatusCode::GATEWAY_TIMEOUT,
261            _ => StatusCode::BAD_GATEWAY,
262        };
263        let kind = match status {
264            StatusCode::UNAUTHORIZED => "authentication_error",
265            StatusCode::FORBIDDEN => "permission_error",
266            StatusCode::TOO_MANY_REQUESTS => "rate_limit_error",
267            _ => "api_error",
268        };
269        Self::new(
270            status,
271            kind,
272            error.detail.unwrap_or(error.message),
273            None,
274            None,
275        )
276    }
277
278    pub fn value(&self) -> Value {
279        json!({"error":{"message":self.message,"type":self.kind,"param":self.param,"code":self.code}})
280    }
281
282    pub fn response(self) -> Response {
283        (self.status, Json(self.value())).into_response()
284    }
285}
286
287#[cfg(test)]
288mod tests {
289    use super::*;
290    use crate::{
291        monitor::{EndpointKind, MonitorHandle},
292        providers::codex::auth::token_store::StoredAuth,
293    };
294    use tokio::{
295        io::{AsyncReadExt, AsyncWriteExt},
296        net::TcpListener,
297    };
298
299    fn context(monitor: MonitorHandle) -> RequestContext {
300        monitor.request_started(
301            "chat-test",
302            Some("session".into()),
303            None,
304            EndpointKind::ChatCompletions,
305        );
306        RequestContext {
307            req_id: "chat-test".into(),
308            session_id: Some("session".into()),
309            session_seq: None,
310            provider: "codex".into(),
311            traffic: None,
312            monitor: Some(monitor),
313            passthrough: None,
314        }
315    }
316
317    async fn mock_backend(
318        sse_body: &'static [u8],
319    ) -> (ChatCompletionsBackend, tokio::task::JoinHandle<Value>) {
320        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
321        let address = listener.local_addr().unwrap();
322        let server = tokio::spawn(async move {
323            let (mut socket, _) = listener.accept().await.unwrap();
324            let mut request = Vec::new();
325            let mut buffer = [0_u8; 4096];
326            loop {
327                let read = socket.read(&mut buffer).await.unwrap();
328                if read == 0 {
329                    break;
330                }
331                request.extend_from_slice(&buffer[..read]);
332                let Some(header_end) = request.windows(4).position(|window| window == b"\r\n\r\n")
333                else {
334                    continue;
335                };
336                let headers = String::from_utf8_lossy(&request[..header_end]);
337                let length = headers
338                    .lines()
339                    .find_map(|line| {
340                        line.to_ascii_lowercase()
341                            .strip_prefix("content-length:")?
342                            .trim()
343                            .parse::<usize>()
344                            .ok()
345                    })
346                    .unwrap_or(0);
347                if request.len() >= header_end + 4 + length {
348                    break;
349                }
350            }
351            let header_end = request
352                .windows(4)
353                .position(|window| window == b"\r\n\r\n")
354                .unwrap();
355            let body: Value = serde_json::from_slice(&request[header_end + 4..]).unwrap();
356            let response = format!(
357                "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\nx-request-id: upstream-1\r\nconnection: close\r\n\r\n",
358                sse_body.len()
359            );
360            socket.write_all(response.as_bytes()).await.unwrap();
361            socket.write_all(sse_body).await.unwrap();
362            body
363        });
364        let client = CodexHttpClient::new_for_test(
365            reqwest::Client::new(),
366            format!("http://{address}/v1/responses"),
367            1_000,
368            1_000,
369            0,
370        );
371        client.auth_manager().set_test_auth(StoredAuth {
372            access: "test-token".into(),
373            refresh: String::new(),
374            account_id: Some("account".into()),
375            expires: u64::MAX,
376        });
377        (ChatCompletionsBackend::with_client(client), server)
378    }
379
380    #[tokio::test]
381    async fn buffered_request_translates_upstream_and_downstream() {
382        const SSE: &[u8] = b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"{\\\"answer\\\":\\\"yes\\\"}\"}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_buffered\",\"model\":\"gpt-5.6-sol\",\"status\":\"completed\",\"usage\":{\"input_tokens\":8,\"output_tokens\":4}}}\n\n";
383        let (backend, server) = mock_backend(SSE).await;
384        let request = request::translate_request(json!({
385            "model":"gpt-5.6-sol",
386            "messages":[{"role":"system","content":"JSON only"},{"role":"user","content":"answer"}],
387            "reasoning_effort":"low",
388            "response_format":{"type":"json_schema","json_schema":{"name":"answer","strict":true,"schema":{"type":"object"}}}
389        })).unwrap();
390        let monitor = MonitorHandle::new(10);
391        let response = backend.handle(request, context(monitor.clone())).await;
392        assert_eq!(response.status(), StatusCode::OK);
393        assert_eq!(response.headers()["x-request-id"], "upstream-1");
394        let value: Value = serde_json::from_slice(
395            &axum::body::to_bytes(response.into_body(), usize::MAX)
396                .await
397                .unwrap(),
398        )
399        .unwrap();
400        assert_eq!(value["object"], "chat.completion");
401        assert_eq!(
402            value["choices"][0]["message"]["content"],
403            r#"{"answer":"yes"}"#
404        );
405        assert_eq!(value["usage"]["total_tokens"], 12);
406
407        let upstream = server.await.unwrap();
408        assert_eq!(upstream["store"], false);
409        assert_eq!(upstream["stream"], true);
410        assert_eq!(upstream["input"][0]["role"], "developer");
411        assert_eq!(upstream["reasoning"]["effort"], "low");
412        assert_eq!(upstream["reasoning"]["context"], "all_turns");
413        assert_eq!(upstream["text"]["format"]["name"], "answer");
414        let snapshot = monitor.snapshot();
415        assert_eq!(snapshot.active[0].input_tokens, Some(8));
416        assert_eq!(snapshot.active[0].output_tokens, Some(4));
417    }
418
419    #[tokio::test]
420    async fn streaming_request_emits_chat_chunks_usage_and_done() {
421        const SSE: &[u8] = b"data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_stream\",\"model\":\"gpt-5.6-sol\"}}\n\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"hello\"}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_stream\",\"status\":\"completed\",\"usage\":{\"input_tokens\":3,\"output_tokens\":1}}}\n\n";
422        let (backend, server) = mock_backend(SSE).await;
423        let request = request::translate_request(json!({
424            "model":"gpt-5.6-sol",
425            "messages":[{"role":"user","content":"hello"}],
426            "stream":true,
427            "stream_options":{"include_usage":true}
428        }))
429        .unwrap();
430        let response = backend
431            .handle(request, context(MonitorHandle::new(10)))
432            .await;
433        assert_eq!(response.status(), StatusCode::OK);
434        assert_eq!(response.headers()["content-type"], "text/event-stream");
435        let body = String::from_utf8(
436            axum::body::to_bytes(response.into_body(), usize::MAX)
437                .await
438                .unwrap()
439                .to_vec(),
440        )
441        .unwrap();
442        assert!(body.contains(r#""delta":{"role":"assistant"}"#));
443        assert!(body.contains(r#""delta":{"content":"hello"}"#));
444        assert!(body.contains(r#""finish_reason":"stop""#));
445        assert!(body.contains(r#""prompt_tokens":3"#));
446        assert!(body.ends_with("data: [DONE]\n\n"));
447        server.await.unwrap();
448    }
449
450    #[test]
451    fn codex_errors_map_status_and_preserve_retry_metadata() {
452        let response = codex_error_response(CodexError {
453            status: 429,
454            message: "Rate limited".into(),
455            detail: Some("Try later".into()),
456            retry_after: Some("7".into()),
457            origin: super::super::client::CodexErrorOrigin::Http,
458        });
459        assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
460        assert_eq!(response.headers()[http::header::RETRY_AFTER], "7");
461
462        let auth = ChatError::from_codex(CodexError {
463            status: 401,
464            message: "Auth error".into(),
465            detail: None,
466            retry_after: None,
467            origin: super::super::client::CodexErrorOrigin::Auth,
468        });
469        assert_eq!(auth.status, StatusCode::UNAUTHORIZED);
470        assert_eq!(auth.kind, "authentication_error");
471    }
472}