Skip to main content

synapse/
server.rs

1//! axum surface: a thin HTTP layer that delegates to the in-process `Gateway`.
2
3use 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        // Gemini-native passthrough: `@google/generative-ai` / `@google/genai`
30        // clients pointed at the gateway via GOOGLE_VERTEX_BASE_URL. The three
31        // prefixes cover the SDK `apiVersion` values in use (default v1beta,
32        // Vertex-style `google`, and v1).
33        .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    // Prefer an explicit message id for correlation when the client stamps
72    // x-synapse-message; otherwise mint a UUID shared by ledger + response body.
73    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
101/// Gemini-native passthrough (`POST .../models/{model}:{action}`): forward the
102/// request to Vertex via the native provider and meter usage from the
103/// response's `usageMetadata`. Lets Gemini SDK clients (`GOOGLE_VERTEX_BASE_URL`)
104/// route through the gateway without translating to the OpenAI surface.
105///
106/// Generation actions walk consecutive Vertex legs from `routes.toml` on
107/// 429/5xx (see `RouteTable::vertex_fallback_chain`). Non-generation actions
108/// remain a single forward.
109async fn gemini_passthrough(
110    State(st): State<AppState>,
111    axum::extract::Path(model_action): axum::extract::Path<String>,
112    axum::extract::RawQuery(query): axum::extract::RawQuery,
113    headers: HeaderMap,
114    Json(body): Json<serde_json::Value>,
115) -> Result<Response, GatewayError> {
116    use axum::body::Body;
117    use axum::http::{header, StatusCode};
118
119    let provider = st
120        .gateway
121        .vertex_native
122        .as_ref()
123        .ok_or_else(|| {
124            GatewayError::BadRequest("gemini passthrough requires the native vertex lane".into())
125        })?
126        .clone();
127
128    let (model, action) = model_action.rsplit_once(':').ok_or_else(|| {
129        GatewayError::BadRequest(format!(
130            "expected models/{{model}}:{{action}}, got '{model_action}'"
131        ))
132    })?;
133    let alt_sse = query.as_deref().is_some_and(|q| q.contains("alt=sse"));
134    let streaming = action == "streamGenerateContent";
135    let metered = action == "generateContent" || streaming;
136
137    // countTokens etc.: pure single forward, no chain / metering.
138    if !metered {
139        let resp = provider
140            .passthrough_request(model, action, false, body, None)
141            .await?;
142        let status = resp.status();
143        metrics::counter!(
144            "synapse_passthrough_total",
145            "model" => model.to_string(),
146            "action" => action.to_string(),
147            "status" => if status.is_success() { "ok" } else { "error" },
148        )
149        .increment(1);
150        let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
151            status: 502,
152            body: e.to_string(),
153        })?;
154        return passthrough_response(status.as_u16(), "application/json", bytes);
155    }
156
157    let chain = st.gateway.routes.vertex_fallback_chain(model);
158    let ctx = request_ctx(&headers);
159    let route_alias = chain.route.as_deref();
160    let mut prev_model: Option<String> = None;
161    let mut last_failure: Option<(u16, axum::body::Bytes)> = None;
162
163    for (i, leg) in chain.legs.iter().enumerate() {
164        if let Some(from) = prev_model.take() {
165            metrics::counter!(
166                "synapse_passthrough_fallback_total",
167                "from_model" => from,
168                "to_model" => leg.model.clone(),
169            )
170            .increment(1);
171        }
172
173        let attempt = provider
174            .passthrough_request(
175                &leg.model,
176                action,
177                alt_sse && streaming,
178                body.clone(),
179                leg.region.as_deref(),
180            )
181            .await;
182
183        let resp = match attempt {
184            Ok(r) => r,
185            Err(e) => {
186                meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
187                metrics::counter!(
188                    "synapse_passthrough_total",
189                    "model" => leg.model.clone(),
190                    "action" => action.to_string(),
191                    "status" => "error",
192                )
193                .increment(1);
194                if i + 1 < chain.legs.len() {
195                    prev_model = Some(leg.model.clone());
196                    continue;
197                }
198                return Err(e);
199            }
200        };
201
202        let status = resp.status();
203        metrics::counter!(
204            "synapse_passthrough_total",
205            "model" => leg.model.clone(),
206            "action" => action.to_string(),
207            "status" => if status.is_success() { "ok" } else { "error" },
208        )
209        .increment(1);
210
211        if status.is_success() {
212            let mut guard = PassthroughUsageGuard::new(&st.gateway, &ctx, &leg.model, route_alias);
213            if streaming && alt_sse {
214                let content_type = resp
215                    .headers()
216                    .get(header::CONTENT_TYPE)
217                    .and_then(|v| v.to_str().ok())
218                    .unwrap_or("text/event-stream")
219                    .to_string();
220                let metered_stream = MeteredSseStream {
221                    inner: resp.bytes_stream(),
222                    guard,
223                    line_buf: String::new(),
224                };
225                let response = Response::builder()
226                    .status(StatusCode::OK)
227                    .header(header::CONTENT_TYPE, content_type)
228                    .body(Body::from_stream(metered_stream))
229                    .map_err(|e| GatewayError::Upstream {
230                        status: 502,
231                        body: e.to_string(),
232                    })?;
233                return Ok(response);
234            }
235
236            let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
237                status: 502,
238                body: e.to_string(),
239            })?;
240            if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
241                guard.observe_usage_metadata(&value);
242            }
243            drop(guard);
244            return passthrough_response(status.as_u16(), "application/json", bytes);
245        }
246
247        let bytes = resp.bytes().await.unwrap_or_default();
248        meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
249
250        if passthrough_status_retryable(status) && i + 1 < chain.legs.len() {
251            prev_model = Some(leg.model.clone());
252            last_failure = Some((status.as_u16(), bytes));
253            continue;
254        }
255        return passthrough_response(status.as_u16(), "application/json", bytes);
256    }
257
258    if let Some((code, bytes)) = last_failure {
259        return passthrough_response(code, "application/json", bytes);
260    }
261    Err(GatewayError::Upstream {
262        status: 502,
263        body: "gemini passthrough: empty vertex fallback chain".into(),
264    })
265}
266
267fn passthrough_status_retryable(status: reqwest::StatusCode) -> bool {
268    status.is_server_error()
269        || status == reqwest::StatusCode::TOO_MANY_REQUESTS
270        || status == reqwest::StatusCode::REQUEST_TIMEOUT
271}
272
273fn meter_passthrough_error(gateway: &Gateway, ctx: &RequestCtx, model: &str, route: Option<&str>) {
274    let mut guard = PassthroughUsageGuard::new(gateway, ctx, model, route);
275    guard.status = "error";
276    drop(guard);
277}
278
279fn passthrough_response(
280    status: u16,
281    content_type: &str,
282    body: axum::body::Bytes,
283) -> Result<Response, GatewayError> {
284    Response::builder()
285        .status(status)
286        .header(axum::http::header::CONTENT_TYPE, content_type)
287        .body(axum::body::Body::from(body))
288        .map_err(|e| GatewayError::Upstream {
289            status: 502,
290            body: e.to_string(),
291        })
292}
293
294/// Accumulates `usageMetadata` token counts and fires exactly one ledger row on
295/// drop — every termination path (completion, error, client disconnect) meters.
296struct PassthroughUsageGuard {
297    ledger: crate::ledger::LedgerHandle,
298    pricing: std::sync::Arc<crate::pricing::PricingTable>,
299    tenant: String,
300    workspace: Option<String>,
301    user: Option<String>,
302    thread: Option<String>,
303    message: Option<String>,
304    route: String,
305    model: String,
306    request_id: String,
307    input_tokens: u64,
308    output_tokens: u64,
309    status: &'static str,
310}
311
312impl PassthroughUsageGuard {
313    fn new(gateway: &Gateway, ctx: &RequestCtx, model: &str, route: Option<&str>) -> Self {
314        Self {
315            ledger: gateway.ledger.clone(),
316            pricing: gateway.pricing.clone(),
317            tenant: ctx
318                .tenant
319                .clone()
320                .unwrap_or_else(|| gateway.default_tenant.clone()),
321            workspace: ctx.workspace.clone(),
322            user: ctx.user.clone(),
323            thread: ctx.thread.clone(),
324            message: ctx.message.clone(),
325            route: route.unwrap_or(model).to_string(),
326            model: model.to_string(),
327            request_id: ctx
328                .resolved_request_id()
329                .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
330            input_tokens: 0,
331            output_tokens: 0,
332            status: "ok",
333        }
334    }
335
336    /// Fold a Gemini response chunk's `usageMetadata` into the running totals.
337    /// Counts are cumulative per Vertex semantics, so later chunks overwrite.
338    fn observe_usage_metadata(&mut self, value: &serde_json::Value) {
339        let usage = &value["usageMetadata"];
340        if let Some(n) = usage["promptTokenCount"].as_u64() {
341            self.input_tokens = n;
342        }
343        if let Some(n) = usage["candidatesTokenCount"].as_u64() {
344            self.output_tokens = n;
345        }
346    }
347
348    fn observe_sse_line(&mut self, line: &str) {
349        if let Some(data) = line.strip_prefix("data:") {
350            if let Ok(value) = serde_json::from_str::<serde_json::Value>(data.trim()) {
351                self.observe_usage_metadata(&value);
352            }
353        }
354    }
355}
356
357impl Drop for PassthroughUsageGuard {
358    fn drop(&mut self) {
359        let cost =
360            self.pricing
361                .cost_usd("vertex", &self.model, self.input_tokens, self.output_tokens);
362        self.ledger.enqueue(crate::ledger::UsageEntry {
363            ts: chrono::Utc::now(),
364            tenant: self.tenant.clone(),
365            workspace: self.workspace.clone(),
366            user: self.user.clone(),
367            thread: self.thread.clone(),
368            message: self.message.clone(),
369            route: self.route.clone(),
370            provider: "vertex".into(),
371            model: self.model.clone(),
372            lane: "passthrough".into(),
373            input_tokens: self.input_tokens,
374            output_tokens: self.output_tokens,
375            cost_usd: cost,
376            request_id: self.request_id.clone(),
377            status: self.status.to_string(),
378            op: "chat".into(),
379        });
380    }
381}
382
383/// Tee a Vertex SSE byte stream to the client while scanning complete
384/// `data:` lines for `usageMetadata` (metered by the owned guard on drop).
385struct MeteredSseStream<S> {
386    inner: S,
387    guard: PassthroughUsageGuard,
388    line_buf: String,
389}
390
391impl<S> futures::Stream for MeteredSseStream<S>
392where
393    S: futures::Stream<Item = Result<axum::body::Bytes, reqwest::Error>> + Unpin,
394{
395    type Item = Result<axum::body::Bytes, std::io::Error>;
396
397    fn poll_next(
398        self: std::pin::Pin<&mut Self>,
399        cx: &mut std::task::Context<'_>,
400    ) -> std::task::Poll<Option<Self::Item>> {
401        use futures::StreamExt;
402        let this = self.get_mut();
403        match this.inner.poll_next_unpin(cx) {
404            std::task::Poll::Ready(Some(Ok(bytes))) => {
405                this.line_buf.push_str(&String::from_utf8_lossy(&bytes));
406                while let Some(pos) = this.line_buf.find('\n') {
407                    let line: String = this.line_buf.drain(..=pos).collect();
408                    this.guard.observe_sse_line(line.trim_end());
409                }
410                std::task::Poll::Ready(Some(Ok(bytes)))
411            }
412            std::task::Poll::Ready(Some(Err(e))) => {
413                this.guard.status = "error";
414                std::task::Poll::Ready(Some(Err(std::io::Error::other(e.to_string()))))
415            }
416            std::task::Poll::Ready(None) => {
417                // Flush a final unterminated data line before the guard drops.
418                let rest = std::mem::take(&mut this.line_buf);
419                this.guard.observe_sse_line(rest.trim_end());
420                std::task::Poll::Ready(None)
421            }
422            std::task::Poll::Pending => std::task::Poll::Pending,
423        }
424    }
425}
426
427/// Build the OpenAI `chat.completion` JSON from a buffered `Completion`
428/// (content OR tool_calls) via an `Accumulator`.
429fn openai_json(c: &crate::routing::executor::Completion, request_id: &str) -> serde_json::Value {
430    let mut acc = Accumulator::default();
431    if c.tool_calls.is_empty() {
432        acc.push(StreamItem::Delta(c.content.clone()));
433    } else {
434        for (i, tc) in c.tool_calls.iter().enumerate() {
435            acc.push(StreamItem::ToolCallDelta {
436                index: i as u32,
437                id: Some(tc.id.clone()),
438                name: Some(tc.name.clone()),
439                args_fragment: tc.arguments.clone(),
440            });
441        }
442    }
443    acc.push(StreamItem::Done {
444        input_tokens: c.input_tokens,
445        output_tokens: c.output_tokens,
446        finish_reason: c.finish_reason,
447    });
448    acc.to_openai_response(request_id, &c.model)
449}
450
451/// Render a `GuardedStream` as OpenAI SSE (`chat.completion.chunk` … `[DONE]`).
452fn sse_body(
453    stream: GuardedStream,
454    request_id: String,
455) -> impl futures::Stream<Item = Result<Event, std::convert::Infallible>> {
456    use futures::StreamExt;
457    let model = stream.model().to_string();
458    stream
459        .map(move |item| match item {
460            Ok(it) => {
461                let json = stream_item_to_sse_json(&it, &request_id, &model);
462                Ok(Event::default().data(json.to_string()))
463            }
464            Err(e) => {
465                let err = json!({
466                    "error": { "type": "upstream_error", "message": e.to_string(), "code": "upstream_error" }
467                });
468                Ok(Event::default().data(err.to_string()))
469            }
470        })
471        .chain(futures::stream::once(async { Ok(Event::default().data("[DONE]")) }))
472}