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(¶m),
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}