1use std::sync::Arc;
4
5use axum::extract::State;
6use axum::http::{HeaderMap, HeaderValue};
7use axum::response::sse::{Event, Sse};
8use axum::response::{IntoResponse, Response};
9use axum::routing::{get, post};
10use axum::{Json, Router};
11use serde_json::json;
12use tap::Tap;
13
14use crate::error::GatewayError;
15use crate::gateway::{Gateway, GuardedStream, RequestCtx};
16use crate::routing::jev_router::RoutingReport;
17use crate::routing::request::ChatRequest;
18use crate::routing::stream::{stream_item_to_sse_json, Accumulator, StreamItem};
19
20#[derive(Clone)]
21pub struct AppState {
22 pub gateway: Arc<Gateway>,
23}
24
25pub fn router(gateway: Arc<Gateway>) -> Router {
26 Router::new()
27 .route("/health", get(|| async { "ok" }))
28 .route("/v1/models", get(list_models))
29 .route("/v1/chat/completions", post(chat_completions))
30 .route("/v1/embeddings", post(embeddings))
31 .route("/v1beta/models/{model_action}", post(gemini_passthrough))
36 .route("/v1/models/{model_action}", post(gemini_passthrough))
37 .route("/google/models/{model_action}", post(gemini_passthrough))
38 .route("/typesafe/v1/systemone", post(jev_passthrough))
41 .with_state(AppState { gateway })
42}
43
44fn request_ctx(headers: &HeaderMap) -> RequestCtx {
45 let header = |name: &str| {
46 headers
47 .get(name)
48 .and_then(|v| v.to_str().ok())
49 .map(str::to_string)
50 };
51 RequestCtx {
52 tenant: header("x-synapse-tenant"),
53 workspace: header("x-synapse-workspace"),
54 user: header("x-synapse-user"),
55 thread: header("x-synapse-thread"),
56 message: header("x-synapse-message"),
57 user_task_type: header("x-synapse-user-task-type"),
58 ai_task_type: header("x-synapse-ai-task-type"),
59 request_id: None,
60 }
61}
62
63async fn list_models(State(st): State<AppState>) -> impl IntoResponse {
64 let data = st
65 .gateway
66 .model_aliases()
67 .into_iter()
68 .map(|id| json!({ "id": id, "object": "model", "owned_by": "synapse" }))
69 .collect::<Vec<_>>();
70 Json(json!({ "object": "list", "data": data }))
71}
72
73async fn chat_completions(
74 State(st): State<AppState>,
75 headers: HeaderMap,
76 Json(req): Json<ChatRequest>,
77) -> Result<Response, GatewayError> {
78 let headers_ctx = request_ctx(&headers);
81 let request_id = headers_ctx
82 .resolved_request_id()
83 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
84 let ctx = RequestCtx {
85 request_id: Some(request_id.clone()),
86 ..headers_ctx
87 };
88
89 if req.stream == Some(true) {
90 let stream = st.gateway.chat_stream(req, &ctx).await?;
91 let routing = stream.routing().clone();
92 return Ok(with_routing_headers(
93 Sse::new(sse_body(stream, request_id)).into_response(),
94 &routing,
95 ));
96 }
97
98 let (outcome, routing) = st.gateway.chat_routed(req, &ctx).await?;
99 let response = match outcome {
100 crate::gateway::ChatOutcome::Plain(completion) => {
101 Json(openai_json(&completion, &request_id)).into_response()
102 }
103 crate::gateway::ChatOutcome::Hybrid(h) => {
104 Json(hybrid_json(&h, &request_id)).into_response()
105 }
106 };
107 Ok(with_routing_headers(response, &routing))
108}
109
110fn with_routing_headers(response: Response, routing: &RoutingReport) -> Response {
113 routing
114 .headers()
115 .into_iter()
116 .filter_map(|(name, value)| HeaderValue::from_str(&value).ok().map(|v| (name, v)))
117 .fold(response, |r, (name, v)| {
118 r.tap_mut(|r| {
119 r.headers_mut().insert(name, v);
120 })
121 })
122}
123
124async fn embeddings(
125 State(st): State<AppState>,
126 headers: HeaderMap,
127 Json(req): Json<crate::embeddings::EmbeddingRequest>,
128) -> Result<Response, GatewayError> {
129 let ctx = request_ctx(&headers);
130 let resp = st.gateway.embed(req, ctx).await?;
131 Ok(Json(resp).into_response())
132}
133
134async fn gemini_passthrough(
143 State(st): State<AppState>,
144 axum::extract::Path(model_action): axum::extract::Path<String>,
145 axum::extract::RawQuery(query): axum::extract::RawQuery,
146 headers: HeaderMap,
147 Json(body): Json<serde_json::Value>,
148) -> Result<Response, GatewayError> {
149 use axum::body::Body;
150 use axum::http::{header, StatusCode};
151
152 let provider = st
153 .gateway
154 .vertex_native
155 .as_ref()
156 .ok_or_else(|| {
157 GatewayError::BadRequest("gemini passthrough requires the native vertex lane".into())
158 })?
159 .clone();
160
161 let (model, action) = model_action.rsplit_once(':').ok_or_else(|| {
162 GatewayError::BadRequest(format!(
163 "expected models/{{model}}:{{action}}, got '{model_action}'"
164 ))
165 })?;
166 let alt_sse = query.as_deref().is_some_and(|q| q.contains("alt=sse"));
167 let streaming = action == "streamGenerateContent";
168 let metered = action == "generateContent" || streaming;
169
170 if !metered {
172 let resp = provider
173 .passthrough_request(model, action, false, body, None)
174 .await?;
175 let status = resp.status();
176 st.gateway
177 .metrics
178 .passthrough("vertex", model, action, status.is_success());
179 let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
180 status: 502,
181 body: e.to_string(),
182 })?;
183 return passthrough_response(status.as_u16(), "application/json", bytes);
184 }
185
186 let chain = st.gateway.routes.vertex_fallback_chain(model);
187 let ctx = request_ctx(&headers);
188 let route_alias = chain.route.as_deref();
189 let mut prev_model: Option<String> = None;
190 let mut last_failure: Option<(u16, axum::body::Bytes)> = None;
191
192 for (i, leg) in chain.legs.iter().enumerate() {
193 if let Some(from) = prev_model.take() {
194 st.gateway.metrics.passthrough_fallback(&from, &leg.model);
195 }
196
197 let attempt = provider
198 .passthrough_request(
199 &leg.model,
200 action,
201 alt_sse && streaming,
202 body.clone(),
203 leg.region.as_deref(),
204 )
205 .await;
206
207 let resp = match attempt {
208 Ok(r) => r,
209 Err(e) => {
210 meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
211 st.gateway
212 .metrics
213 .passthrough("vertex", &leg.model, action, false);
214 if i + 1 < chain.legs.len() {
215 prev_model = Some(leg.model.clone());
216 continue;
217 }
218 return Err(e);
219 }
220 };
221
222 let status = resp.status();
223 st.gateway
224 .metrics
225 .passthrough("vertex", &leg.model, action, status.is_success());
226
227 if status.is_success() {
228 let mut guard = PassthroughUsageGuard::new(
229 &st.gateway,
230 &ctx,
231 &leg.model,
232 route_alias,
233 "vertex",
234 "chat",
235 );
236 if streaming && alt_sse {
237 let content_type = resp
238 .headers()
239 .get(header::CONTENT_TYPE)
240 .and_then(|v| v.to_str().ok())
241 .unwrap_or("text/event-stream")
242 .to_string();
243 let metered_stream = MeteredSseStream {
244 inner: resp.bytes_stream(),
245 guard,
246 line_buf: String::new(),
247 };
248 let response = Response::builder()
249 .status(StatusCode::OK)
250 .header(header::CONTENT_TYPE, content_type)
251 .body(Body::from_stream(metered_stream))
252 .map_err(|e| GatewayError::Upstream {
253 status: 502,
254 body: e.to_string(),
255 })?;
256 return Ok(response);
257 }
258
259 let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
260 status: 502,
261 body: e.to_string(),
262 })?;
263 if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
264 guard.observe_usage_metadata(&value);
265 }
266 drop(guard);
267 return passthrough_response(status.as_u16(), "application/json", bytes);
268 }
269
270 let bytes = resp.bytes().await.unwrap_or_default();
271 meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
272
273 if passthrough_status_retryable(status) && i + 1 < chain.legs.len() {
274 prev_model = Some(leg.model.clone());
275 last_failure = Some((status.as_u16(), bytes));
276 continue;
277 }
278 return passthrough_response(status.as_u16(), "application/json", bytes);
279 }
280
281 if let Some((code, bytes)) = last_failure {
282 return passthrough_response(code, "application/json", bytes);
283 }
284 Err(GatewayError::Upstream {
285 status: 502,
286 body: "gemini passthrough: empty vertex fallback chain".into(),
287 })
288}
289
290async fn jev_passthrough(
298 State(st): State<AppState>,
299 headers: HeaderMap,
300 Json(mut body): Json<serde_json::Value>,
301) -> Result<Response, GatewayError> {
302 let provider = st
303 .gateway
304 .jev_native
305 .as_ref()
306 .ok_or_else(|| {
307 GatewayError::BadRequest(
308 "jev passthrough requires TYPESAFE_API_KEY to be configured".into(),
309 )
310 })?
311 .clone();
312
313 if !body.is_object() {
314 return Err(GatewayError::BadRequest(
315 "expected a JSON object body".into(),
316 ));
317 }
318 if body.get("model").is_none() {
319 body["model"] = serde_json::Value::from(crate::jev_native::DEFAULT_MODEL);
320 }
321 let model = body["model"]
322 .as_str()
323 .unwrap_or(crate::jev_native::DEFAULT_MODEL)
324 .to_string();
325
326 let ctx = request_ctx(&headers);
327 let mut guard =
328 PassthroughUsageGuard::new(&st.gateway, &ctx, &model, None, "typesafe", "systemone");
329
330 let resp = match provider.evaluate(body).await {
331 Ok(r) => r,
332 Err(e) => {
333 guard.status = "error";
334 return Err(e);
335 }
336 };
337
338 let status = resp.status();
339 st.gateway
340 .metrics
341 .passthrough("typesafe", &model, "systemone", status.is_success());
342
343 if status.is_success() {
344 let bytes = resp.bytes().await.map_err(|e| {
345 guard.status = "error";
346 GatewayError::Upstream {
347 status: 502,
348 body: e.to_string(),
349 }
350 })?;
351 if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
352 guard.observe_usage_metadata(&value);
353 }
354 return passthrough_response(status.as_u16(), "application/json", bytes);
355 }
356
357 let bytes = resp.bytes().await.unwrap_or_default();
358 guard.status = "error";
359 passthrough_response(status.as_u16(), "application/json", bytes)
360}
361
362fn passthrough_status_retryable(status: reqwest::StatusCode) -> bool {
363 status.is_server_error()
364 || status == reqwest::StatusCode::TOO_MANY_REQUESTS
365 || status == reqwest::StatusCode::REQUEST_TIMEOUT
366}
367
368fn meter_passthrough_error(gateway: &Gateway, ctx: &RequestCtx, model: &str, route: Option<&str>) {
369 let mut guard = PassthroughUsageGuard::new(gateway, ctx, model, route, "vertex", "chat");
370 guard.status = "error";
371 drop(guard);
372}
373
374fn passthrough_response(
375 status: u16,
376 content_type: &str,
377 body: axum::body::Bytes,
378) -> Result<Response, GatewayError> {
379 Response::builder()
380 .status(status)
381 .header(axum::http::header::CONTENT_TYPE, content_type)
382 .body(axum::body::Body::from(body))
383 .map_err(|e| GatewayError::Upstream {
384 status: 502,
385 body: e.to_string(),
386 })
387}
388
389struct PassthroughUsageGuard {
392 ledger: crate::ledger::LedgerHandle,
393 pricing: std::sync::Arc<crate::pricing::PricingTable>,
394 provider: &'static str,
395 tenant: String,
396 attribution: crate::gateway::Attribution,
397 route: String,
398 model: String,
399 request_id: String,
400 input_tokens: u64,
401 output_tokens: u64,
402 status: &'static str,
403 op: &'static str,
404}
405
406impl PassthroughUsageGuard {
407 fn new(
408 gateway: &Gateway,
409 ctx: &RequestCtx,
410 model: &str,
411 route: Option<&str>,
412 provider: &'static str,
413 op: &'static str,
414 ) -> Self {
415 Self {
416 ledger: gateway.ledger.clone(),
417 pricing: gateway.pricing.clone(),
418 provider,
419 tenant: ctx
420 .tenant
421 .clone()
422 .unwrap_or_else(|| gateway.default_tenant.clone()),
423 attribution: gateway.attribution_of(ctx, route.unwrap_or(model)),
424 route: route.unwrap_or(model).to_string(),
425 model: model.to_string(),
426 request_id: ctx
427 .resolved_request_id()
428 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
429 input_tokens: 0,
430 output_tokens: 0,
431 status: "ok",
432 op,
433 }
434 }
435
436 fn observe_usage_metadata(&mut self, value: &serde_json::Value) {
440 if let Some(n) = value["usage"]["input_tokens"].as_u64() {
441 self.input_tokens = n;
442 }
443 if let Some(n) = value["usage"]["output_tokens"].as_u64() {
444 self.output_tokens = n;
445 }
446 if let Some(n) = value["usageMetadata"]["promptTokenCount"].as_u64() {
447 self.input_tokens = n;
448 }
449 if let Some(n) = value["usageMetadata"]["candidatesTokenCount"].as_u64() {
450 self.output_tokens = n;
451 }
452 }
453
454 fn observe_sse_line(&mut self, line: &str) {
455 if let Some(data) = line.strip_prefix("data:") {
456 if let Ok(value) = serde_json::from_str::<serde_json::Value>(data.trim()) {
457 self.observe_usage_metadata(&value);
458 }
459 }
460 }
461}
462
463impl Drop for PassthroughUsageGuard {
464 fn drop(&mut self) {
465 let cost = self.pricing.cost_usd(
466 self.provider,
467 &self.model,
468 self.input_tokens,
469 self.output_tokens,
470 );
471 self.ledger.enqueue(crate::ledger::UsageEntry {
472 ts: chrono::Utc::now(),
473 tenant: self.tenant.clone(),
474 workspace: self.attribution.workspace.clone(),
475 user: self.attribution.user.clone(),
476 thread: self.attribution.thread.clone(),
477 message: self.attribution.message.clone(),
478 route: self.route.clone(),
479 provider: self.provider.into(),
480 model: self.model.clone(),
481 lane: "passthrough".into(),
482 input_tokens: self.input_tokens,
483 output_tokens: self.output_tokens,
484 cost_usd: cost,
485 request_id: self.request_id.clone(),
486 status: self.status.to_string(),
487 op: self.op.into(),
488 user_task_type: self.attribution.user_task_type.clone(),
489 ai_task_type: self.attribution.ai_task_type.clone(),
490 });
491 }
492}
493
494struct MeteredSseStream<S> {
497 inner: S,
498 guard: PassthroughUsageGuard,
499 line_buf: String,
500}
501
502impl<S> futures::Stream for MeteredSseStream<S>
503where
504 S: futures::Stream<Item = Result<axum::body::Bytes, reqwest::Error>> + Unpin,
505{
506 type Item = Result<axum::body::Bytes, std::io::Error>;
507
508 fn poll_next(
509 self: std::pin::Pin<&mut Self>,
510 cx: &mut std::task::Context<'_>,
511 ) -> std::task::Poll<Option<Self::Item>> {
512 use futures::StreamExt;
513 let this = self.get_mut();
514 match this.inner.poll_next_unpin(cx) {
515 std::task::Poll::Ready(Some(Ok(bytes))) => {
516 this.line_buf.push_str(&String::from_utf8_lossy(&bytes));
517 while let Some(pos) = this.line_buf.find('\n') {
518 let line: String = this.line_buf.drain(..=pos).collect();
519 this.guard.observe_sse_line(line.trim_end());
520 }
521 std::task::Poll::Ready(Some(Ok(bytes)))
522 }
523 std::task::Poll::Ready(Some(Err(e))) => {
524 this.guard.status = "error";
525 std::task::Poll::Ready(Some(Err(std::io::Error::other(e.to_string()))))
526 }
527 std::task::Poll::Ready(None) => {
528 let rest = std::mem::take(&mut this.line_buf);
530 this.guard.observe_sse_line(rest.trim_end());
531 std::task::Poll::Ready(None)
532 }
533 std::task::Poll::Pending => std::task::Poll::Pending,
534 }
535 }
536}
537
538fn openai_json(c: &crate::routing::executor::Completion, request_id: &str) -> serde_json::Value {
541 let mut acc = Accumulator::default();
542 if c.tool_calls.is_empty() {
543 acc.push(StreamItem::Delta(c.content.clone()));
544 } else {
545 for (i, tc) in c.tool_calls.iter().enumerate() {
546 acc.push(StreamItem::ToolCallDelta {
547 index: i as u32,
548 id: Some(tc.id.clone()),
549 name: Some(tc.name.clone()),
550 args_fragment: tc.arguments.clone(),
551 });
552 }
553 }
554 acc.push(StreamItem::Done {
555 input_tokens: c.input_tokens,
556 output_tokens: c.output_tokens,
557 finish_reason: c.finish_reason,
558 });
559 acc.to_openai_response(request_id, &c.model)
560}
561
562fn hybrid_json(h: &crate::gateway::HybridOutcome, request_id: &str) -> serde_json::Value {
567 let content = h
568 .extraction_ran
569 .then(|| serde_json::to_string(&h.extractions).unwrap());
570 let mut message = json!({ "role": "assistant" });
571 if let Some(content) = content {
572 message["content"] = json!(content);
573 }
574 json!({
575 "id": format!("chatcmpl-{request_id}"),
576 "object": "chat.completion",
577 "created": chrono::Utc::now().timestamp(),
578 "model": h.model,
579 "jev": {
580 "answers": h.answers,
581 "survivors": h.survivors,
582 "degraded": h.degraded,
583 },
584 "choices": [{
585 "index": 0,
586 "message": message,
587 "finish_reason": "stop",
588 }],
589 "usage": {
590 "prompt_tokens": h.input_tokens,
591 "completion_tokens": h.output_tokens,
592 "total_tokens": h.input_tokens + h.output_tokens,
593 },
594 })
595}
596
597fn sse_body(
599 stream: GuardedStream,
600 request_id: String,
601) -> impl futures::Stream<Item = Result<Event, std::convert::Infallible>> {
602 use futures::StreamExt;
603 let model = stream.model().to_string();
604 stream
605 .map(move |item| match item {
606 Ok(it) => {
607 let json = stream_item_to_sse_json(&it, &request_id, &model);
608 Ok(Event::default().data(json.to_string()))
609 }
610 Err(e) => {
611 let err = json!({
612 "error": { "type": "upstream_error", "message": e.to_string(), "code": "upstream_error" }
613 });
614 Ok(Event::default().data(err.to_string()))
615 }
616 })
617 .chain(futures::stream::once(async { Ok(Event::default().data("[DONE]")) }))
618}