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        // TypeSafe System One (Jev) passthrough: forwards `{state, questions}`
37        // bodies verbatim — Jev has no OpenAI-shaped equivalent.
38        .route("/typesafe/v1/systemone", post(jev_passthrough))
39        .with_state(AppState { gateway })
40}
41
42fn request_ctx(headers: &HeaderMap) -> RequestCtx {
43    let header = |name: &str| {
44        headers
45            .get(name)
46            .and_then(|v| v.to_str().ok())
47            .map(str::to_string)
48    };
49    RequestCtx {
50        tenant: header("x-synapse-tenant"),
51        workspace: header("x-synapse-workspace"),
52        user: header("x-synapse-user"),
53        thread: header("x-synapse-thread"),
54        message: header("x-synapse-message"),
55        user_task_type: header("x-synapse-user-task-type"),
56        ai_task_type: header("x-synapse-ai-task-type"),
57        request_id: None,
58    }
59}
60
61async fn list_models(State(st): State<AppState>) -> impl IntoResponse {
62    let data = st
63        .gateway
64        .model_aliases()
65        .into_iter()
66        .map(|id| json!({ "id": id, "object": "model", "owned_by": "synapse" }))
67        .collect::<Vec<_>>();
68    Json(json!({ "object": "list", "data": data }))
69}
70
71async fn chat_completions(
72    State(st): State<AppState>,
73    headers: HeaderMap,
74    Json(req): Json<ChatRequest>,
75) -> Result<Response, GatewayError> {
76    // Prefer an explicit message id for correlation when the client stamps
77    // x-synapse-message; otherwise mint a UUID shared by ledger + response body.
78    let headers_ctx = request_ctx(&headers);
79    let request_id = headers_ctx
80        .resolved_request_id()
81        .unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
82    let ctx = RequestCtx {
83        request_id: Some(request_id.clone()),
84        ..headers_ctx
85    };
86
87    if req.stream == Some(true) {
88        let stream = st.gateway.chat_stream(req, &ctx).await?;
89        return Ok(Sse::new(sse_body(stream, request_id)).into_response());
90    }
91
92    let outcome = st.gateway.chat(req, &ctx).await?;
93    match outcome {
94        crate::gateway::ChatOutcome::Plain(completion) => {
95            Ok(Json(openai_json(&completion, &request_id)).into_response())
96        }
97        crate::gateway::ChatOutcome::Hybrid(h) => {
98            Ok(Json(hybrid_json(&h, &request_id)).into_response())
99        }
100    }
101}
102
103async fn embeddings(
104    State(st): State<AppState>,
105    headers: HeaderMap,
106    Json(req): Json<crate::embeddings::EmbeddingRequest>,
107) -> Result<Response, GatewayError> {
108    let ctx = request_ctx(&headers);
109    let resp = st.gateway.embed(req, ctx).await?;
110    Ok(Json(resp).into_response())
111}
112
113/// Gemini-native passthrough (`POST .../models/{model}:{action}`): forward the
114/// request to Vertex via the native provider and meter usage from the
115/// response's `usageMetadata`. Lets Gemini SDK clients (`GOOGLE_VERTEX_BASE_URL`)
116/// route through the gateway without translating to the OpenAI surface.
117///
118/// Generation actions walk consecutive Vertex legs from `routes.toml` on
119/// 429/5xx (see `RouteTable::vertex_fallback_chain`). Non-generation actions
120/// remain a single forward.
121async fn gemini_passthrough(
122    State(st): State<AppState>,
123    axum::extract::Path(model_action): axum::extract::Path<String>,
124    axum::extract::RawQuery(query): axum::extract::RawQuery,
125    headers: HeaderMap,
126    Json(body): Json<serde_json::Value>,
127) -> Result<Response, GatewayError> {
128    use axum::body::Body;
129    use axum::http::{header, StatusCode};
130
131    let provider = st
132        .gateway
133        .vertex_native
134        .as_ref()
135        .ok_or_else(|| {
136            GatewayError::BadRequest("gemini passthrough requires the native vertex lane".into())
137        })?
138        .clone();
139
140    let (model, action) = model_action.rsplit_once(':').ok_or_else(|| {
141        GatewayError::BadRequest(format!(
142            "expected models/{{model}}:{{action}}, got '{model_action}'"
143        ))
144    })?;
145    let alt_sse = query.as_deref().is_some_and(|q| q.contains("alt=sse"));
146    let streaming = action == "streamGenerateContent";
147    let metered = action == "generateContent" || streaming;
148
149    // countTokens etc.: pure single forward, no chain / metering.
150    if !metered {
151        let resp = provider
152            .passthrough_request(model, action, false, body, None)
153            .await?;
154        let status = resp.status();
155        metrics::counter!(
156            "synapse_passthrough_total",
157            "model" => model.to_string(),
158            "action" => action.to_string(),
159            "status" => if status.is_success() { "ok" } else { "error" },
160        )
161        .increment(1);
162        let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
163            status: 502,
164            body: e.to_string(),
165        })?;
166        return passthrough_response(status.as_u16(), "application/json", bytes);
167    }
168
169    let chain = st.gateway.routes.vertex_fallback_chain(model);
170    let ctx = request_ctx(&headers);
171    let route_alias = chain.route.as_deref();
172    let mut prev_model: Option<String> = None;
173    let mut last_failure: Option<(u16, axum::body::Bytes)> = None;
174
175    for (i, leg) in chain.legs.iter().enumerate() {
176        if let Some(from) = prev_model.take() {
177            metrics::counter!(
178                "synapse_passthrough_fallback_total",
179                "from_model" => from,
180                "to_model" => leg.model.clone(),
181            )
182            .increment(1);
183        }
184
185        let attempt = provider
186            .passthrough_request(
187                &leg.model,
188                action,
189                alt_sse && streaming,
190                body.clone(),
191                leg.region.as_deref(),
192            )
193            .await;
194
195        let resp = match attempt {
196            Ok(r) => r,
197            Err(e) => {
198                meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
199                metrics::counter!(
200                    "synapse_passthrough_total",
201                    "model" => leg.model.clone(),
202                    "action" => action.to_string(),
203                    "status" => "error",
204                )
205                .increment(1);
206                if i + 1 < chain.legs.len() {
207                    prev_model = Some(leg.model.clone());
208                    continue;
209                }
210                return Err(e);
211            }
212        };
213
214        let status = resp.status();
215        metrics::counter!(
216            "synapse_passthrough_total",
217            "model" => leg.model.clone(),
218            "action" => action.to_string(),
219            "status" => if status.is_success() { "ok" } else { "error" },
220        )
221        .increment(1);
222
223        if status.is_success() {
224            let mut guard = PassthroughUsageGuard::new(
225                &st.gateway,
226                &ctx,
227                &leg.model,
228                route_alias,
229                "vertex",
230                "chat",
231            );
232            if streaming && alt_sse {
233                let content_type = resp
234                    .headers()
235                    .get(header::CONTENT_TYPE)
236                    .and_then(|v| v.to_str().ok())
237                    .unwrap_or("text/event-stream")
238                    .to_string();
239                let metered_stream = MeteredSseStream {
240                    inner: resp.bytes_stream(),
241                    guard,
242                    line_buf: String::new(),
243                };
244                let response = Response::builder()
245                    .status(StatusCode::OK)
246                    .header(header::CONTENT_TYPE, content_type)
247                    .body(Body::from_stream(metered_stream))
248                    .map_err(|e| GatewayError::Upstream {
249                        status: 502,
250                        body: e.to_string(),
251                    })?;
252                return Ok(response);
253            }
254
255            let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
256                status: 502,
257                body: e.to_string(),
258            })?;
259            if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
260                guard.observe_usage_metadata(&value);
261            }
262            drop(guard);
263            return passthrough_response(status.as_u16(), "application/json", bytes);
264        }
265
266        let bytes = resp.bytes().await.unwrap_or_default();
267        meter_passthrough_error(&st.gateway, &ctx, &leg.model, route_alias);
268
269        if passthrough_status_retryable(status) && i + 1 < chain.legs.len() {
270            prev_model = Some(leg.model.clone());
271            last_failure = Some((status.as_u16(), bytes));
272            continue;
273        }
274        return passthrough_response(status.as_u16(), "application/json", bytes);
275    }
276
277    if let Some((code, bytes)) = last_failure {
278        return passthrough_response(code, "application/json", bytes);
279    }
280    Err(GatewayError::Upstream {
281        status: 502,
282        body: "gemini passthrough: empty vertex fallback chain".into(),
283    })
284}
285
286/// TypeSafe System One (Jev) passthrough (`POST /typesafe/v1/systemone`):
287/// forwards the `{state, questions}` body verbatim to TypeSafe's API,
288/// re-authenticated with the gateway's own key, and meters from the response's
289/// `usage`. Jev answers typed questions with structured decisions and has no
290/// chat surface, so — unlike the Gemini passthrough — there is no fallback
291/// chain and no SSE variant; upstream statuses and error bodies pass through
292/// untranslated.
293async fn jev_passthrough(
294    State(st): State<AppState>,
295    headers: HeaderMap,
296    Json(mut body): Json<serde_json::Value>,
297) -> Result<Response, GatewayError> {
298    let provider = st
299        .gateway
300        .jev_native
301        .as_ref()
302        .ok_or_else(|| {
303            GatewayError::BadRequest(
304                "jev passthrough requires TYPESAFE_API_KEY to be configured".into(),
305            )
306        })?
307        .clone();
308
309    if !body.is_object() {
310        return Err(GatewayError::BadRequest(
311            "expected a JSON object body".into(),
312        ));
313    }
314    if body.get("model").is_none() {
315        body["model"] = serde_json::Value::from(crate::jev_native::DEFAULT_MODEL);
316    }
317    let model = body["model"]
318        .as_str()
319        .unwrap_or(crate::jev_native::DEFAULT_MODEL)
320        .to_string();
321
322    let ctx = request_ctx(&headers);
323    let mut guard =
324        PassthroughUsageGuard::new(&st.gateway, &ctx, &model, None, "typesafe", "systemone");
325
326    let resp = match provider.evaluate(body).await {
327        Ok(r) => r,
328        Err(e) => {
329            guard.status = "error";
330            return Err(e);
331        }
332    };
333
334    let status = resp.status();
335    metrics::counter!(
336        "synapse_passthrough_total",
337        "provider" => "typesafe",
338        "model" => model.clone(),
339        "action" => "systemone",
340        "status" => if status.is_success() { "ok" } else { "error" },
341    )
342    .increment(1);
343
344    if status.is_success() {
345        let bytes = resp.bytes().await.map_err(|e| {
346            guard.status = "error";
347            GatewayError::Upstream {
348                status: 502,
349                body: e.to_string(),
350            }
351        })?;
352        if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) {
353            guard.observe_usage_metadata(&value);
354        }
355        return passthrough_response(status.as_u16(), "application/json", bytes);
356    }
357
358    let bytes = resp.bytes().await.unwrap_or_default();
359    guard.status = "error";
360    passthrough_response(status.as_u16(), "application/json", bytes)
361}
362
363fn passthrough_status_retryable(status: reqwest::StatusCode) -> bool {
364    status.is_server_error()
365        || status == reqwest::StatusCode::TOO_MANY_REQUESTS
366        || status == reqwest::StatusCode::REQUEST_TIMEOUT
367}
368
369fn meter_passthrough_error(gateway: &Gateway, ctx: &RequestCtx, model: &str, route: Option<&str>) {
370    let mut guard = PassthroughUsageGuard::new(gateway, ctx, model, route, "vertex", "chat");
371    guard.status = "error";
372    drop(guard);
373}
374
375fn passthrough_response(
376    status: u16,
377    content_type: &str,
378    body: axum::body::Bytes,
379) -> Result<Response, GatewayError> {
380    Response::builder()
381        .status(status)
382        .header(axum::http::header::CONTENT_TYPE, content_type)
383        .body(axum::body::Body::from(body))
384        .map_err(|e| GatewayError::Upstream {
385            status: 502,
386            body: e.to_string(),
387        })
388}
389
390/// Accumulates response usage counts and fires exactly one ledger row on
391/// drop — every termination path (completion, error, client disconnect) meters.
392struct PassthroughUsageGuard {
393    ledger: crate::ledger::LedgerHandle,
394    pricing: std::sync::Arc<crate::pricing::PricingTable>,
395    provider: &'static str,
396    tenant: String,
397    attribution: crate::gateway::Attribution,
398    route: String,
399    model: String,
400    request_id: String,
401    input_tokens: u64,
402    output_tokens: u64,
403    status: &'static str,
404    op: &'static str,
405}
406
407impl PassthroughUsageGuard {
408    fn new(
409        gateway: &Gateway,
410        ctx: &RequestCtx,
411        model: &str,
412        route: Option<&str>,
413        provider: &'static str,
414        op: &'static str,
415    ) -> Self {
416        Self {
417            ledger: gateway.ledger.clone(),
418            pricing: gateway.pricing.clone(),
419            provider,
420            tenant: ctx
421                .tenant
422                .clone()
423                .unwrap_or_else(|| gateway.default_tenant.clone()),
424            attribution: gateway.attribution_of(ctx, route.unwrap_or(model)),
425            route: route.unwrap_or(model).to_string(),
426            model: model.to_string(),
427            request_id: ctx
428                .resolved_request_id()
429                .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
430            input_tokens: 0,
431            output_tokens: 0,
432            status: "ok",
433            op,
434        }
435    }
436
437    /// Fold a response's usage into the running totals, overwriting earlier
438    /// counts. Both metered shapes are understood: Vertex `usageMetadata`
439    /// (counts are cumulative per chunk) and TypeSafe `usage` (single shot).
440    fn observe_usage_metadata(&mut self, value: &serde_json::Value) {
441        if let Some(n) = value["usage"]["input_tokens"].as_u64() {
442            self.input_tokens = n;
443        }
444        if let Some(n) = value["usage"]["output_tokens"].as_u64() {
445            self.output_tokens = n;
446        }
447        if let Some(n) = value["usageMetadata"]["promptTokenCount"].as_u64() {
448            self.input_tokens = n;
449        }
450        if let Some(n) = value["usageMetadata"]["candidatesTokenCount"].as_u64() {
451            self.output_tokens = n;
452        }
453    }
454
455    fn observe_sse_line(&mut self, line: &str) {
456        if let Some(data) = line.strip_prefix("data:") {
457            if let Ok(value) = serde_json::from_str::<serde_json::Value>(data.trim()) {
458                self.observe_usage_metadata(&value);
459            }
460        }
461    }
462}
463
464impl Drop for PassthroughUsageGuard {
465    fn drop(&mut self) {
466        let cost = self.pricing.cost_usd(
467            self.provider,
468            &self.model,
469            self.input_tokens,
470            self.output_tokens,
471        );
472        self.ledger.enqueue(crate::ledger::UsageEntry {
473            ts: chrono::Utc::now(),
474            tenant: self.tenant.clone(),
475            workspace: self.attribution.workspace.clone(),
476            user: self.attribution.user.clone(),
477            thread: self.attribution.thread.clone(),
478            message: self.attribution.message.clone(),
479            route: self.route.clone(),
480            provider: self.provider.into(),
481            model: self.model.clone(),
482            lane: "passthrough".into(),
483            input_tokens: self.input_tokens,
484            output_tokens: self.output_tokens,
485            cost_usd: cost,
486            request_id: self.request_id.clone(),
487            status: self.status.to_string(),
488            op: self.op.into(),
489            user_task_type: self.attribution.user_task_type.clone(),
490            ai_task_type: self.attribution.ai_task_type.clone(),
491        });
492    }
493}
494
495/// Tee a Vertex SSE byte stream to the client while scanning complete
496/// `data:` lines for `usageMetadata` (metered by the owned guard on drop).
497struct MeteredSseStream<S> {
498    inner: S,
499    guard: PassthroughUsageGuard,
500    line_buf: String,
501}
502
503impl<S> futures::Stream for MeteredSseStream<S>
504where
505    S: futures::Stream<Item = Result<axum::body::Bytes, reqwest::Error>> + Unpin,
506{
507    type Item = Result<axum::body::Bytes, std::io::Error>;
508
509    fn poll_next(
510        self: std::pin::Pin<&mut Self>,
511        cx: &mut std::task::Context<'_>,
512    ) -> std::task::Poll<Option<Self::Item>> {
513        use futures::StreamExt;
514        let this = self.get_mut();
515        match this.inner.poll_next_unpin(cx) {
516            std::task::Poll::Ready(Some(Ok(bytes))) => {
517                this.line_buf.push_str(&String::from_utf8_lossy(&bytes));
518                while let Some(pos) = this.line_buf.find('\n') {
519                    let line: String = this.line_buf.drain(..=pos).collect();
520                    this.guard.observe_sse_line(line.trim_end());
521                }
522                std::task::Poll::Ready(Some(Ok(bytes)))
523            }
524            std::task::Poll::Ready(Some(Err(e))) => {
525                this.guard.status = "error";
526                std::task::Poll::Ready(Some(Err(std::io::Error::other(e.to_string()))))
527            }
528            std::task::Poll::Ready(None) => {
529                // Flush a final unterminated data line before the guard drops.
530                let rest = std::mem::take(&mut this.line_buf);
531                this.guard.observe_sse_line(rest.trim_end());
532                std::task::Poll::Ready(None)
533            }
534            std::task::Poll::Pending => std::task::Poll::Pending,
535        }
536    }
537}
538
539/// Build the OpenAI `chat.completion` JSON from a buffered `Completion`
540/// (content OR tool_calls) via an `Accumulator`.
541fn openai_json(c: &crate::routing::executor::Completion, request_id: &str) -> serde_json::Value {
542    let mut acc = Accumulator::default();
543    if c.tool_calls.is_empty() {
544        acc.push(StreamItem::Delta(c.content.clone()));
545    } else {
546        for (i, tc) in c.tool_calls.iter().enumerate() {
547            acc.push(StreamItem::ToolCallDelta {
548                index: i as u32,
549                id: Some(tc.id.clone()),
550                name: Some(tc.name.clone()),
551                args_fragment: tc.arguments.clone(),
552            });
553        }
554    }
555    acc.push(StreamItem::Done {
556        input_tokens: c.input_tokens,
557        output_tokens: c.output_tokens,
558        finish_reason: c.finish_reason,
559    });
560    acc.to_openai_response(request_id, &c.model)
561}
562
563/// Render a Jev hybrid-extraction outcome as the stable envelope from the
564/// spec: a `jev` block beside the OpenAI-shaped choices/usage, `content` a
565/// JSON map of candidate key → extraction result (absent when no survivor
566/// was attempted).
567fn hybrid_json(h: &crate::gateway::HybridOutcome, request_id: &str) -> serde_json::Value {
568    let content = h
569        .extraction_ran
570        .then(|| serde_json::to_string(&h.extractions).unwrap());
571    let mut message = json!({ "role": "assistant" });
572    if let Some(content) = content {
573        message["content"] = json!(content);
574    }
575    json!({
576        "id": format!("chatcmpl-{request_id}"),
577        "object": "chat.completion",
578        "created": chrono::Utc::now().timestamp(),
579        "model": h.model,
580        "jev": {
581            "answers": h.answers,
582            "survivors": h.survivors,
583            "degraded": h.degraded,
584        },
585        "choices": [{
586            "index": 0,
587            "message": message,
588            "finish_reason": "stop",
589        }],
590        "usage": {
591            "prompt_tokens": h.input_tokens,
592            "completion_tokens": h.output_tokens,
593            "total_tokens": h.input_tokens + h.output_tokens,
594        },
595    })
596}
597
598/// Render a `GuardedStream` as OpenAI SSE (`chat.completion.chunk` … `[DONE]`).
599fn sse_body(
600    stream: GuardedStream,
601    request_id: String,
602) -> impl futures::Stream<Item = Result<Event, std::convert::Infallible>> {
603    use futures::StreamExt;
604    let model = stream.model().to_string();
605    stream
606        .map(move |item| match item {
607            Ok(it) => {
608                let json = stream_item_to_sse_json(&it, &request_id, &model);
609                Ok(Event::default().data(json.to_string()))
610            }
611            Err(e) => {
612                let err = json!({
613                    "error": { "type": "upstream_error", "message": e.to_string(), "code": "upstream_error" }
614                });
615                Ok(Event::default().data(err.to_string()))
616            }
617        })
618        .chain(futures::stream::once(async { Ok(Event::default().data("[DONE]")) }))
619}