Skip to main content

claude_codex/
server.rs

1use crate::{
2    anthropic::json_error,
3    logging::{Logger, REDACT_KEYS, create_logger},
4    monitor::{EndpointKind, MonitorHandle},
5    provider::RequestContext,
6    registry::{Registry, normalize_incoming_model},
7    session::{self, SessionState},
8    traffic::{TrafficCaptureOptions, create_traffic_capture},
9};
10use axum::{
11    Json, Router,
12    body::Body,
13    extract::State,
14    http::{Request, StatusCode},
15    response::Response,
16    routing::{get, post},
17};
18use http_body_util::{BodyExt, StreamBody};
19use serde::de::DeserializeOwned;
20use serde_json::{Map, Value, json};
21use std::fs::{self, File};
22use std::future::Future;
23use std::io::Write;
24use std::path::{Path, PathBuf};
25use std::sync::Arc;
26use std::time::Instant;
27use tokio::net::TcpListener;
28use uuid::Uuid;
29
30pub struct ServerConfig {
31    pub port: u16,
32    pub monitor: Option<MonitorHandle>,
33}
34
35pub async fn serve(config: ServerConfig) -> anyhow::Result<()> {
36    serve_inner(config, std::future::pending::<()>()).await
37}
38
39pub async fn serve_with_shutdown(
40    config: ServerConfig,
41    shutdown: impl Future<Output = ()> + Send + 'static,
42) -> anyhow::Result<()> {
43    serve_inner(config, shutdown).await
44}
45
46async fn serve_inner(
47    config: ServerConfig,
48    shutdown: impl Future<Output = ()> + Send + 'static,
49) -> anyhow::Result<()> {
50    let listener = bind_proxy_listener(config.port).await?;
51    serve_listener(listener, config.monitor, shutdown).await
52}
53
54pub async fn bind_proxy_listener(port: u16) -> anyhow::Result<TcpListener> {
55    let addr = format!("127.0.0.1:{port}");
56    TcpListener::bind(&addr)
57        .await
58        .map_err(|err| anyhow::anyhow!("failed to bind proxy listener on {addr}: {err}"))
59}
60
61pub async fn serve_listener(
62    listener: TcpListener,
63    monitor: Option<MonitorHandle>,
64    shutdown: impl Future<Output = ()> + Send + 'static,
65) -> anyhow::Result<()> {
66    let port = listener.local_addr()?.port();
67    create_logger("server").info(
68        "server listening",
69        Some(serde_json::Map::from_iter([
70            ("port".to_string(), json!(port)),
71            (
72                "logDir".to_string(),
73                json!(
74                    crate::paths::log_file()
75                        .parent()
76                        .map(|path| path.display().to_string())
77                ),
78            ),
79        ])),
80    );
81    let app = app_with_monitor(Arc::new(Registry::with_default_alias()), monitor);
82    axum::serve(listener, app)
83        .with_graceful_shutdown(shutdown)
84        .await?;
85    Ok(())
86}
87
88pub fn app(registry: Arc<Registry>) -> Router {
89    app_with_monitor(registry, None)
90}
91
92pub fn app_with_monitor(registry: Arc<Registry>, monitor: Option<MonitorHandle>) -> Router {
93    let state = Arc::new(AppState { registry, monitor });
94    Router::new()
95        .route("/healthz", get(healthz))
96        .route("/v1/messages", post(handler_messages))
97        .route("/v1/messages/count_tokens", post(handler_count_tokens))
98        .fallback(fallback_handler)
99        .with_state(state)
100}
101
102#[derive(Clone)]
103struct AppState {
104    registry: Arc<Registry>,
105    monitor: Option<MonitorHandle>,
106}
107
108async fn healthz() -> Json<serde_json::Value> {
109    Json(json!({ "ok": true }))
110}
111
112async fn handler_messages(State(state): State<Arc<AppState>>, req: Request<Body>) -> Response {
113    dispatch_request(state, req, false).await
114}
115
116async fn handler_count_tokens(State(state): State<Arc<AppState>>, req: Request<Body>) -> Response {
117    dispatch_request(state, req, true).await
118}
119
120async fn dispatch_request(
121    state: Arc<AppState>,
122    req: Request<Body>,
123    count_tokens: bool,
124) -> Response {
125    let started_at = Instant::now();
126    let log = create_logger("server");
127    let req_id = Uuid::new_v4().to_string();
128    let method = req.method().clone();
129    let uri = req.uri().clone();
130    let headers = req.headers().clone();
131    let path = uri.path().to_string();
132    let query = redacted_query(&uri);
133    let endpoint = if count_tokens {
134        EndpointKind::CountTokens
135    } else {
136        EndpointKind::Messages
137    };
138    log.info(
139        "request",
140        Some(serde_json::Map::from_iter([
141            ("reqId".to_string(), json!(&req_id)),
142            ("method".to_string(), json!(method.as_str())),
143            ("path".to_string(), json!(&path)),
144            ("query".to_string(), json!(&query)),
145        ])),
146    );
147    let session_id = req
148        .headers()
149        .get("x-claude-code-session-id")
150        .and_then(|value| value.to_str().ok())
151        .map(std::string::ToString::to_string);
152    if let Some(monitor) = state.monitor.as_ref() {
153        monitor.request_started(&req_id, session_id.clone(), None, endpoint);
154    }
155    let request_guard = RequestMonitorGuard::new(state.monitor.clone(), req_id.clone());
156    let now = current_millis();
157    let body_bytes = match axum::body::to_bytes(req.into_body(), usize::MAX).await {
158        Ok(bytes) => bytes,
159        Err(err) => {
160            let response = json_error(
161                StatusCode::BAD_REQUEST,
162                "invalid_request_error",
163                format!("Invalid JSON: {err}"),
164            );
165            log_request_completed(
166                &log,
167                RequestLogContext {
168                    req_id: &req_id,
169                    provider: None,
170                    model: None,
171                    count_tokens,
172                    status: response.status(),
173                    started_at,
174                },
175            );
176            let (response, details) = record_failed_response(
177                &log,
178                FailedResponseLogContext {
179                    req_id: &req_id,
180                    provider: None,
181                    model: None,
182                    count_tokens,
183                    started_at,
184                },
185                response,
186            )
187            .await;
188            monitor_failed(
189                state.monitor.as_ref(),
190                &req_id,
191                Some(response.status()),
192                details
193                    .as_ref()
194                    .map(|details| details.message.as_str())
195                    .unwrap_or("Invalid JSON"),
196            );
197            return response;
198        }
199    };
200
201    let body: crate::anthropic::schema::MessagesRequest = match parse_json_body(&body_bytes) {
202        Ok(body) => body,
203        Err(response) => {
204            let status = response.status();
205            log_request_completed(
206                &log,
207                RequestLogContext {
208                    req_id: &req_id,
209                    provider: None,
210                    model: None,
211                    count_tokens,
212                    status: response.status(),
213                    started_at,
214                },
215            );
216            let (response, details) = record_failed_response(
217                &log,
218                FailedResponseLogContext {
219                    req_id: &req_id,
220                    provider: None,
221                    model: None,
222                    count_tokens,
223                    started_at,
224                },
225                *response,
226            )
227            .await;
228            monitor_failed(
229                state.monitor.as_ref(),
230                &req_id,
231                Some(status),
232                details
233                    .as_ref()
234                    .map(|details| details.message.as_str())
235                    .unwrap_or("Invalid JSON"),
236            );
237            return response;
238        }
239    };
240
241    let model = match body.model.as_deref() {
242        Some(model) => model,
243        None => {
244            let response = json_error(
245                StatusCode::BAD_REQUEST,
246                "invalid_request_error",
247                format!(
248                    "Missing \"model\" in request body. {}",
249                    state.registry.unknown_model_message()
250                ),
251            );
252            log_request_completed(
253                &log,
254                RequestLogContext {
255                    req_id: &req_id,
256                    provider: None,
257                    model: None,
258                    count_tokens,
259                    status: response.status(),
260                    started_at,
261                },
262            );
263            let (response, details) = record_failed_response(
264                &log,
265                FailedResponseLogContext {
266                    req_id: &req_id,
267                    provider: None,
268                    model: None,
269                    count_tokens,
270                    started_at,
271                },
272                response,
273            )
274            .await;
275            monitor_failed(
276                state.monitor.as_ref(),
277                &req_id,
278                Some(response.status()),
279                details
280                    .as_ref()
281                    .map(|details| details.message.as_str())
282                    .unwrap_or("Missing model"),
283            );
284            return response;
285        }
286    };
287
288    let normalized_model = normalize_incoming_model(model);
289    let session_state = if let Some(session_id) = session_id.as_deref() {
290        session::existing_session(Some(session_id), now)
291    } else {
292        None
293    };
294
295    let provider = state.registry.provider_for_model(
296        &normalized_model,
297        session_state
298            .as_ref()
299            .and_then(|state| state.affinity_provider.as_ref()),
300    );
301
302    let provider = match provider {
303        Some(provider) => provider,
304        None => {
305            log.warn(
306                "unknown model",
307                Some(serde_json::Map::from_iter([
308                    ("reqId".to_string(), json!(&req_id)),
309                    ("model".to_string(), json!(&normalized_model)),
310                ])),
311            );
312            let response = json_error(
313                StatusCode::BAD_REQUEST,
314                "invalid_request_error",
315                format!(
316                    "Unknown model \"{normalized_model}\". {}",
317                    state.registry.unknown_model_message()
318                ),
319            );
320            log_request_completed(
321                &log,
322                RequestLogContext {
323                    req_id: &req_id,
324                    provider: None,
325                    model: Some(&normalized_model),
326                    count_tokens,
327                    status: response.status(),
328                    started_at,
329                },
330            );
331            let (response, details) = record_failed_response(
332                &log,
333                FailedResponseLogContext {
334                    req_id: &req_id,
335                    provider: None,
336                    model: Some(&normalized_model),
337                    count_tokens,
338                    started_at,
339                },
340                response,
341            )
342            .await;
343            monitor_failed(
344                state.monitor.as_ref(),
345                &req_id,
346                Some(response.status()),
347                details
348                    .as_ref()
349                    .map(|details| details.message.as_str())
350                    .unwrap_or("Unknown model"),
351            );
352            return response;
353        }
354    };
355
356    let effort = crate::providers::translate_shared::read_effort(&body)
357        .ok()
358        .flatten()
359        .map(str::to_string);
360    let current = session::record_session_request(
361        session_id.as_deref(),
362        session_state.as_ref(),
363        provider.name(),
364        &normalized_model,
365        now,
366    );
367    if let Some(monitor) = state.monitor.as_ref() {
368        if let Some(current) = current.as_ref() {
369            monitor.request_started(&req_id, session_id.clone(), Some(current.seq), endpoint);
370        }
371        monitor.provider_selected(&req_id, provider.name(), &normalized_model, effort);
372    }
373
374    let traffic = create_traffic_capture(TrafficCaptureOptions {
375        req_id: req_id.clone(),
376        session_id: session_id.clone(),
377        session_seq: current.as_ref().map(|s| s.seq),
378        provider: Some(provider.name().to_string()),
379        state_dir_override: None,
380    })
381    .map(Arc::new);
382
383    if let Some(capture) = traffic.as_ref() {
384        if let Some(monitor) = state.monitor.as_ref() {
385            monitor.traffic_capture_path(&req_id, capture.root().to_path_buf());
386        }
387        capture.write_json(
388            "000-metadata",
389            &json!({
390                "reqId": &req_id,
391                "sessionId": &session_id,
392                "sessionSeq": current.as_ref().map(|s| s.seq),
393                "kind": if count_tokens { "count_tokens" } else { "messages" },
394                "provider": provider.name(),
395                "model": &normalized_model,
396                "method": method.as_str(),
397                "path": &path,
398                "query": &query,
399                "headers": headers_to_record(&headers),
400            }),
401        );
402        capture.write_json(
403            "010-anthropic-request",
404            &serde_json::to_value(&body).unwrap_or_else(|_| json!({})),
405        );
406    }
407
408    let context = RequestContext {
409        req_id: req_id.clone(),
410        session_id,
411        session_seq: current.map(|s| s.seq),
412        provider: provider.name().to_string(),
413        traffic,
414        monitor: state.monitor.clone(),
415        passthrough: Some(crate::provider::Passthrough {
416            raw_body: body_bytes,
417            headers,
418            path_and_query: uri
419                .path_and_query()
420                .map(|pq| pq.as_str().to_string())
421                .unwrap_or_else(|| path.clone()),
422        }),
423    };
424
425    let response = if count_tokens {
426        provider.handle_count_tokens(body, context).await
427    } else {
428        provider.handle_messages(body, context).await
429    };
430    log_request_completed(
431        &log,
432        RequestLogContext {
433            req_id: &req_id,
434            provider: Some(provider.name()),
435            model: Some(&normalized_model),
436            count_tokens,
437            status: response.status(),
438            started_at,
439        },
440    );
441    let status = response.status();
442    if status.is_success() {
443        return monitor_response_body(response, request_guard);
444    }
445
446    let (response, details) = record_failed_response(
447        &log,
448        FailedResponseLogContext {
449            req_id: &req_id,
450            provider: Some(provider.name()),
451            model: Some(&normalized_model),
452            count_tokens,
453            started_at,
454        },
455        response,
456    )
457    .await;
458    if let Some(details) = details.as_ref() {
459        monitor_failed(
460            state.monitor.as_ref(),
461            &req_id,
462            Some(status),
463            details.message.as_str(),
464        );
465    } else {
466        monitor_failed(
467            state.monitor.as_ref(),
468            &req_id,
469            Some(status),
470            format!("HTTP {}", status.as_u16()),
471        );
472    }
473    response
474}
475
476fn monitor_response_body(response: Response, guard: RequestMonitorGuard) -> Response {
477    let status = response.status();
478    let (parts, body) = response.into_parts();
479    let stream =
480        futures_util::stream::unfold((body, guard), move |(mut body, mut guard)| async move {
481            match body.frame().await {
482                Some(Ok(frame)) => Some((Ok(frame), (body, guard))),
483                Some(Err(err)) => {
484                    guard.failed(status, err.to_string());
485                    Some((Err(err), (body, guard)))
486                }
487                None => {
488                    guard.completed(status);
489                    None
490                }
491            }
492        });
493    Response::from_parts(parts, Body::new(StreamBody::new(stream)))
494}
495
496struct RequestLogContext<'a> {
497    req_id: &'a str,
498    provider: Option<&'a str>,
499    model: Option<&'a str>,
500    count_tokens: bool,
501    status: StatusCode,
502    started_at: Instant,
503}
504
505fn log_request_completed(log: &Logger, ctx: RequestLogContext<'_>) {
506    log.info(
507        "request_completed",
508        Some(serde_json::Map::from_iter([
509            ("reqId".to_string(), json!(ctx.req_id)),
510            ("provider".to_string(), json!(ctx.provider)),
511            ("model".to_string(), json!(ctx.model)),
512            ("countTokens".to_string(), json!(ctx.count_tokens)),
513            ("status".to_string(), json!(ctx.status.as_u16())),
514            (
515                "ms".to_string(),
516                json!(ctx.started_at.elapsed().as_millis()),
517            ),
518        ])),
519    );
520}
521
522struct FailedResponseLogContext<'a> {
523    req_id: &'a str,
524    provider: Option<&'a str>,
525    model: Option<&'a str>,
526    count_tokens: bool,
527    started_at: Instant,
528}
529
530struct FailedResponseDetails {
531    message: String,
532}
533
534async fn record_failed_response(
535    log: &Logger,
536    ctx: FailedResponseLogContext<'_>,
537    response: Response,
538) -> (Response, Option<FailedResponseDetails>) {
539    if response.status().is_success() {
540        return (response, None);
541    }
542
543    let status = response.status();
544    let (parts, body) = response.into_parts();
545    let bytes = match body.collect().await {
546        Ok(collected) => collected.to_bytes(),
547        Err(err) => {
548            log.info(
549                "request_failed",
550                Some(serde_json::Map::from_iter([
551                    ("reqId".to_string(), json!(ctx.req_id)),
552                    ("provider".to_string(), json!(ctx.provider)),
553                    ("model".to_string(), json!(ctx.model)),
554                    ("countTokens".to_string(), json!(ctx.count_tokens)),
555                    ("status".to_string(), json!(status.as_u16())),
556                    (
557                        "ms".to_string(),
558                        json!(ctx.started_at.elapsed().as_millis()),
559                    ),
560                    ("bodyReadError".to_string(), json!(err.to_string())),
561                ])),
562            );
563            return (Response::from_parts(parts, Body::empty()), None);
564        }
565    };
566
567    let response_body = response_body_value(&bytes);
568    let message = error_message_from_response(&response_body)
569        .unwrap_or_else(|| format!("HTTP {}", status.as_u16()));
570    let document = json!({
571        "reqId": ctx.req_id,
572        "provider": ctx.provider,
573        "model": ctx.model,
574        "countTokens": ctx.count_tokens,
575        "status": status.as_u16(),
576        "elapsedMs": ctx.started_at.elapsed().as_millis(),
577        "message": message,
578        "response": response_body,
579    });
580    let error_file = write_error_capture(ctx.req_id, &redact_error_value(document));
581
582    let mut fields = serde_json::Map::from_iter([
583        ("reqId".to_string(), json!(ctx.req_id)),
584        ("provider".to_string(), json!(ctx.provider)),
585        ("model".to_string(), json!(ctx.model)),
586        ("countTokens".to_string(), json!(ctx.count_tokens)),
587        ("status".to_string(), json!(status.as_u16())),
588        (
589            "ms".to_string(),
590            json!(ctx.started_at.elapsed().as_millis()),
591        ),
592        ("message".to_string(), json!(message)),
593    ]);
594    if let Some(path) = error_file.as_ref() {
595        fields.insert("errorFile".to_string(), json!(path.display().to_string()));
596    }
597    log.info("request_failed", Some(fields));
598
599    (
600        Response::from_parts(parts, Body::from(bytes)),
601        Some(FailedResponseDetails { message }),
602    )
603}
604
605fn response_body_value(bytes: &[u8]) -> Value {
606    match serde_json::from_slice::<Value>(bytes) {
607        Ok(value) => json!({ "json": value }),
608        Err(_) => json!({ "text": String::from_utf8_lossy(bytes) }),
609    }
610}
611
612fn error_message_from_response(response_body: &Value) -> Option<String> {
613    response_body
614        .get("json")
615        .and_then(|body| body.get("error"))
616        .and_then(|error| error.get("message"))
617        .and_then(Value::as_str)
618        .or_else(|| {
619            response_body
620                .get("text")
621                .and_then(Value::as_str)
622                .map(str::trim)
623                .filter(|text| !text.is_empty())
624        })
625        .map(std::string::ToString::to_string)
626}
627
628fn write_error_capture(req_id: &str, document: &Value) -> Option<PathBuf> {
629    let dir = crate::paths::state_dir().join("errors");
630    fs::create_dir_all(&dir).ok()?;
631    set_mode(&dir, 0o700);
632    let path = dir.join(format!(
633        "{}-{}.json",
634        current_millis(),
635        sanitize_path_part(req_id)
636    ));
637    let mut file = File::create(&path).ok()?;
638    set_mode(&path, 0o600);
639    let payload = serde_json::to_vec_pretty(document).ok()?;
640    file.write_all(&payload).ok()?;
641    file.write_all(b"\n").ok()?;
642    Some(path)
643}
644
645fn sanitize_path_part(raw: &str) -> String {
646    let sanitized: String = raw
647        .chars()
648        .map(|ch| {
649            if ch.is_ascii_alphanumeric() || matches!(ch, '-' | '_' | '.') {
650                ch
651            } else {
652                '_'
653            }
654        })
655        .collect();
656    if sanitized.is_empty() {
657        "unknown".to_string()
658    } else {
659        sanitized
660    }
661}
662
663fn redact_error_value(value: Value) -> Value {
664    match value {
665        Value::Array(values) => Value::Array(values.into_iter().map(redact_error_value).collect()),
666        Value::Object(fields) => {
667            let mut out = Map::new();
668            for (key, value) in fields {
669                if REDACT_KEYS.contains(&key.to_lowercase().as_str()) {
670                    out.insert(key, redact_error_key(value));
671                } else {
672                    out.insert(key, redact_error_value(value));
673                }
674            }
675            Value::Object(out)
676        }
677        value => value,
678    }
679}
680
681fn redact_error_key(value: Value) -> Value {
682    match value {
683        Value::String(value) => Value::String(format!("[redacted len={}]", value.len())),
684        _ => Value::String("[redacted]".to_string()),
685    }
686}
687
688struct RequestMonitorGuard {
689    monitor: Option<MonitorHandle>,
690    req_id: String,
691}
692
693impl RequestMonitorGuard {
694    fn new(monitor: Option<MonitorHandle>, req_id: String) -> Self {
695        Self { monitor, req_id }
696    }
697
698    fn completed(&mut self, status: StatusCode) {
699        if let Some(monitor) = self.monitor.take() {
700            monitor.request_completed(&self.req_id, status.as_u16(), None, None);
701        }
702    }
703
704    fn failed(&mut self, status: StatusCode, error: String) {
705        if let Some(monitor) = self.monitor.take() {
706            monitor.request_failed(&self.req_id, Some(status.as_u16()), error);
707        }
708    }
709}
710
711impl Drop for RequestMonitorGuard {
712    fn drop(&mut self) {
713        if let Some(monitor) = self.monitor.as_ref() {
714            monitor.request_abandoned(&self.req_id, "Request future ended before completion");
715        }
716    }
717}
718
719fn monitor_failed(
720    monitor: Option<&MonitorHandle>,
721    req_id: &str,
722    status: Option<StatusCode>,
723    error: impl Into<String>,
724) {
725    if let Some(monitor) = monitor {
726        monitor.request_failed(req_id, status.map(|status| status.as_u16()), error);
727    }
728}
729
730fn headers_to_record(headers: &http::HeaderMap) -> Value {
731    let mut out = Map::new();
732    for (key, value) in headers {
733        if let Ok(raw) = value.to_str() {
734            let recorded = if REDACT_KEYS.contains(&key.as_str().to_lowercase().as_str()) {
735                format!("[redacted len={}]", raw.len())
736            } else {
737                raw.to_string()
738            };
739            out.insert(key.as_str().to_string(), Value::String(recorded));
740        }
741    }
742    Value::Object(out)
743}
744
745fn redacted_query(uri: &http::Uri) -> Value {
746    let mut out = Map::new();
747    let Some(query) = uri.query() else {
748        return Value::Object(out);
749    };
750    for (key, value) in url::form_urlencoded::parse(query.as_bytes()) {
751        let key = key.into_owned();
752        let lower = key.to_lowercase();
753        let value = if REDACT_KEYS.contains(&lower.as_str()) {
754            Value::String(format!("[redacted len={}]", value.len()))
755        } else {
756            Value::String(value.into_owned())
757        };
758        out.insert(key, value);
759    }
760    Value::Object(out)
761}
762
763fn parse_json_body<T>(body: &[u8]) -> Result<T, Box<Response>>
764where
765    T: DeserializeOwned,
766{
767    if body.is_empty() {
768        return Err(Box::new(json_error(
769            StatusCode::BAD_REQUEST,
770            "invalid_request_error",
771            "Invalid JSON: empty body",
772        )));
773    }
774
775    serde_json::from_slice::<T>(body).map_err(|err| {
776        Box::new(json_error(
777            StatusCode::BAD_REQUEST,
778            "invalid_request_error",
779            format!("Invalid JSON: {err}"),
780        ))
781    })
782}
783
784async fn fallback_handler(method: axum::http::Method, uri: axum::http::Uri) -> Response {
785    json_error(
786        StatusCode::NOT_FOUND,
787        "not_found",
788        format!("No route for {method} {}", uri.path()),
789    )
790}
791
792fn current_millis() -> u64 {
793    use std::time::{SystemTime, UNIX_EPOCH};
794    SystemTime::now()
795        .duration_since(UNIX_EPOCH)
796        .unwrap_or_default()
797        .as_millis() as u64
798}
799
800fn set_mode(path: &Path, mode: u32) {
801    #[cfg(unix)]
802    {
803        use std::os::unix::fs::PermissionsExt;
804        if let Ok(meta) = fs::metadata(path) {
805            let mut perm = meta.permissions();
806            perm.set_mode(mode);
807            let _ = fs::set_permissions(path, perm);
808        }
809    }
810    #[cfg(not(unix))]
811    {
812        let _ = (path, mode);
813    }
814}
815
816#[allow(dead_code)]
817fn _unused(session_state: Option<&SessionState>) {
818    let _ = session_state;
819}