Skip to main content

lean_ctx/proxy/
chatgpt.rs

1use axum::{
2    body::Body,
3    extract::State,
4    http::{HeaderName, Request, StatusCode},
5    response::Response,
6};
7
8use super::{ProxyState, forward, openai_responses};
9
10/// Codex subscription model turns hit ChatGPT's Responses-compatible rail:
11/// `/backend-api/codex/responses`. Forward through the same compressor/metering
12/// path as OpenAI Responses, but target `https://chatgpt.com`.
13pub 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
30/// ChatGPT's Codex rail rejects WS-only continuation fields such as
31/// `previous_response_id`; ask Codex to retry through the HTTP/SSE path.
32pub 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
50/// ChatGPT backend calls outside the model rail are not model JSON and must not be
51/// compressed or cost-metered. They are credential-preserving passthroughs.
52pub 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}