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