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}