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