lean_ctx/proxy/
chatgpt.rs1use axum::{
2 body::Body,
3 extract::State,
4 http::{HeaderName, Request, StatusCode},
5 response::Response,
6};
7
8use super::{ProxyState, forward, openai_responses};
9
10pub async fn codex_responses_handler(
14 State(state): State<ProxyState>,
15 req: Request<Body>,
16) -> Result<Response, StatusCode> {
17 let upstream = state.chatgpt_upstream();
18 forward::forward_request(
19 State(state),
20 req,
21 &upstream,
22 "/backend-api/codex/responses",
23 openai_responses::compress_request_body,
24 "ChatGPT",
25 &[],
26 )
27 .await
28}
29
30pub async fn codex_responses_ws_handler(
33 State(_state): State<ProxyState>,
34 _headers: axum::http::HeaderMap,
35 _ws: axum::extract::ws::WebSocketUpgrade,
36) -> Response {
37 chatgpt_responses_ws_fallback_response()
38}
39
40fn chatgpt_responses_ws_fallback_response() -> Response {
41 Response::builder()
42 .status(StatusCode::UPGRADE_REQUIRED)
43 .header("content-type", "application/json")
44 .body(Body::from(
45 r#"{"error":{"type":"unsupported_transport","message":"ChatGPT codex responses use HTTP/SSE; retry without WebSocket."}}"#,
46 ))
47 .expect("static response is valid")
48}
49
50pub async fn backend_api_handler(
53 State(state): State<ProxyState>,
54 req: Request<Body>,
55) -> Result<Response, StatusCode> {
56 let (parts, body) = req.into_parts();
57 let body_bytes = axum::body::to_bytes(body, forward::max_body_bytes())
58 .await
59 .map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?;
60 let upstream = state.chatgpt_upstream();
61 let path = parts
62 .uri
63 .path_and_query()
64 .map_or("/backend-api", axum::http::uri::PathAndQuery::as_str);
65 let url = format!("{upstream}{path}");
66
67 let mut upstream_req = state.client.request(parts.method.clone(), &url);
68 for (key, value) in &parts.headers {
69 if is_backend_passthrough_request_header(key) {
70 upstream_req = upstream_req.header(key.clone(), value.clone());
71 }
72 }
73
74 let response = upstream_req
75 .body(body_bytes.to_vec())
76 .send()
77 .await
78 .map_err(|e| {
79 tracing::error!("lean-ctx proxy: ChatGPT backend upstream error: {e}");
80 StatusCode::BAD_GATEWAY
81 })?;
82
83 let status = StatusCode::from_u16(response.status().as_u16()).unwrap_or(StatusCode::OK);
84 let headers = response.headers().clone();
85 let is_stream = headers
86 .get("content-type")
87 .and_then(|v| v.to_str().ok())
88 .is_some_and(|ct| ct.contains("text/event-stream"));
89
90 let mut out = Response::builder().status(status);
91 for (key, value) in &headers {
92 if is_backend_passthrough_response_header(key) {
93 out = out.header(key, value);
94 }
95 }
96
97 if is_stream {
98 return out
99 .body(Body::from_stream(response.bytes_stream()))
100 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR);
101 }
102
103 let bytes = response
104 .bytes()
105 .await
106 .map_err(|_| StatusCode::BAD_GATEWAY)?;
107
108 out.body(Body::from(bytes))
109 .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)
110}
111
112fn is_backend_passthrough_request_header(name: &HeaderName) -> bool {
113 let lower = name.as_str().to_ascii_lowercase();
114 !matches!(
115 lower.as_str(),
116 "host"
117 | "connection"
118 | "content-length"
119 | "transfer-encoding"
120 | "upgrade"
121 | "keep-alive"
122 | "proxy-authenticate"
123 | "proxy-authorization"
124 | "te"
125 | "trailer"
126 | "accept-encoding"
127 )
128}
129
130fn is_backend_passthrough_response_header(name: &HeaderName) -> bool {
131 let lower = name.as_str().to_ascii_lowercase();
132 !matches!(
133 lower.as_str(),
134 "connection"
135 | "content-length"
136 | "transfer-encoding"
137 | "upgrade"
138 | "keep-alive"
139 | "proxy-authenticate"
140 | "proxy-authorization"
141 | "te"
142 | "trailer"
143 )
144}
145
146#[cfg(test)]
147mod tests {
148 use std::sync::Arc;
149 use std::time::Duration;
150
151 use tokio::io::{AsyncReadExt, AsyncWriteExt};
152
153 use super::*;
154 use crate::core::config::Upstreams;
155
156 fn proxy_state(chatgpt_upstream: String) -> ProxyState {
157 let (_tx, rx) = tokio::sync::watch::channel(Arc::new(Upstreams {
158 anthropic: "https://api.anthropic.com".into(),
159 openai: "https://api.openai.com".into(),
160 chatgpt: chatgpt_upstream,
161 gemini: "https://generativelanguage.googleapis.com".into(),
162 }));
163 ProxyState {
164 client: reqwest::Client::new(),
165 port: 0,
166 stats: Arc::new(crate::proxy::ProxyStats::default()),
167 introspect: Arc::new(crate::proxy::introspect::IntrospectState::default()),
168 upstreams: rx,
169 }
170 }
171
172 #[test]
173 fn codex_responses_ws_requests_trigger_http_fallback() {
174 let response = chatgpt_responses_ws_fallback_response();
175 assert_eq!(response.status(), StatusCode::UPGRADE_REQUIRED);
176 assert_eq!(
177 response
178 .headers()
179 .get(axum::http::header::CONTENT_TYPE)
180 .unwrap(),
181 "application/json"
182 );
183 }
184
185 async fn spawn_streaming_upstream() -> (String, tokio::sync::oneshot::Receiver<String>) {
186 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
187 let addr = listener.local_addr().unwrap();
188 let (tx, rx) = tokio::sync::oneshot::channel();
189 tokio::spawn(async move {
190 let (mut socket, _) = listener.accept().await.unwrap();
191 let mut buf = Vec::new();
192 loop {
193 let mut chunk = [0_u8; 1024];
194 let n = socket.read(&mut chunk).await.unwrap();
195 if n == 0 {
196 break;
197 }
198 buf.extend_from_slice(&chunk[..n]);
199 if buf.windows(4).any(|w| w == b"\r\n\r\n") {
200 break;
201 }
202 }
203 let _ = tx.send(String::from_utf8_lossy(&buf).into_owned());
204 socket
205 .write_all(
206 b"HTTP/1.1 200 OK\r\n\
207 content-type: text/event-stream\r\n\
208 mcp-session-id: server-session\r\n\
209 cache-control: no-cache\r\n\
210 x-custom-backend-state: passthrough\r\n\
211 \r\n\
212 event: message\n\
213 data: {\"jsonrpc\":\"2.0\"}\n\n",
214 )
215 .await
216 .unwrap();
217 tokio::time::sleep(Duration::from_secs(2)).await;
218 });
219 (format!("http://{addr}"), rx)
220 }
221
222 #[tokio::test]
223 async fn backend_api_streams_mcp_sse_and_preserves_session_headers() {
224 let (upstream, seen_request) = spawn_streaming_upstream().await;
225 let state = proxy_state(upstream);
226 let req = Request::builder()
227 .method("POST")
228 .uri("/backend-api/ps/mcp?transport=streamable")
229 .header("Authorization", "Bearer codex-token")
230 .header("Mcp-Session-Id", "client-session")
231 .header("Last-Event-ID", "event-7")
232 .header("X-OpenAI-Product-Sku", "codex")
233 .header("X-OpenAI-Internal-Codex-Residency", "us")
234 .header("Originator", "codex_cli_rs")
235 .header("Accept", "application/json, text/event-stream")
236 .body(Body::empty())
237 .unwrap();
238
239 let response = tokio::time::timeout(
240 Duration::from_millis(500),
241 backend_api_handler(State(state), req),
242 )
243 .await
244 .expect("SSE passthrough must return after upstream headers")
245 .expect("backend request should succeed");
246
247 assert_eq!(response.status(), StatusCode::OK);
248 assert_eq!(
249 response.headers().get("mcp-session-id").unwrap(),
250 "server-session"
251 );
252 assert_eq!(
253 response.headers().get("x-custom-backend-state").unwrap(),
254 "passthrough"
255 );
256
257 let request = seen_request.await.unwrap().to_ascii_lowercase();
258 assert!(request.contains("post /backend-api/ps/mcp?transport=streamable http/1.1"));
259 assert!(request.contains("authorization: bearer codex-token"));
260 assert!(request.contains("mcp-session-id: client-session"));
261 assert!(request.contains("last-event-id: event-7"));
262 assert!(request.contains("x-openai-product-sku: codex"));
263 assert!(request.contains("x-openai-internal-codex-residency: us"));
264 assert!(request.contains("originator: codex_cli_rs"));
265 }
266}