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