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