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        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    // Share one id between the ledger row (via the gateway) and the response body.
71    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
95/// Gemini-native passthrough (`POST .../models/{model}:{action}`): forward the
96/// request to Vertex via the native provider and meter usage from the
97/// response's `usageMetadata`. Lets Gemini SDK clients (`GOOGLE_VERTEX_BASE_URL`)
98/// route through the gateway without translating to the OpenAI surface.
99async 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    // Only generation calls are metered; countTokens etc. are pure forwards.
139    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); // meter the failed call
155        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
207/// Accumulates `usageMetadata` token counts and fires exactly one ledger row on
208/// drop — every termination path (completion, error, client disconnect) meters.
209struct 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    /// Fold a Gemini response chunk's `usageMetadata` into the running totals.
242    /// Counts are cumulative per Vertex semantics, so later chunks overwrite.
243    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
286/// Tee a Vertex SSE byte stream to the client while scanning complete
287/// `data:` lines for `usageMetadata` (metered by the owned guard on drop).
288struct 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                // Flush a final unterminated data line before the guard drops.
321                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
330/// Build the OpenAI `chat.completion` JSON from a buffered `Completion`
331/// (content OR tool_calls) via an `Accumulator`.
332fn 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
354/// Render a `GuardedStream` as OpenAI SSE (`chat.completion.chunk` … `[DONE]`).
355fn 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}