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 .with_state(AppState { gateway })
37}
38
39fn request_ctx(headers: &HeaderMap) -> RequestCtx {
40 let header = |name: &str| {
41 headers
42 .get(name)
43 .and_then(|v| v.to_str().ok())
44 .map(str::to_string)
45 };
46 RequestCtx {
47 tenant: header("x-synapse-tenant"),
48 workspace: header("x-synapse-workspace"),
49 user: header("x-synapse-user"),
50 thread: header("x-synapse-thread"),
51 message: header("x-synapse-message"),
52 request_id: None,
53 }
54}
55
56async fn list_models(State(st): State<AppState>) -> impl IntoResponse {
57 let data = st
58 .gateway
59 .model_aliases()
60 .into_iter()
61 .map(|id| json!({ "id": id, "object": "model", "owned_by": "synapse" }))
62 .collect::<Vec<_>>();
63 Json(json!({ "object": "list", "data": data }))
64}
65
66async fn chat_completions(
67 State(st): State<AppState>,
68 headers: HeaderMap,
69 Json(req): Json<ChatRequest>,
70) -> Result<Response, GatewayError> {
71 let headers_ctx = request_ctx(&headers);
74 let request_id = headers_ctx
75 .resolved_request_id()
76 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
77 let ctx = RequestCtx {
78 request_id: Some(request_id.clone()),
79 ..headers_ctx
80 };
81
82 if req.stream == Some(true) {
83 let stream = st.gateway.chat_stream(req, &ctx).await?;
84 return Ok(Sse::new(sse_body(stream, request_id)).into_response());
85 }
86
87 let completion = st.gateway.chat(req, &ctx).await?;
88 Ok(Json(openai_json(&completion, &request_id)).into_response())
89}
90
91async fn embeddings(
92 State(st): State<AppState>,
93 headers: HeaderMap,
94 Json(req): Json<crate::embeddings::EmbeddingRequest>,
95) -> Result<Response, GatewayError> {
96 let ctx = request_ctx(&headers);
97 let resp = st.gateway.embed(req, ctx).await?;
98 Ok(Json(resp).into_response())
99}
100
101async fn gemini_passthrough(
106 State(st): State<AppState>,
107 axum::extract::Path(model_action): axum::extract::Path<String>,
108 axum::extract::RawQuery(query): axum::extract::RawQuery,
109 headers: HeaderMap,
110 Json(body): Json<serde_json::Value>,
111) -> Result<Response, GatewayError> {
112 use axum::body::Body;
113 use axum::http::{header, StatusCode};
114
115 let provider = st
116 .gateway
117 .vertex_native
118 .as_ref()
119 .ok_or_else(|| {
120 GatewayError::BadRequest("gemini passthrough requires the native vertex lane".into())
121 })?
122 .clone();
123
124 let (model, action) = model_action.rsplit_once(':').ok_or_else(|| {
125 GatewayError::BadRequest(format!(
126 "expected models/{{model}}:{{action}}, got '{model_action}'"
127 ))
128 })?;
129 let alt_sse = query.as_deref().is_some_and(|q| q.contains("alt=sse"));
130 let streaming = action == "streamGenerateContent";
131
132 let resp = provider
133 .passthrough_request(model, action, alt_sse && streaming, body)
134 .await?;
135 let status = resp.status();
136 metrics::counter!(
137 "synapse_passthrough_total",
138 "model" => model.to_string(),
139 "action" => action.to_string(),
140 "status" => if status.is_success() { "ok" } else { "error" },
141 )
142 .increment(1);
143
144 let metered = action == "generateContent" || streaming;
146 if !metered {
147 let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
148 status: 502,
149 body: e.to_string(),
150 })?;
151 return passthrough_response(status.as_u16(), "application/json", bytes);
152 }
153
154 let ctx = request_ctx(&headers);
155 let mut guard = PassthroughUsageGuard::new(&st.gateway, &ctx, model);
156
157 if !status.is_success() {
158 let bytes = resp.bytes().await.unwrap_or_default();
159 guard.status = "error";
160 drop(guard); return passthrough_response(status.as_u16(), "application/json", bytes);
162 }
163
164 if streaming && alt_sse {
165 let content_type = resp
166 .headers()
167 .get(header::CONTENT_TYPE)
168 .and_then(|v| v.to_str().ok())
169 .unwrap_or("text/event-stream")
170 .to_string();
171 let metered_stream = MeteredSseStream {
172 inner: resp.bytes_stream(),
173 guard,
174 line_buf: String::new(),
175 };
176 let response = Response::builder()
177 .status(StatusCode::OK)
178 .header(header::CONTENT_TYPE, content_type)
179 .body(Body::from_stream(metered_stream))
180 .map_err(|e| GatewayError::Upstream {
181 status: 502,
182 body: e.to_string(),
183 })?;
184 return Ok(response);
185 }
186
187 let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
188 status: 502,
189 body: e.to_string(),
190 })?;
191 if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
192 guard.observe_usage_metadata(&value);
193 }
194 drop(guard);
195 passthrough_response(status.as_u16(), "application/json", bytes)
196}
197
198fn passthrough_response(
199 status: u16,
200 content_type: &str,
201 body: axum::body::Bytes,
202) -> Result<Response, GatewayError> {
203 Response::builder()
204 .status(status)
205 .header(axum::http::header::CONTENT_TYPE, content_type)
206 .body(axum::body::Body::from(body))
207 .map_err(|e| GatewayError::Upstream {
208 status: 502,
209 body: e.to_string(),
210 })
211}
212
213struct PassthroughUsageGuard {
216 ledger: crate::ledger::LedgerHandle,
217 pricing: std::sync::Arc<crate::pricing::PricingTable>,
218 tenant: String,
219 workspace: Option<String>,
220 user: Option<String>,
221 thread: Option<String>,
222 message: Option<String>,
223 model: String,
224 request_id: String,
225 input_tokens: u64,
226 output_tokens: u64,
227 status: &'static str,
228}
229
230impl PassthroughUsageGuard {
231 fn new(gateway: &Gateway, ctx: &RequestCtx, model: &str) -> Self {
232 Self {
233 ledger: gateway.ledger.clone(),
234 pricing: gateway.pricing.clone(),
235 tenant: ctx
236 .tenant
237 .clone()
238 .unwrap_or_else(|| gateway.default_tenant.clone()),
239 workspace: ctx.workspace.clone(),
240 user: ctx.user.clone(),
241 thread: ctx.thread.clone(),
242 message: ctx.message.clone(),
243 model: model.to_string(),
244 request_id: ctx
245 .resolved_request_id()
246 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
247 input_tokens: 0,
248 output_tokens: 0,
249 status: "ok",
250 }
251 }
252
253 fn observe_usage_metadata(&mut self, value: &serde_json::Value) {
256 let usage = &value["usageMetadata"];
257 if let Some(n) = usage["promptTokenCount"].as_u64() {
258 self.input_tokens = n;
259 }
260 if let Some(n) = usage["candidatesTokenCount"].as_u64() {
261 self.output_tokens = n;
262 }
263 }
264
265 fn observe_sse_line(&mut self, line: &str) {
266 if let Some(data) = line.strip_prefix("data:") {
267 if let Ok(value) = serde_json::from_str::<serde_json::Value>(data.trim()) {
268 self.observe_usage_metadata(&value);
269 }
270 }
271 }
272}
273
274impl Drop for PassthroughUsageGuard {
275 fn drop(&mut self) {
276 let cost =
277 self.pricing
278 .cost_usd("vertex", &self.model, self.input_tokens, self.output_tokens);
279 self.ledger.enqueue(crate::ledger::UsageEntry {
280 ts: chrono::Utc::now(),
281 tenant: self.tenant.clone(),
282 workspace: self.workspace.clone(),
283 user: self.user.clone(),
284 thread: self.thread.clone(),
285 message: self.message.clone(),
286 route: self.model.clone(),
287 provider: "vertex".into(),
288 model: self.model.clone(),
289 lane: "passthrough".into(),
290 input_tokens: self.input_tokens,
291 output_tokens: self.output_tokens,
292 cost_usd: cost,
293 request_id: self.request_id.clone(),
294 status: self.status.to_string(),
295 op: "chat".into(),
296 });
297 }
298}
299
300struct MeteredSseStream<S> {
303 inner: S,
304 guard: PassthroughUsageGuard,
305 line_buf: String,
306}
307
308impl<S> futures::Stream for MeteredSseStream<S>
309where
310 S: futures::Stream<Item = Result<axum::body::Bytes, reqwest::Error>> + Unpin,
311{
312 type Item = Result<axum::body::Bytes, std::io::Error>;
313
314 fn poll_next(
315 self: std::pin::Pin<&mut Self>,
316 cx: &mut std::task::Context<'_>,
317 ) -> std::task::Poll<Option<Self::Item>> {
318 use futures::StreamExt;
319 let this = self.get_mut();
320 match this.inner.poll_next_unpin(cx) {
321 std::task::Poll::Ready(Some(Ok(bytes))) => {
322 this.line_buf.push_str(&String::from_utf8_lossy(&bytes));
323 while let Some(pos) = this.line_buf.find('\n') {
324 let line: String = this.line_buf.drain(..=pos).collect();
325 this.guard.observe_sse_line(line.trim_end());
326 }
327 std::task::Poll::Ready(Some(Ok(bytes)))
328 }
329 std::task::Poll::Ready(Some(Err(e))) => {
330 this.guard.status = "error";
331 std::task::Poll::Ready(Some(Err(std::io::Error::other(e.to_string()))))
332 }
333 std::task::Poll::Ready(None) => {
334 let rest = std::mem::take(&mut this.line_buf);
336 this.guard.observe_sse_line(rest.trim_end());
337 std::task::Poll::Ready(None)
338 }
339 std::task::Poll::Pending => std::task::Poll::Pending,
340 }
341 }
342}
343
344fn openai_json(c: &crate::routing::executor::Completion, request_id: &str) -> serde_json::Value {
347 let mut acc = Accumulator::default();
348 if c.tool_calls.is_empty() {
349 acc.push(StreamItem::Delta(c.content.clone()));
350 } else {
351 for (i, tc) in c.tool_calls.iter().enumerate() {
352 acc.push(StreamItem::ToolCallDelta {
353 index: i as u32,
354 id: Some(tc.id.clone()),
355 name: Some(tc.name.clone()),
356 args_fragment: tc.arguments.clone(),
357 });
358 }
359 }
360 acc.push(StreamItem::Done {
361 input_tokens: c.input_tokens,
362 output_tokens: c.output_tokens,
363 finish_reason: c.finish_reason,
364 });
365 acc.to_openai_response(request_id, &c.model)
366}
367
368fn sse_body(
370 stream: GuardedStream,
371 request_id: String,
372) -> impl futures::Stream<Item = Result<Event, std::convert::Infallible>> {
373 use futures::StreamExt;
374 let model = stream.model().to_string();
375 stream
376 .map(move |item| match item {
377 Ok(it) => {
378 let json = stream_item_to_sse_json(&it, &request_id, &model);
379 Ok(Event::default().data(json.to_string()))
380 }
381 Err(e) => {
382 let err = json!({
383 "error": { "type": "upstream_error", "message": e.to_string(), "code": "upstream_error" }
384 });
385 Ok(Event::default().data(err.to_string()))
386 }
387 })
388 .chain(futures::stream::once(async { Ok(Event::default().data("[DONE]")) }))
389}