use std::sync::Arc;
use axum::extract::State;
use axum::http::{HeaderMap, HeaderValue};
use axum::response::sse::{Event, Sse};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::{Json, Router};
use serde_json::json;
use tap::Tap;
use crate::error::GatewayError;
use crate::gateway::{Gateway, GuardedStream, RequestCtx};
use crate::routing::jev_router::RoutingReport;
use crate::routing::request::ChatRequest;
use crate::routing::stream::{stream_item_to_sse_json, Accumulator, StreamItem};
#[derive(Clone)]
pub struct AppState {
pub gateway: Arc<Gateway>,
}
pub fn router(gateway: Arc<Gateway>) -> Router {
Router::new()
.route("/health", get(|| async { "ok" }))
.route("/v1/models", get(list_models))
.route("/v1/chat/completions", post(chat_completions))
.route("/v1/embeddings", post(embeddings))
.route("/v1beta/models/{model_action}", post(gemini_passthrough))
.route("/v1/models/{model_action}", post(gemini_passthrough))
.route("/google/models/{model_action}", post(gemini_passthrough))
.route("/typesafe/v1/systemone", post(jev_passthrough))
.with_state(AppState { gateway })
}
fn request_ctx(headers: &HeaderMap) -> RequestCtx {
let header = |name: &str| {
headers
.get(name)
.and_then(|v| v.to_str().ok())
.map(str::to_string)
};
RequestCtx {
tenant: header("x-synapse-tenant"),
workspace: header("x-synapse-workspace"),
user: header("x-synapse-user"),
thread: header("x-synapse-thread"),
message: header("x-synapse-message"),
user_task_type: header("x-synapse-user-task-type"),
ai_task_type: header("x-synapse-ai-task-type"),
request_id: None,
}
}
async fn list_models(State(st): State<AppState>) -> impl IntoResponse {
let data = st
.gateway
.model_aliases()
.into_iter()
.map(|id| json!({ "id": id, "object": "model", "owned_by": "synapse" }))
.collect::<Vec<_>>();
Json(json!({ "object": "list", "data": data }))
}
async fn chat_completions(
State(st): State<AppState>,
headers: HeaderMap,
Json(req): Json<ChatRequest>,
) -> Result<Response, GatewayError> {
let headers_ctx = request_ctx(&headers);
let request_id = headers_ctx
.resolved_request_id()
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
let ctx = RequestCtx {
request_id: Some(request_id.clone()),
..headers_ctx
};
if req.stream == Some(true) {
let stream = st.gateway.chat_stream(req, &ctx).await?;
let routing = stream.routing().clone();
return Ok(with_routing_headers(
Sse::new(sse_body(stream, request_id)).into_response(),
&routing,
));
}
let (outcome, routing) = st.gateway.chat_routed(req, &ctx).await?;
let response = match outcome {
crate::gateway::ChatOutcome::Plain(completion) => {
Json(openai_json(&completion, &request_id)).into_response()
}
crate::gateway::ChatOutcome::Hybrid(h) => {
Json(hybrid_json(&h, &request_id)).into_response()
}
};
Ok(with_routing_headers(response, &routing))
}
fn with_routing_headers(response: Response, routing: &RoutingReport) -> Response {
routing
.headers()
.into_iter()
.filter_map(|(name, value)| HeaderValue::from_str(&value).ok().map(|v| (name, v)))
.fold(response, |r, (name, v)| {
r.tap_mut(|r| {
r.headers_mut().insert(name, v);
})
})
}
async fn embeddings(
State(st): State<AppState>,
headers: HeaderMap,
Json(req): Json<crate::embeddings::EmbeddingRequest>,
) -> Result<Response, GatewayError> {
let ctx = request_ctx(&headers);
let resp = st.gateway.embed(req, ctx).await?;
Ok(Json(resp).into_response())
}
async fn gemini_passthrough(
State(st): State<AppState>,
axum::extract::Path(model_action): axum::extract::Path<String>,
axum::extract::RawQuery(query): axum::extract::RawQuery,
headers: HeaderMap,
Json(body): Json<serde_json::Value>,
) -> Result<Response, GatewayError> {
use axum::body::Body;
use axum::http::{header, StatusCode};
let provider = st
.gateway
.vertex_native
.as_ref()
.ok_or_else(|| {
GatewayError::BadRequest("gemini passthrough requires the native vertex lane".into())
})?
.clone();
let (model, action) = model_action.rsplit_once(':').ok_or_else(|| {
GatewayError::BadRequest(format!(
"expected models/{{model}}:{{action}}, got '{model_action}'"
))
})?;
let alt_sse = query.as_deref().is_some_and(|q| q.contains("alt=sse"));
let streaming = action == "streamGenerateContent";
let metered = action == "generateContent" || streaming;
if !metered {
let resp = provider
.passthrough_request(model, action, false, body, None)
.await?;
let status = resp.status();
st.gateway
.metrics
.passthrough("vertex", model, action, status.is_success());
let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
status: 502,
body: e.to_string(),
})?;
return passthrough_response(status.as_u16(), "application/json", bytes);
}
let chain = st.gateway.routes.vertex_fallback_chain(model);
let ctx = request_ctx(&headers);
let route_alias = chain.route.as_deref();
let mut prev_model: Option<String> = None;
let mut last_failure: Option<(u16, axum::body::Bytes)> = None;
for (i, leg) in chain.legs.iter().enumerate() {
if let Some(from) = prev_model.take() {
st.gateway.metrics.passthrough_fallback(&from, &leg.model);
}
let attempt = provider
.passthrough_request(
&leg.model,
action,
alt_sse && streaming,
body.clone(),
leg.region.as_deref(),
)
.await;
let resp = match attempt {
Ok(r) => r,
Err(e) => {
meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
st.gateway
.metrics
.passthrough("vertex", &leg.model, action, false);
if i + 1 < chain.legs.len() {
prev_model = Some(leg.model.clone());
continue;
}
return Err(e);
}
};
let status = resp.status();
st.gateway
.metrics
.passthrough("vertex", &leg.model, action, status.is_success());
if status.is_success() {
let mut guard = PassthroughUsageGuard::new(
&st.gateway,
&ctx,
&leg.model,
route_alias,
"vertex",
"chat",
);
if streaming && alt_sse {
let content_type = resp
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("text/event-stream")
.to_string();
let metered_stream = MeteredSseStream {
inner: resp.bytes_stream(),
guard,
line_buf: String::new(),
};
let response = Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, content_type)
.body(Body::from_stream(metered_stream))
.map_err(|e| GatewayError::Upstream {
status: 502,
body: e.to_string(),
})?;
return Ok(response);
}
let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
status: 502,
body: e.to_string(),
})?;
if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
guard.observe_usage_metadata(&value);
}
drop(guard);
return passthrough_response(status.as_u16(), "application/json", bytes);
}
let bytes = resp.bytes().await.unwrap_or_default();
meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
if passthrough_status_retryable(status) && i + 1 < chain.legs.len() {
prev_model = Some(leg.model.clone());
last_failure = Some((status.as_u16(), bytes));
continue;
}
return passthrough_response(status.as_u16(), "application/json", bytes);
}
if let Some((code, bytes)) = last_failure {
return passthrough_response(code, "application/json", bytes);
}
Err(GatewayError::Upstream {
status: 502,
body: "gemini passthrough: empty vertex fallback chain".into(),
})
}
async fn jev_passthrough(
State(st): State<AppState>,
headers: HeaderMap,
Json(mut body): Json<serde_json::Value>,
) -> Result<Response, GatewayError> {
let provider = st
.gateway
.jev_native
.as_ref()
.ok_or_else(|| {
GatewayError::BadRequest(
"jev passthrough requires TYPESAFE_API_KEY to be configured".into(),
)
})?
.clone();
if !body.is_object() {
return Err(GatewayError::BadRequest(
"expected a JSON object body".into(),
));
}
if body.get("model").is_none() {
body["model"] = serde_json::Value::from(crate::jev_native::DEFAULT_MODEL);
}
let model = body["model"]
.as_str()
.unwrap_or(crate::jev_native::DEFAULT_MODEL)
.to_string();
let ctx = request_ctx(&headers);
let mut guard =
PassthroughUsageGuard::new(&st.gateway, &ctx, &model, None, "typesafe", "systemone");
let resp = match provider.evaluate(body).await {
Ok(r) => r,
Err(e) => {
guard.status = "error";
return Err(e);
}
};
let status = resp.status();
st.gateway
.metrics
.passthrough("typesafe", &model, "systemone", status.is_success());
if status.is_success() {
let bytes = resp.bytes().await.map_err(|e| {
guard.status = "error";
GatewayError::Upstream {
status: 502,
body: e.to_string(),
}
})?;
if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
guard.observe_usage_metadata(&value);
}
return passthrough_response(status.as_u16(), "application/json", bytes);
}
let bytes = resp.bytes().await.unwrap_or_default();
guard.status = "error";
passthrough_response(status.as_u16(), "application/json", bytes)
}
fn passthrough_status_retryable(status: reqwest::StatusCode) -> bool {
status.is_server_error()
|| status == reqwest::StatusCode::TOO_MANY_REQUESTS
|| status == reqwest::StatusCode::REQUEST_TIMEOUT
}
fn meter_passthrough_error(gateway: &Gateway, ctx: &RequestCtx, model: &str, route: Option<&str>) {
let mut guard = PassthroughUsageGuard::new(gateway, ctx, model, route, "vertex", "chat");
guard.status = "error";
drop(guard);
}
fn passthrough_response(
status: u16,
content_type: &str,
body: axum::body::Bytes,
) -> Result<Response, GatewayError> {
Response::builder()
.status(status)
.header(axum::http::header::CONTENT_TYPE, content_type)
.body(axum::body::Body::from(body))
.map_err(|e| GatewayError::Upstream {
status: 502,
body: e.to_string(),
})
}
struct PassthroughUsageGuard {
ledger: crate::ledger::LedgerHandle,
pricing: std::sync::Arc<crate::pricing::PricingTable>,
provider: &'static str,
tenant: String,
attribution: crate::gateway::Attribution,
route: String,
model: String,
request_id: String,
input_tokens: u64,
output_tokens: u64,
status: &'static str,
op: &'static str,
}
impl PassthroughUsageGuard {
fn new(
gateway: &Gateway,
ctx: &RequestCtx,
model: &str,
route: Option<&str>,
provider: &'static str,
op: &'static str,
) -> Self {
Self {
ledger: gateway.ledger.clone(),
pricing: gateway.pricing.clone(),
provider,
tenant: ctx
.tenant
.clone()
.unwrap_or_else(|| gateway.default_tenant.clone()),
attribution: gateway.attribution_of(ctx, route.unwrap_or(model)),
route: route.unwrap_or(model).to_string(),
model: model.to_string(),
request_id: ctx
.resolved_request_id()
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
input_tokens: 0,
output_tokens: 0,
status: "ok",
op,
}
}
fn observe_usage_metadata(&mut self, value: &serde_json::Value) {
if let Some(n) = value["usage"]["input_tokens"].as_u64() {
self.input_tokens = n;
}
if let Some(n) = value["usage"]["output_tokens"].as_u64() {
self.output_tokens = n;
}
if let Some(n) = value["usageMetadata"]["promptTokenCount"].as_u64() {
self.input_tokens = n;
}
if let Some(n) = value["usageMetadata"]["candidatesTokenCount"].as_u64() {
self.output_tokens = n;
}
}
fn observe_sse_line(&mut self, line: &str) {
if let Some(data) = line.strip_prefix("data:") {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(data.trim()) {
self.observe_usage_metadata(&value);
}
}
}
}
impl Drop for PassthroughUsageGuard {
fn drop(&mut self) {
let cost = self.pricing.cost_usd(
self.provider,
&self.model,
self.input_tokens,
self.output_tokens,
);
self.ledger.enqueue(crate::ledger::UsageEntry {
ts: chrono::Utc::now(),
tenant: self.tenant.clone(),
workspace: self.attribution.workspace.clone(),
user: self.attribution.user.clone(),
thread: self.attribution.thread.clone(),
message: self.attribution.message.clone(),
route: self.route.clone(),
provider: self.provider.into(),
model: self.model.clone(),
lane: "passthrough".into(),
input_tokens: self.input_tokens,
output_tokens: self.output_tokens,
cost_usd: cost,
request_id: self.request_id.clone(),
status: self.status.to_string(),
op: self.op.into(),
user_task_type: self.attribution.user_task_type.clone(),
ai_task_type: self.attribution.ai_task_type.clone(),
});
}
}
struct MeteredSseStream<S> {
inner: S,
guard: PassthroughUsageGuard,
line_buf: String,
}
impl<S> futures::Stream for MeteredSseStream<S>
where
S: futures::Stream<Item = Result<axum::body::Bytes, reqwest::Error>> + Unpin,
{
type Item = Result<axum::body::Bytes, std::io::Error>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
use futures::StreamExt;
let this = self.get_mut();
match this.inner.poll_next_unpin(cx) {
std::task::Poll::Ready(Some(Ok(bytes))) => {
this.line_buf.push_str(&String::from_utf8_lossy(&bytes));
while let Some(pos) = this.line_buf.find('\n') {
let line: String = this.line_buf.drain(..=pos).collect();
this.guard.observe_sse_line(line.trim_end());
}
std::task::Poll::Ready(Some(Ok(bytes)))
}
std::task::Poll::Ready(Some(Err(e))) => {
this.guard.status = "error";
std::task::Poll::Ready(Some(Err(std::io::Error::other(e.to_string()))))
}
std::task::Poll::Ready(None) => {
let rest = std::mem::take(&mut this.line_buf);
this.guard.observe_sse_line(rest.trim_end());
std::task::Poll::Ready(None)
}
std::task::Poll::Pending => std::task::Poll::Pending,
}
}
}
fn openai_json(c: &crate::routing::executor::Completion, request_id: &str) -> serde_json::Value {
let mut acc = Accumulator::default();
if c.tool_calls.is_empty() {
acc.push(StreamItem::Delta(c.content.clone()));
} else {
for (i, tc) in c.tool_calls.iter().enumerate() {
acc.push(StreamItem::ToolCallDelta {
index: i as u32,
id: Some(tc.id.clone()),
name: Some(tc.name.clone()),
args_fragment: tc.arguments.clone(),
});
}
}
acc.push(StreamItem::Done {
input_tokens: c.input_tokens,
output_tokens: c.output_tokens,
finish_reason: c.finish_reason,
});
acc.to_openai_response(request_id, &c.model)
}
fn hybrid_json(h: &crate::gateway::HybridOutcome, request_id: &str) -> serde_json::Value {
let content = h
.extraction_ran
.then(|| serde_json::to_string(&h.extractions).unwrap());
let mut message = json!({ "role": "assistant" });
if let Some(content) = content {
message["content"] = json!(content);
}
json!({
"id": format!("chatcmpl-{request_id}"),
"object": "chat.completion",
"created": chrono::Utc::now().timestamp(),
"model": h.model,
"jev": {
"answers": h.answers,
"survivors": h.survivors,
"degraded": h.degraded,
},
"choices": [{
"index": 0,
"message": message,
"finish_reason": "stop",
}],
"usage": {
"prompt_tokens": h.input_tokens,
"completion_tokens": h.output_tokens,
"total_tokens": h.input_tokens + h.output_tokens,
},
})
}
fn sse_body(
stream: GuardedStream,
request_id: String,
) -> impl futures::Stream<Item = Result<Event, std::convert::Infallible>> {
use futures::StreamExt;
let model = stream.model().to_string();
stream
.map(move |item| match item {
Ok(it) => {
let json = stream_item_to_sse_json(&it, &request_id, &model);
Ok(Event::default().data(json.to_string()))
}
Err(e) => {
let err = json!({
"error": { "type": "upstream_error", "message": e.to_string(), "code": "upstream_error" }
});
Ok(Event::default().data(err.to_string()))
}
})
.chain(futures::stream::once(async { Ok(Event::default().data("[DONE]")) }))
}