Skip to main content

synapse/
gateway.rs

1//! In-process LLM gateway: routing, fallback, native Vertex, ledger, metrics.
2//! Transport-independent core; the axum HTTP layer (`server`) delegates here.
3
4use std::pin::Pin;
5use std::sync::Arc;
6use std::task::{Context, Poll};
7use std::time::Instant;
8
9use chrono::Utc;
10use futures::stream::{BoxStream, Stream};
11use uuid::Uuid;
12
13use crate::error::{GatewayError, LegFailure};
14use crate::guard::GuardEngine;
15use crate::jev_native::JevNativeProvider;
16use crate::ledger::{LedgerHandle, UsageEntry};
17use crate::observability::GenAiSpan;
18use crate::pricing::PricingTable;
19use crate::providers::Catalog;
20use crate::routing::classify::{classify, vertex_triggers, Lane};
21use crate::routing::executor::{
22    execute_buffered_with_timeouts, CommittedStream, Completion, LegError, StreamTimeouts,
23};
24use crate::routing::request::ChatRequest;
25use crate::routing::stream::{Accumulator, FinishReason, StreamItem};
26use crate::routing::table::{ChainLeg, RouteTable};
27use crate::vertex_native::VertexNativeProvider;
28
29/// Embeddable gateway handle. Construct with [`Gateway::builder`].
30#[derive(Clone)]
31pub struct Gateway {
32    pub(crate) routes: Arc<RouteTable>,
33    pub(crate) catalog: Arc<Catalog>,
34    pub(crate) pricing: Arc<PricingTable>,
35    pub(crate) ledger: LedgerHandle,
36    pub(crate) vertex_native: Option<Arc<VertexNativeProvider>>,
37    pub(crate) jev_native: Option<Arc<JevNativeProvider>>,
38    pub(crate) timeouts: StreamTimeouts,
39    pub(crate) default_tenant: String,
40    pub(crate) embed_routes: Arc<crate::routing::embeddings::EmbeddingRouteTable>,
41    pub(crate) embedders:
42        std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>,
43    pub(crate) embed_default_input_per_mtok: f64,
44    pub(crate) guard: Arc<GuardEngine>,
45    pub(crate) ai_task_types: Arc<crate::ai_task_type::AiTaskTypeTable>,
46}
47
48/// Per-call identity for an in-process request (replaces HTTP headers).
49#[derive(Debug, Clone, Default)]
50pub struct RequestCtx {
51    pub tenant: Option<String>,
52    pub workspace: Option<String>,
53    /// End-user attribution within the tenant (`x-synapse-user` over HTTP).
54    pub user: Option<String>,
55    /// Conversation / agent thread (`x-synapse-thread` over HTTP).
56    pub thread: Option<String>,
57    /// Chat / work message within the thread (`x-synapse-message` over HTTP).
58    pub message: Option<String>,
59    /// Caller-supplied classification of the work this request serves
60    /// (`x-synapse-user-task-type` over HTTP). Free-form and never interpreted
61    /// by the gateway; recorded on the ledger row for downstream reporting.
62    pub user_task_type: Option<String>,
63    /// Overrides the route-alias inference for the AI task type
64    /// (`x-synapse-ai-task-type` over HTTP). When `None` or empty, the task type
65    /// is inferred from the route alias, falling back to `"simple"`.
66    pub ai_task_type: Option<String>,
67    /// Correlation id for the ledger row. When `None`, falls back to `message`
68    /// (so a chat turn correlates with usage), then a generated UUID.
69    /// Embedders (and the HTTP layer) can pass their own so the ledger row and
70    /// any client-facing response id share one value.
71    pub request_id: Option<String>,
72}
73
74impl RequestCtx {
75    /// Prefer an explicit `request_id`, else the chat `message` id, else `None`
76    /// (callers generate a UUID).
77    pub fn resolved_request_id(&self) -> Option<String> {
78        self.request_id
79            .clone()
80            .or_else(|| self.message.clone())
81            .filter(|s| !s.is_empty())
82    }
83}
84
85/// Outcome of running a route's `typesafe` legs on the Jev lane.
86pub(crate) enum JevAttempt {
87    /// Jev answered; the completion's content is the JSON-encoded `answers`
88    /// map, and `answers` is the parsed map (hybrid extraction reads it).
89    Decided {
90        completion: Completion,
91        answers: serde_json::Value,
92    },
93    /// Every typesafe leg failed retryably — the caller falls through to
94    /// the remaining legs on the standard/native machinery.
95    Exhausted(Vec<LegFailure>),
96}
97
98/// Result of a buffered chat call. `Plain` is a normal chat completion;
99/// `Hybrid` is a Jev hybrid-extraction response (envelope per the spec).
100#[derive(Debug)]
101pub enum ChatOutcome {
102    Plain(Completion),
103    Hybrid(HybridOutcome),
104}
105
106/// A completed Jev hybrid extraction: Jev's answers plus per-survivor
107/// extraction results. Rendered by `server::hybrid_json`.
108#[derive(Debug)]
109pub struct HybridOutcome {
110    /// Jev build that answered, or `jev-latest` when the lane was exhausted.
111    pub model: String,
112    /// All Jev answers verbatim (`{}` when degraded).
113    pub answers: serde_json::Value,
114    /// Candidate keys that met the floor (all candidates when degraded).
115    pub survivors: Vec<String>,
116    /// Jev lane exhausted, or at least one survivor failed its legs.
117    pub degraded: bool,
118    /// True when at least one survivor extraction was attempted.
119    pub extraction_ran: bool,
120    /// Candidate key → extraction result (failed survivors omitted).
121    pub extractions: serde_json::Map<String, serde_json::Value>,
122    /// Jev + extraction token totals.
123    pub input_tokens: u64,
124    pub output_tokens: u64,
125}
126
127/// The attribution fields a ledger row carries, resolved once per request.
128///
129/// Grouped so the streaming and passthrough meters take a single named argument
130/// instead of six positional `Option<String>`s, where a swapped pair would
131/// compile cleanly and silently mis-attribute usage.
132#[derive(Debug, Clone)]
133pub struct Attribution {
134    pub workspace: Option<String>,
135    pub user: Option<String>,
136    pub thread: Option<String>,
137    pub message: Option<String>,
138    pub user_task_type: Option<String>,
139    /// Always resolved — see [`Gateway::ai_task_type_of`].
140    pub ai_task_type: String,
141}
142
143impl Default for Attribution {
144    fn default() -> Self {
145        Self {
146            workspace: None,
147            user: None,
148            thread: None,
149            message: None,
150            user_task_type: None,
151            ai_task_type: crate::ai_task_type::DEFAULT_AI_TASK_TYPE.to_string(),
152        }
153    }
154}
155
156#[derive(Default)]
157pub struct GatewayBuilder {
158    routes: Option<RouteTable>,
159    catalog: Option<Catalog>,
160    pricing: Option<PricingTable>,
161    ledger: Option<LedgerHandle>,
162    vertex_native: Option<VertexNativeProvider>,
163    jev_native: Option<JevNativeProvider>,
164    timeouts: Option<StreamTimeouts>,
165    default_tenant: Option<String>,
166    embed_routes: Option<crate::routing::embeddings::EmbeddingRouteTable>,
167    embedders:
168        Option<std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>>,
169    embed_default_input_per_mtok: Option<f64>,
170    guard: Option<GuardEngine>,
171    ai_task_types: Option<crate::ai_task_type::AiTaskTypeTable>,
172}
173
174impl Gateway {
175    pub fn builder() -> GatewayBuilder {
176        GatewayBuilder::default()
177    }
178
179    /// Client-facing model aliases (for `/v1/models` and discovery).
180    pub fn model_aliases(&self) -> Vec<String> {
181        self.routes.aliases()
182    }
183
184    fn resolve_legs(&self, req: &ChatRequest) -> Result<Vec<ChainLeg>, GatewayError> {
185        self.routes
186            .legs(&req.model)
187            .ok_or_else(|| GatewayError::UnknownModel(req.model.clone()))
188            .map(<[ChainLeg]>::to_vec)
189    }
190
191    /// Run the route's guardrail policy over the request input. Falls back to
192    /// the `default` policy; a no-op when neither is configured.
193    fn guard_input(&self, req: &ChatRequest) -> Result<(), GatewayError> {
194        let policy = self.routes.policy_of(&req.model).unwrap_or("default");
195        self.guard.guard(policy, req)
196    }
197
198    fn tenant_of<'a>(&'a self, ctx: &'a RequestCtx) -> &'a str {
199        ctx.tenant.as_deref().unwrap_or(&self.default_tenant)
200    }
201
202    /// Resolve the AI task type for a call: the caller's header when non-empty,
203    /// else the task type configured for `alias`, else `"simple"`.
204    pub(crate) fn ai_task_type_of(&self, ctx: &RequestCtx, alias: &str) -> String {
205        ctx.ai_task_type
206            .as_deref()
207            .filter(|s| !s.is_empty())
208            .unwrap_or_else(|| self.ai_task_types.resolve(alias))
209            .to_string()
210    }
211
212    /// Everything a ledger row needs from the request, resolved once so every
213    /// record site (buffered, streaming, embedding, passthrough) agrees.
214    pub(crate) fn attribution_of(&self, ctx: &RequestCtx, alias: &str) -> Attribution {
215        Attribution {
216            workspace: ctx.workspace.clone(),
217            user: ctx.user.clone(),
218            thread: ctx.thread.clone(),
219            message: ctx.message.clone(),
220            user_task_type: ctx.user_task_type.clone(),
221            ai_task_type: self.ai_task_type_of(ctx, alias),
222        }
223    }
224
225    /// Fire cost + ledger + metrics for a completed call (buffered side-effects).
226    #[allow(clippy::too_many_arguments)]
227    pub(crate) fn record(
228        &self,
229        ctx: &RequestCtx,
230        route: &str,
231        lane: &'static str,
232        request_id: &str,
233        c: &Completion,
234        legs: u32,
235        started: Instant,
236    ) {
237        let cost = self
238            .pricing
239            .cost_usd(&c.provider, &c.model, c.input_tokens, c.output_tokens);
240        let attr = self.attribution_of(ctx, route);
241        self.ledger.enqueue(UsageEntry {
242            ts: Utc::now(),
243            tenant: self.tenant_of(ctx).to_string(),
244            workspace: attr.workspace,
245            user: attr.user,
246            thread: attr.thread,
247            message: attr.message,
248            route: route.to_string(),
249            provider: c.provider.clone(),
250            model: c.model.clone(),
251            lane: lane.to_string(),
252            input_tokens: c.input_tokens,
253            output_tokens: c.output_tokens,
254            cost_usd: cost,
255            request_id: request_id.to_string(),
256            status: "ok".into(),
257            op: "chat".into(),
258            user_task_type: attr.user_task_type,
259            ai_task_type: attr.ai_task_type,
260        });
261        let span_lane = match lane {
262            "native" => Lane::NativeVertex,
263            "jev" => Lane::Jev,
264            _ => Lane::Standard,
265        };
266        GenAiSpan::from_completion(
267            c,
268            span_lane,
269            route,
270            self.tenant_of(ctx),
271            ctx.workspace.as_deref(),
272            legs,
273            false,
274        )
275        .emit_metrics(started.elapsed().as_secs_f64());
276    }
277
278    /// Resolve the native Vertex committed stream (shared by chat/chat_stream).
279    ///
280    /// Commit the first Vertex leg that can open a stream. On retryable upstream
281    /// failures (5xx, 429, 408, transport), advance to the next `provider=vertex`
282    /// leg. Non-retryable 4xx abort immediately. Non-vertex legs are skipped
283    /// (native features cannot be expressed on other providers).
284    ///
285    /// NOTE: bounded only by the reqwest client timeout (`config.request_timeout`);
286    /// the first-chunk/idle `StreamTimeouts` are not applied here yet.
287    pub(crate) async fn native_committed(
288        &self,
289        req: &ChatRequest,
290        legs: &[ChainLeg],
291    ) -> Result<CommittedStream, GatewayError> {
292        let provider = self
293            .vertex_native
294            .as_ref()
295            .ok_or_else(|| GatewayError::BadRequest("native vertex lane not configured".into()))?;
296
297        let vertex_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "vertex").collect();
298        if vertex_legs.is_empty() {
299            return Err(GatewayError::NativeFeatureUnsupported {
300                feature: "native-vertex".into(),
301                route: req.model.clone(),
302            });
303        }
304
305        let mut failures: Vec<LegFailure> = Vec::new();
306        let mut last_retryable: Option<GatewayError> = None;
307        for leg in &vertex_legs {
308            match provider
309                .stream_generate(&leg.model, req, leg.region.as_deref())
310                .await
311            {
312                Ok(stream) => {
313                    return Ok(CommittedStream::single(
314                        "vertex".into(),
315                        leg.model.clone(),
316                        stream,
317                    ));
318                }
319                Err(e) if native_start_retryable(&e) => {
320                    failures.push(LegFailure {
321                        provider: leg.provider.clone(),
322                        model: leg.model.clone(),
323                        message: e.to_string(),
324                    });
325                    last_retryable = Some(e);
326                }
327                Err(e) => return Err(e),
328            }
329        }
330        Err(
331            last_retryable.unwrap_or_else(|| GatewayError::AllLegsFailed {
332                route: req.model.clone(),
333                failures,
334            }),
335        )
336    }
337
338    /// Run the route's `typesafe` legs against TypeSafe System One (Jev):
339    /// state + typed questions in, structured decisions out. Non-retryable
340    /// upstream failures (malformed questions, 4xx other than 429/408) abort —
341    /// a chat fallback cannot fix them. Retryable failures (429/408/5xx,
342    /// transport) advance to the next typesafe leg, then report `Exhausted`
343    /// so the caller can fall back to the route's chat legs.
344    pub(crate) async fn jev_attempt(
345        &self,
346        req: &ChatRequest,
347        legs: &[ChainLeg],
348    ) -> Result<JevAttempt, GatewayError> {
349        let jev_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "typesafe").collect();
350        if jev_legs.is_empty() {
351            return Ok(JevAttempt::Exhausted(Vec::new()));
352        }
353        let ext = req
354            .jev
355            .as_ref()
356            .filter(|j| !j.questions.is_empty())
357            .ok_or_else(|| {
358                GatewayError::BadRequest(format!(
359                    "route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
360                    req.model
361                ))
362            })?;
363        let provider = self.jev_native.as_ref().ok_or_else(|| {
364            GatewayError::BadRequest(
365                "typesafe legs require TYPESAFE_API_KEY to be configured".into(),
366            )
367        })?;
368
369        let state = ext
370            .state
371            .clone()
372            .or_else(|| serde_json::to_string(&req.messages).ok())
373            .unwrap_or_default();
374        let mut body = serde_json::json!({
375            "state": state,
376            "questions": ext.questions.clone(),
377        });
378
379        let mut failures: Vec<LegFailure> = Vec::new();
380        for leg in jev_legs {
381            let model = if leg.model.is_empty() {
382                crate::jev_native::DEFAULT_MODEL.to_string()
383            } else {
384                leg.model.clone()
385            };
386            body["model"] = serde_json::Value::String(model.clone());
387
388            let resp = match provider.evaluate(body.clone()).await {
389                Ok(r) => r,
390                Err(e) if native_start_retryable(&e) => {
391                    failures.push(LegFailure {
392                        provider: leg.provider.clone(),
393                        model,
394                        message: e.to_string(),
395                    });
396                    continue;
397                }
398                Err(e) => return Err(e),
399            };
400
401            let status = resp.status();
402            let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
403                status: 502,
404                body: e.to_string(),
405            })?;
406            if status.is_success() {
407                let value: serde_json::Value =
408                    serde_json::from_slice(&bytes).map_err(|e| GatewayError::Upstream {
409                        status: 502,
410                        body: format!("jev response is not JSON: {e}"),
411                    })?;
412                return Ok(JevAttempt::Decided {
413                    answers: value["answers"].clone(),
414                    completion: Completion {
415                        provider: "typesafe".into(),
416                        model: value["model"].as_str().unwrap_or(&model).to_string(),
417                        content: serde_json::to_string(&value["answers"])
418                            .unwrap_or_else(|_| "{}".into()),
419                        tool_calls: Vec::new(),
420                        finish_reason: FinishReason::Stop,
421                        input_tokens: value["usage"]["input_tokens"].as_u64().unwrap_or(0),
422                        output_tokens: value["usage"]["output_tokens"].as_u64().unwrap_or(0),
423                    },
424                });
425            }
426
427            let message = String::from_utf8_lossy(&bytes).into_owned();
428            if status.is_server_error()
429                || status == reqwest::StatusCode::TOO_MANY_REQUESTS
430                || status == reqwest::StatusCode::REQUEST_TIMEOUT
431            {
432                failures.push(LegFailure {
433                    provider: leg.provider.clone(),
434                    model,
435                    message: format!("{status}: {message}"),
436                });
437                continue;
438            }
439            return Err(GatewayError::BadRequest(format!(
440                "typesafe {}: {message}",
441                status.as_u16()
442            )));
443        }
444        Ok(JevAttempt::Exhausted(failures))
445    }
446
447    /// Jev answers are unary: present a buffered decision as a single content
448    /// delta + terminal item so the streaming surface renders one SSE chunk.
449    fn jev_committed(c: Completion) -> CommittedStream {
450        let items: Vec<Result<StreamItem, LegError>> = vec![
451            Ok(StreamItem::Delta(c.content.clone())),
452            Ok(StreamItem::Done {
453                input_tokens: c.input_tokens,
454                output_tokens: c.output_tokens,
455                finish_reason: c.finish_reason,
456            }),
457        ];
458        CommittedStream::single("typesafe".into(), c.model, futures::stream::iter(items))
459    }
460
461    /// Streaming in-process chat: commit a leg on the first item; returns a
462    /// `GuardedStream` that yields items and fires side-effects on completion/drop.
463    pub async fn chat_stream(
464        &self,
465        req: ChatRequest,
466        ctx: &RequestCtx,
467    ) -> Result<GuardedStream, GatewayError> {
468        use crate::routing::executor::execute_streaming_with_timeouts;
469        let started = Instant::now();
470        let legs = self.resolve_legs(&req)?;
471        self.guard_input(&req)?;
472        let request_id = ctx
473            .resolved_request_id()
474            .unwrap_or_else(|| Uuid::new_v4().to_string());
475        let vertex_leg_count = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
476        require_jev_block(&req, &legs)?;
477        if let Err(message) = crate::routing::jev_extract::validate_extract(
478            &req,
479            legs.iter().any(|l| l.provider != "typesafe"),
480        ) {
481            return Err(GatewayError::BadRequest(message));
482        }
483        let (committed, lane_str, legs_attempted) = match classify(&req) {
484            Lane::Standard => (
485                execute_streaming_with_timeouts(
486                    &self.catalog,
487                    &req.model,
488                    &legs,
489                    &req,
490                    self.timeouts,
491                )
492                .await?,
493                "standard",
494                legs.len() as u32,
495            ),
496            Lane::NativeVertex => (
497                self.native_committed(&req, &legs).await?,
498                "native",
499                vertex_leg_count.max(1),
500            ),
501            Lane::Jev => match self.jev_attempt(&req, &legs).await? {
502                JevAttempt::Decided { completion, .. } => (
503                    Self::jev_committed(completion),
504                    "jev",
505                    legs.iter()
506                        .filter(|l| l.provider == "typesafe")
507                        .count()
508                        .max(1) as u32,
509                ),
510                JevAttempt::Exhausted(failures) => {
511                    let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
512                    if rest.is_empty() {
513                        return Err(GatewayError::AllLegsFailed {
514                            route: req.model.clone(),
515                            failures,
516                        });
517                    }
518                    if vertex_triggers(&req) {
519                        let vertex_n =
520                            rest.iter().filter(|l| l.provider == "vertex").count() as u32;
521                        (
522                            self.native_committed(&req, &rest).await?,
523                            "native",
524                            vertex_n.max(1),
525                        )
526                    } else {
527                        (
528                            execute_streaming_with_timeouts(
529                                &self.catalog,
530                                &req.model,
531                                &rest,
532                                &req,
533                                self.timeouts,
534                            )
535                            .await?,
536                            "standard",
537                            rest.len() as u32,
538                        )
539                    }
540                }
541            },
542        };
543        let model = committed.model.clone();
544        let guard = StreamSideEffects::new(
545            self.ledger.clone(),
546            self.pricing.clone(),
547            req.model.clone(),
548            self.tenant_of(ctx).to_string(),
549            self.attribution_of(ctx, &req.model),
550            committed.provider.clone(),
551            committed.model.clone(),
552            lane_str,
553            request_id,
554            legs_attempted,
555            started,
556        );
557        Ok(GuardedStream::new(committed.stream, model, guard))
558    }
559
560    /// Buffered in-process chat: stream the chain internally, aggregate, fire
561    /// side-effects, return a `ChatOutcome` (plain or hybrid).
562    pub async fn chat(
563        &self,
564        req: ChatRequest,
565        ctx: &RequestCtx,
566    ) -> Result<ChatOutcome, GatewayError> {
567        let started = Instant::now();
568        let legs = self.resolve_legs(&req)?;
569        self.guard_input(&req)?;
570        let request_id = ctx
571            .resolved_request_id()
572            .unwrap_or_else(|| Uuid::new_v4().to_string());
573        require_jev_block(&req, &legs)?;
574        if let Err(message) = crate::routing::jev_extract::validate_extract(
575            &req,
576            legs.iter().any(|l| l.provider != "typesafe"),
577        ) {
578            return Err(GatewayError::BadRequest(message));
579        }
580        if req.jev.as_ref().and_then(|j| j.extract.as_ref()).is_some() {
581            let outcome = self
582                .chat_hybrid(req, ctx, &legs, started, &request_id)
583                .await?;
584            return Ok(ChatOutcome::Hybrid(outcome));
585        }
586        let (completion, lane_str, legs_n) = match classify(&req) {
587            Lane::Standard => execute_buffered_with_timeouts(
588                &self.catalog,
589                &req.model,
590                &legs,
591                &req,
592                self.timeouts,
593            )
594            .await
595            .map(|c| (c, "standard", legs.len() as u32))?,
596            Lane::NativeVertex => {
597                let vertex_n = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
598                let committed = self.native_committed(&req, &legs).await?;
599                (
600                    collect_committed(committed).await?,
601                    "native",
602                    vertex_n.max(1),
603                )
604            }
605            Lane::Jev => match self.jev_attempt(&req, &legs).await? {
606                JevAttempt::Decided { completion, .. } => (
607                    completion,
608                    "jev",
609                    legs.iter()
610                        .filter(|l| l.provider == "typesafe")
611                        .count()
612                        .max(1) as u32,
613                ),
614                JevAttempt::Exhausted(failures) => {
615                    let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
616                    if rest.is_empty() {
617                        return Err(GatewayError::AllLegsFailed {
618                            route: req.model.clone(),
619                            failures,
620                        });
621                    }
622                    if vertex_triggers(&req) {
623                        let committed = self.native_committed(&req, &rest).await?;
624                        let vertex_n =
625                            rest.iter().filter(|l| l.provider == "vertex").count() as u32;
626                        (
627                            collect_committed(committed).await?,
628                            "native",
629                            vertex_n.max(1),
630                        )
631                    } else {
632                        execute_buffered_with_timeouts(
633                            &self.catalog,
634                            &req.model,
635                            &rest,
636                            &req,
637                            self.timeouts,
638                        )
639                        .await
640                        .map(|c| (c, "standard", rest.len() as u32))?
641                    }
642                }
643            },
644        };
645        self.record(
646            ctx,
647            &req.model,
648            lane_str,
649            &request_id,
650            &completion,
651            legs_n,
652            started,
653        );
654        Ok(ChatOutcome::Plain(completion))
655    }
656
657    /// Embed `req.input` against the embedding alias `req.model`, pinning output to
658    /// the alias's declared dimension. Tries each leg in order, falling through on
659    /// upstream error (simple ordered fallback). Records one ledger row on success.
660    ///
661    // TODO(embeddings): wrap legs in resilience breakers like the chat path.
662    pub async fn embed(
663        &self,
664        req: crate::embeddings::EmbeddingRequest,
665        ctx: RequestCtx,
666    ) -> Result<crate::embeddings::EmbeddingResponse, GatewayError> {
667        let alias = req.model.clone();
668        let dims = self
669            .embed_routes
670            .dimensions(&alias)
671            .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?;
672        if let Some(d) = req.dimensions {
673            if d != dims {
674                return Err(GatewayError::BadRequest(format!(
675                    "dimensions {d} does not match embedding alias '{alias}' dimension {dims}"
676                )));
677            }
678        }
679        let inputs = req.input.into_vec();
680        let legs = self
681            .embed_routes
682            .legs(&alias)
683            .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?
684            .to_vec();
685
686        let mut last_err: Option<GatewayError> = None;
687        for leg in &legs {
688            let Some(embedder) = self.embedders.get(&leg.provider) else {
689                continue;
690            };
691            let limit = match leg.provider.as_str() {
692                "vertex" => crate::embeddings::vertex::VERTEX_EMBED_BATCH,
693                _ => crate::embeddings::openai::OPENAI_EMBED_BATCH,
694            };
695            let started = std::time::Instant::now();
696            match self
697                .embed_all_batches(embedder.as_ref(), &leg.model, &inputs, dims, limit)
698                .await
699            {
700                Ok(out) => {
701                    metrics::counter!(
702                        "synapse_embeddings_total",
703                        "route" => alias.clone(),
704                        "model" => leg.model.clone(),
705                        "provider" => leg.provider.clone(),
706                    )
707                    .increment(1);
708                    metrics::histogram!(
709                        "synapse_embedding_duration_seconds",
710                        "route" => alias.clone(),
711                        "model" => leg.model.clone(),
712                        "provider" => leg.provider.clone(),
713                    )
714                    .record(started.elapsed().as_secs_f64());
715                    self.record_embed_usage(&ctx, &alias, leg, out.input_tokens);
716                    return Ok(crate::embeddings::build_response(alias, out));
717                }
718                Err(e) => last_err = Some(e),
719            }
720        }
721        Err(last_err.unwrap_or(GatewayError::AllLegsFailed {
722            route: alias,
723            failures: Vec::new(),
724        }))
725    }
726
727    /// Embed `inputs` in provider-sized batches, concatenating vectors (input order)
728    /// and summing input tokens into a single `EmbedOut`.
729    async fn embed_all_batches(
730        &self,
731        embedder: &dyn crate::embeddings::EmbeddingProvider,
732        model: &str,
733        inputs: &[String],
734        dims: u32,
735        limit: usize,
736    ) -> Result<crate::embeddings::EmbedOut, GatewayError> {
737        let mut vectors = Vec::with_capacity(inputs.len());
738        let mut input_tokens = 0u64;
739        for batch in crate::embeddings::split_batches(inputs, limit) {
740            let out = embedder.embed(model, batch, dims).await?;
741            input_tokens += out.input_tokens;
742            vectors.extend(out.vectors);
743        }
744        Ok(crate::embeddings::EmbedOut {
745            vectors,
746            input_tokens,
747        })
748    }
749
750    /// Fire one cost + ledger row for a completed embedding call (input tokens only).
751    fn record_embed_usage(&self, ctx: &RequestCtx, alias: &str, leg: &ChainLeg, input_tokens: u64) {
752        let cost = self.pricing.embedding_cost_usd(
753            &leg.provider,
754            &leg.model,
755            input_tokens,
756            self.embed_default_input_per_mtok,
757        );
758        let attr = self.attribution_of(ctx, alias);
759        self.ledger.enqueue(UsageEntry {
760            ts: Utc::now(),
761            tenant: self.tenant_of(ctx).to_string(),
762            workspace: attr.workspace,
763            user: attr.user,
764            thread: attr.thread,
765            message: attr.message,
766            route: alias.to_string(),
767            provider: leg.provider.clone(),
768            model: leg.model.clone(),
769            lane: "embedding".into(),
770            input_tokens,
771            output_tokens: 0,
772            cost_usd: cost,
773            request_id: ctx
774                .resolved_request_id()
775                .unwrap_or_else(|| Uuid::new_v4().to_string()),
776            status: "ok".into(),
777            op: "embedding".into(),
778            user_task_type: attr.user_task_type,
779            ai_task_type: attr.ai_task_type,
780        });
781    }
782}
783
784/// True when a native-lane stream **start** failure should advance to the next
785/// Vertex leg (aligned with standard-lane / reqwest retryable policy).
786fn native_start_retryable(e: &GatewayError) -> bool {
787    match e {
788        GatewayError::Upstream { status, .. } => {
789            *status >= 500
790                || *status == reqwest::StatusCode::TOO_MANY_REQUESTS.as_u16()
791                || *status == reqwest::StatusCode::REQUEST_TIMEOUT.as_u16()
792        }
793        GatewayError::UpstreamTimeout => true,
794        _ => false,
795    }
796}
797
798/// Routes with `typesafe` legs answer typed questions, not chat: refuse with a
799/// clear 400 when the request carries no `jev` extension block, instead of
800/// letting the standard executor fail on a provider it has no client for.
801fn require_jev_block(req: &ChatRequest, legs: &[ChainLeg]) -> Result<(), GatewayError> {
802    let needs_block = legs.iter().any(|l| l.provider == "typesafe");
803    let has_block = req.jev.as_ref().is_some_and(|j| !j.questions.is_empty());
804    match (needs_block, has_block) {
805        (true, false) => Err(GatewayError::BadRequest(format!(
806            "route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
807            req.model
808        ))),
809        _ => Ok(()),
810    }
811}
812
813/// Legs remaining after the Jev lane has consumed the `typesafe` ones: the
814/// fallback chain when every Jev leg fails retryably.
815fn non_typesafe_legs(legs: &[ChainLeg]) -> Vec<ChainLeg> {
816    legs.iter()
817        .filter(|l| l.provider != "typesafe")
818        .cloned()
819        .collect()
820}
821
822impl GatewayBuilder {
823    pub fn routes(mut self, routes: RouteTable) -> Self {
824        self.routes = Some(routes);
825        self
826    }
827    pub fn catalog(mut self, catalog: Catalog) -> Self {
828        self.catalog = Some(catalog);
829        self
830    }
831    pub fn pricing(mut self, pricing: PricingTable) -> Self {
832        self.pricing = Some(pricing);
833        self
834    }
835    pub fn ledger(mut self, ledger: LedgerHandle) -> Self {
836        self.ledger = Some(ledger);
837        self
838    }
839    pub fn vertex_native(mut self, v: Option<VertexNativeProvider>) -> Self {
840        self.vertex_native = v;
841        self
842    }
843    pub fn jev_native(mut self, v: Option<JevNativeProvider>) -> Self {
844        self.jev_native = v;
845        self
846    }
847    pub fn timeouts(mut self, t: StreamTimeouts) -> Self {
848        self.timeouts = Some(t);
849        self
850    }
851    pub fn default_tenant(mut self, t: impl Into<String>) -> Self {
852        self.default_tenant = Some(t.into());
853        self
854    }
855    pub fn embed_routes(mut self, t: crate::routing::embeddings::EmbeddingRouteTable) -> Self {
856        self.embed_routes = Some(t);
857        self
858    }
859    pub fn embedder(
860        mut self,
861        id: impl Into<String>,
862        e: Arc<dyn crate::embeddings::EmbeddingProvider>,
863    ) -> Self {
864        self.embedders
865            .get_or_insert_with(Default::default)
866            .insert(id.into(), e);
867        self
868    }
869    pub fn embed_default_input_per_mtok(mut self, v: f64) -> Self {
870        self.embed_default_input_per_mtok = Some(v);
871        self
872    }
873    pub fn guard(mut self, guard: GuardEngine) -> Self {
874        self.guard = Some(guard);
875        self
876    }
877    pub fn ai_task_types(mut self, t: crate::ai_task_type::AiTaskTypeTable) -> Self {
878        self.ai_task_types = Some(t);
879        self
880    }
881
882    pub fn build(self) -> anyhow::Result<Gateway> {
883        Ok(Gateway {
884            routes: Arc::new(
885                self.routes
886                    .ok_or_else(|| anyhow::anyhow!("Gateway: routes required"))?,
887            ),
888            catalog: Arc::new(
889                self.catalog
890                    .ok_or_else(|| anyhow::anyhow!("Gateway: catalog required"))?,
891            ),
892            pricing: Arc::new(
893                self.pricing
894                    .ok_or_else(|| anyhow::anyhow!("Gateway: pricing required"))?,
895            ),
896            ledger: self
897                .ledger
898                .ok_or_else(|| anyhow::anyhow!("Gateway: ledger required"))?,
899            vertex_native: self.vertex_native.map(Arc::new),
900            jev_native: self.jev_native.map(Arc::new),
901            timeouts: self.timeouts.unwrap_or_default(),
902            default_tenant: self.default_tenant.unwrap_or_else(|| "unattributed".into()),
903            embed_routes: Arc::new(self.embed_routes.unwrap_or_default()),
904            embedders: self.embedders.unwrap_or_default(),
905            embed_default_input_per_mtok: self.embed_default_input_per_mtok.unwrap_or(0.10),
906            guard: Arc::new(self.guard.unwrap_or_else(GuardEngine::empty)),
907            ai_task_types: Arc::new(self.ai_task_types.unwrap_or_default()),
908        })
909    }
910}
911
912/// Accumulates running usage during a streamed response and fires cost, ledger,
913/// and metrics exactly once when dropped (normal end, error, or client disconnect).
914/// Emitting from `Drop` guarantees the side-effects run on every termination path,
915/// including client disconnect where the response future is cancelled.
916pub(crate) struct StreamSideEffects {
917    ledger: LedgerHandle,
918    pricing: Arc<PricingTable>,
919    route: String,
920    tenant: String,
921    attribution: Attribution,
922    provider: String,
923    model: String,
924    lane: &'static str, // "standard" | "native"
925    request_id: String,
926    legs_attempted: u32,
927    started: Instant,
928    input_tokens: u64,
929    output_tokens: u64,
930    status: &'static str,
931    fired: bool,
932}
933
934impl StreamSideEffects {
935    #[allow(clippy::too_many_arguments)]
936    pub(crate) fn new(
937        ledger: LedgerHandle,
938        pricing: Arc<PricingTable>,
939        route: String,
940        tenant: String,
941        attribution: Attribution,
942        provider: String,
943        model: String,
944        lane: &'static str,
945        request_id: String,
946        legs_attempted: u32,
947        started: Instant,
948    ) -> Self {
949        Self {
950            ledger,
951            pricing,
952            route,
953            tenant,
954            attribution,
955            provider,
956            model,
957            lane,
958            request_id,
959            legs_attempted,
960            started,
961            input_tokens: 0,
962            output_tokens: 0,
963            status: "ok",
964            fired: false,
965        }
966    }
967
968    /// Fold a streamed item into the running usage totals.
969    pub(crate) fn observe(&mut self, item: &StreamItem) {
970        if let StreamItem::Done {
971            input_tokens,
972            output_tokens,
973            ..
974        } = item
975        {
976            self.input_tokens = *input_tokens;
977            self.output_tokens = *output_tokens;
978        }
979    }
980
981    /// Mark the response as failed so the ledger row records `status = "error"`.
982    pub(crate) fn mark_error(&mut self) {
983        self.status = "error";
984    }
985}
986
987impl Drop for StreamSideEffects {
988    fn drop(&mut self) {
989        if self.fired {
990            return;
991        }
992        self.fired = true;
993
994        // Cost + ledger (fire-and-forget), mirroring the buffered handler.
995        let cost = self.pricing.cost_usd(
996            &self.provider,
997            &self.model,
998            self.input_tokens,
999            self.output_tokens,
1000        );
1001        self.ledger.enqueue(UsageEntry {
1002            ts: Utc::now(),
1003            tenant: self.tenant.clone(),
1004            workspace: self.attribution.workspace.clone(),
1005            user: self.attribution.user.clone(),
1006            thread: self.attribution.thread.clone(),
1007            message: self.attribution.message.clone(),
1008            route: self.route.clone(),
1009            provider: self.provider.clone(),
1010            model: self.model.clone(),
1011            lane: self.lane.to_string(),
1012            input_tokens: self.input_tokens,
1013            output_tokens: self.output_tokens,
1014            cost_usd: cost,
1015            request_id: self.request_id.clone(),
1016            status: self.status.to_string(),
1017            op: "chat".into(),
1018            user_task_type: self.attribution.user_task_type.clone(),
1019            ai_task_type: self.attribution.ai_task_type.clone(),
1020        });
1021
1022        // Metrics via GenAiSpan for parity with the buffered path: build a minimal
1023        // Completion from the running fields and emit with `stream: true`.
1024        let completion = Completion {
1025            provider: self.provider.clone(),
1026            model: self.model.clone(),
1027            content: String::new(),
1028            tool_calls: Vec::new(),
1029            finish_reason: FinishReason::Stop,
1030            input_tokens: self.input_tokens,
1031            output_tokens: self.output_tokens,
1032        };
1033        let lane = if self.lane == "native" {
1034            Lane::NativeVertex
1035        } else {
1036            Lane::Standard
1037        };
1038        GenAiSpan::from_completion(
1039            &completion,
1040            lane,
1041            &self.route,
1042            &self.tenant,
1043            self.attribution.workspace.as_deref(),
1044            self.legs_attempted,
1045            true,
1046        )
1047        .emit_metrics(self.started.elapsed().as_secs_f64());
1048    }
1049}
1050
1051/// A streaming response that yields normalized `StreamItem`s and fires
1052/// cost/ledger/metrics exactly once when fully consumed OR dropped early
1053/// (via the owned `StreamSideEffects` guard). Transport-independent.
1054pub struct GuardedStream {
1055    inner: BoxStream<'static, Result<StreamItem, LegError>>,
1056    model: String,
1057    guard: StreamSideEffects,
1058}
1059
1060impl GuardedStream {
1061    pub(crate) fn new(
1062        inner: BoxStream<'static, Result<StreamItem, LegError>>,
1063        model: String,
1064        guard: StreamSideEffects,
1065    ) -> Self {
1066        Self {
1067            inner,
1068            model,
1069            guard,
1070        }
1071    }
1072
1073    /// The model id of the committed leg (for the OpenAI chunk `model` field).
1074    pub fn model(&self) -> &str {
1075        &self.model
1076    }
1077}
1078
1079impl Stream for GuardedStream {
1080    type Item = Result<StreamItem, LegError>;
1081    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1082        // Both fields are Unpin (BoxStream is Pin<Box<..>>; StreamSideEffects is plain data).
1083        let this = self.get_mut();
1084        match this.inner.as_mut().poll_next(cx) {
1085            Poll::Ready(Some(Ok(item))) => {
1086                this.guard.observe(&item);
1087                Poll::Ready(Some(Ok(item)))
1088            }
1089            Poll::Ready(Some(Err(e))) => {
1090                this.guard.mark_error();
1091                Poll::Ready(Some(Err(e)))
1092            }
1093            Poll::Ready(None) => Poll::Ready(None),
1094            Poll::Pending => Poll::Pending,
1095        }
1096    }
1097}
1098
1099/// Drain a committed stream into a single `Completion`. ROP: map the per-item
1100/// error onto the failure track, `try_fold` items into an `Accumulator`, then map
1101/// to a `Completion`.
1102pub(crate) async fn collect_committed(
1103    committed: CommittedStream,
1104) -> Result<Completion, GatewayError> {
1105    use futures::TryStreamExt;
1106    let CommittedStream {
1107        provider,
1108        model,
1109        stream,
1110    } = committed;
1111    stream
1112        .map_err(|e: LegError| GatewayError::Upstream {
1113            status: 502,
1114            body: e.to_string(),
1115        })
1116        .try_fold(Accumulator::default(), |mut acc, item| async move {
1117            acc.push(item);
1118            Ok(acc)
1119        })
1120        .await
1121        .map(|acc| Completion {
1122            provider,
1123            model,
1124            content: acc.content,
1125            tool_calls: acc.tool_calls,
1126            finish_reason: acc.finish_reason,
1127            input_tokens: acc.input_tokens,
1128            output_tokens: acc.output_tokens,
1129        })
1130}
1131
1132#[cfg(test)]
1133mod tests {
1134    use super::*;
1135    use crate::ledger::{InMemoryLedger, LedgerHandle, LedgerStore};
1136
1137    fn test_gateway() -> Gateway {
1138        let routes = RouteTable::from_toml_str(
1139            r#"[routes."fast"]
1140               legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1141        )
1142        .unwrap();
1143        let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1144        let ledger = LedgerHandle::spawn(
1145            Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1146            16,
1147        );
1148        Gateway::builder()
1149            .routes(routes)
1150            .catalog(catalog)
1151            .pricing(PricingTable::default())
1152            .ledger(ledger)
1153            .default_tenant("acme")
1154            .build()
1155            .unwrap()
1156    }
1157
1158    #[tokio::test]
1159    async fn builder_builds_and_lists_aliases() {
1160        let gw = test_gateway();
1161        assert_eq!(gw.model_aliases(), vec!["fast".to_string()]);
1162        assert_eq!(gw.default_tenant, "acme");
1163    }
1164
1165    #[test]
1166    fn builder_requires_components() {
1167        assert!(Gateway::builder().build().is_err());
1168    }
1169
1170    #[tokio::test]
1171    async fn guard_records_usage_on_drop() {
1172        let store = Arc::new(InMemoryLedger::default());
1173        let ledger = LedgerHandle::spawn(store.clone(), 16);
1174        let pricing = Arc::new(PricingTable::default());
1175        {
1176            let mut guard = StreamSideEffects::new(
1177                ledger.clone(),
1178                pricing,
1179                "route".into(),
1180                "tenant".into(),
1181                Attribution::default(),
1182                "p".into(),
1183                "m".into(),
1184                "standard",
1185                "rid".into(),
1186                1,
1187                Instant::now(),
1188            );
1189            guard.observe(&crate::routing::stream::StreamItem::Done {
1190                input_tokens: 3,
1191                output_tokens: 2,
1192                finish_reason: crate::routing::stream::FinishReason::Stop,
1193            });
1194        } // drop -> records
1195        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1196        let rows = store.entries();
1197        assert_eq!(rows.len(), 1);
1198        assert_eq!(rows[0].input_tokens, 3);
1199        assert_eq!(rows[0].output_tokens, 2);
1200        assert_eq!(rows[0].status, "ok");
1201    }
1202
1203    #[tokio::test]
1204    async fn guarded_stream_yields_items_and_records_on_drop() {
1205        use crate::routing::stream::{FinishReason, StreamItem};
1206        use futures::StreamExt;
1207        let store = Arc::new(InMemoryLedger::default());
1208        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1209        let inner = futures::stream::iter(vec![
1210            Ok(StreamItem::Delta("hi".into())),
1211            Ok(StreamItem::Done {
1212                input_tokens: 4,
1213                output_tokens: 2,
1214                finish_reason: FinishReason::Stop,
1215            }),
1216        ])
1217        .boxed();
1218        let guard = StreamSideEffects::new(
1219            ledger,
1220            Arc::new(PricingTable::default()),
1221            "route".into(),
1222            "acme".into(),
1223            Attribution::default(),
1224            "p".into(),
1225            "m".into(),
1226            "standard",
1227            "rid".into(),
1228            1,
1229            std::time::Instant::now(),
1230        );
1231        {
1232            let mut gs = GuardedStream::new(inner, "m".into(), guard);
1233            let mut n = 0;
1234            while let Some(item) = gs.next().await {
1235                item.unwrap();
1236                n += 1;
1237            }
1238            assert_eq!(n, 2);
1239            assert_eq!(gs.model(), "m");
1240        } // drop -> records
1241        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1242        let rows = store.entries();
1243        assert_eq!(rows.len(), 1);
1244        assert_eq!(rows[0].input_tokens, 4);
1245        assert_eq!(rows[0].tenant, "acme");
1246    }
1247
1248    /// Minimal gateway for exercising resolution logic (no upstream needed).
1249    fn resolver_gateway(ai_task_types: &str) -> Gateway {
1250        let routes = RouteTable::from_toml_str(
1251            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1252        )
1253        .unwrap();
1254        Gateway::builder()
1255            .routes(routes)
1256            .catalog(Catalog::for_test(vec![(
1257                "qwen",
1258                "http://127.0.0.1:1/v1".into(),
1259            )]))
1260            .pricing(PricingTable::default())
1261            .ledger(LedgerHandle::spawn(
1262                Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1263                4,
1264            ))
1265            .ai_task_types(
1266                crate::ai_task_type::AiTaskTypeTable::from_toml_str(ai_task_types).unwrap(),
1267            )
1268            .build()
1269            .unwrap()
1270    }
1271
1272    #[tokio::test]
1273    async fn ai_task_type_prefers_the_caller_header_over_the_alias_mapping() {
1274        let gw = resolver_gateway("conversation = [\"fast\"]");
1275        let ctx = RequestCtx {
1276            ai_task_type: Some("caller-supplied".into()),
1277            ..Default::default()
1278        };
1279        assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "caller-supplied");
1280    }
1281
1282    #[tokio::test]
1283    async fn ai_task_type_is_inferred_from_the_route_alias_when_no_header() {
1284        let gw = resolver_gateway("conversation = [\"fast\", \"planning\"]");
1285        let ctx = RequestCtx::default();
1286        assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
1287        assert_eq!(gw.ai_task_type_of(&ctx, "planning"), "conversation");
1288    }
1289
1290    #[tokio::test]
1291    async fn ai_task_type_defaults_to_simple_for_an_unmapped_alias() {
1292        let gw = resolver_gateway("conversation = [\"fast\"]");
1293        assert_eq!(
1294            gw.ai_task_type_of(&RequestCtx::default(), "graph-llm"),
1295            "simple"
1296        );
1297    }
1298
1299    #[tokio::test]
1300    async fn an_empty_ai_task_type_header_falls_back_to_inference() {
1301        let gw = resolver_gateway("conversation = [\"fast\"]");
1302        let ctx = RequestCtx {
1303            ai_task_type: Some(String::new()),
1304            ..Default::default()
1305        };
1306        assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
1307    }
1308
1309    #[tokio::test]
1310    async fn chat_returns_completion_and_records_ledger() {
1311        use wiremock::matchers::{method, path};
1312        use wiremock::{Mock, MockServer, ResponseTemplate};
1313        let mock = MockServer::start().await;
1314        Mock::given(method("POST")).and(path("/v1/chat/completions"))
1315            .respond_with(ResponseTemplate::new(200)
1316                .insert_header("content-type", "text/event-stream")
1317                .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\n\
1318                                  data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2}}\n\n\
1319                                  data: [DONE]\n\n"))
1320            .mount(&mock).await;
1321        let routes = RouteTable::from_toml_str(
1322            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1323        )
1324        .unwrap();
1325        let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
1326        let store = Arc::new(InMemoryLedger::default());
1327        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1328        let gw = Gateway::builder()
1329            .routes(routes)
1330            .catalog(catalog)
1331            .pricing(PricingTable::default())
1332            .ledger(ledger)
1333            .default_tenant("def")
1334            .ai_task_types(
1335                crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
1336                    .unwrap(),
1337            )
1338            .build()
1339            .unwrap();
1340        let req = serde_json::from_value(serde_json::json!(
1341            {"model":"fast","messages":[{"role":"user","content":"hi"}]}))
1342        .unwrap();
1343        let ctx = RequestCtx {
1344            tenant: Some("acme".into()),
1345            workspace: Some("ws-9".into()),
1346            user: Some("user-42".into()),
1347            thread: Some("thread-9".into()),
1348            message: Some("msg-7".into()),
1349            user_task_type: Some("summarisation".into()),
1350            // No override: the ledger row must show the alias-inferred value.
1351            ai_task_type: None,
1352            request_id: Some("corr-123".into()),
1353        };
1354        let c = match gw.chat(req, &ctx).await.unwrap() {
1355            ChatOutcome::Plain(c) => c,
1356            ChatOutcome::Hybrid(_) => panic!("expected plain completion"),
1357        };
1358        assert_eq!(c.content, "hi");
1359        assert_eq!(c.input_tokens, 3);
1360        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1361        let rows = store.entries();
1362        assert_eq!(rows.len(), 1);
1363        assert_eq!(rows[0].tenant, "acme");
1364        assert_eq!(rows[0].user.as_deref(), Some("user-42"));
1365        assert_eq!(rows[0].thread.as_deref(), Some("thread-9"));
1366        assert_eq!(rows[0].message.as_deref(), Some("msg-7"));
1367        assert_eq!(rows[0].user_task_type.as_deref(), Some("summarisation"));
1368        // No x-synapse-ai-task-type header: inferred from the "fast" route alias.
1369        assert_eq!(rows[0].ai_task_type, "conversation");
1370        // A caller-supplied request_id propagates to the ledger row.
1371        assert_eq!(rows[0].request_id, "corr-123");
1372    }
1373
1374    #[tokio::test]
1375    async fn chat_stream_yields_items_and_records() {
1376        use crate::routing::stream::StreamItem;
1377        use futures::StreamExt;
1378        use wiremock::matchers::{method, path};
1379        use wiremock::{Mock, MockServer, ResponseTemplate};
1380        let mock = MockServer::start().await;
1381        Mock::given(method("POST")).and(path("/v1/chat/completions"))
1382            .respond_with(ResponseTemplate::new(200)
1383                .insert_header("content-type", "text/event-stream")
1384                .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
1385                                  data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
1386                                  data: [DONE]\n\n"))
1387            .mount(&mock).await;
1388        let routes = RouteTable::from_toml_str(
1389            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1390        )
1391        .unwrap();
1392        let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
1393        let store = Arc::new(InMemoryLedger::default());
1394        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1395        let gw = Gateway::builder()
1396            .routes(routes)
1397            .catalog(catalog)
1398            .pricing(PricingTable::default())
1399            .ledger(ledger)
1400            .default_tenant("def")
1401            .ai_task_types(
1402                crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
1403                    .unwrap(),
1404            )
1405            .build()
1406            .unwrap();
1407        let req = serde_json::from_value(serde_json::json!(
1408            {"model":"fast","stream":true,"messages":[{"role":"user","content":"hi"}]}))
1409        .unwrap();
1410        let ctx = RequestCtx {
1411            user_task_type: Some("code-review".into()),
1412            ..Default::default()
1413        };
1414        let mut stream = gw.chat_stream(req, &ctx).await.unwrap();
1415        let mut got = false;
1416        while let Some(i) = stream.next().await {
1417            if matches!(i.unwrap(), StreamItem::Delta(ref t) if t == "go") {
1418                got = true;
1419            }
1420        }
1421        drop(stream);
1422        assert!(got);
1423        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1424        let rows = store.entries();
1425        assert_eq!(rows.len(), 1);
1426        assert_eq!(rows[0].user_task_type.as_deref(), Some("code-review"));
1427        assert_eq!(rows[0].ai_task_type, "conversation");
1428    }
1429
1430    // --- Embeddings ---------------------------------------------------------
1431
1432    use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
1433    use async_trait::async_trait;
1434
1435    /// Always fails — exercises the ordered-fallback path.
1436    struct FlakyEmbedder;
1437    #[async_trait]
1438    impl EmbeddingProvider for FlakyEmbedder {
1439        async fn embed(
1440            &self,
1441            _model: &str,
1442            _inputs: &[String],
1443            _dims: u32,
1444        ) -> Result<EmbedOut, GatewayError> {
1445            Err(GatewayError::Upstream {
1446                status: 500,
1447                body: "x".into(),
1448            })
1449        }
1450    }
1451
1452    /// Returns one zero-vector per input (length `dims`) and a fixed token count.
1453    struct GoodEmbedder;
1454    #[async_trait]
1455    impl EmbeddingProvider for GoodEmbedder {
1456        async fn embed(
1457            &self,
1458            _model: &str,
1459            inputs: &[String],
1460            dims: u32,
1461        ) -> Result<EmbedOut, GatewayError> {
1462            Ok(EmbedOut {
1463                vectors: inputs.iter().map(|_| vec![0.0f32; dims as usize]).collect(),
1464                input_tokens: 6,
1465            })
1466        }
1467    }
1468
1469    fn embed_gateway() -> (Gateway, Arc<InMemoryLedger>) {
1470        let embed_routes = crate::routing::embeddings::EmbeddingRouteTable::from_toml_str(
1471            r#"
1472            [embeddings."default-embed"]
1473            dimensions = 4
1474            legs = [
1475              { provider = "flaky", model = "flaky-embed" },
1476              { provider = "good", model = "good-embed" },
1477            ]
1478            "#,
1479        )
1480        .unwrap();
1481        let routes = RouteTable::from_toml_str(
1482            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1483        )
1484        .unwrap();
1485        let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1486        let store = Arc::new(InMemoryLedger::default());
1487        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1488        let gw = Gateway::builder()
1489            .routes(routes)
1490            .catalog(catalog)
1491            .pricing(PricingTable::default())
1492            .ledger(ledger)
1493            .default_tenant("acme")
1494            .embed_routes(embed_routes)
1495            .embedder(
1496                "flaky",
1497                Arc::new(FlakyEmbedder) as Arc<dyn EmbeddingProvider>,
1498            )
1499            .embedder("good", Arc::new(GoodEmbedder) as Arc<dyn EmbeddingProvider>)
1500            .build()
1501            .unwrap();
1502        (gw, store)
1503    }
1504
1505    #[tokio::test]
1506    async fn embed_falls_through_to_good_leg_and_records_usage() {
1507        let (gw, store) = embed_gateway();
1508        let req = EmbeddingRequest {
1509            input: EmbeddingInput::Many(vec!["a".into(), "b".into()]),
1510            model: "default-embed".into(),
1511            dimensions: None,
1512        };
1513        let ctx = RequestCtx {
1514            user_task_type: Some("retrieval".into()),
1515            ..Default::default()
1516        };
1517        let resp = gw.embed(req, ctx).await.unwrap();
1518        assert_eq!(resp.data.len(), 2);
1519        assert!(resp.data.iter().all(|d| d.embedding.len() == 4));
1520        assert_eq!(resp.data[0].index, 0);
1521        assert_eq!(resp.data[1].index, 1);
1522        assert_eq!(resp.usage.prompt_tokens, 6);
1523
1524        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1525        let rows = store.entries();
1526        assert_eq!(rows.len(), 1);
1527        assert_eq!(rows[0].op, "embedding");
1528        assert_eq!(rows[0].lane, "embedding");
1529        assert_eq!(rows[0].output_tokens, 0);
1530        assert_eq!(rows[0].provider, "good");
1531        assert!(rows[0].cost_usd > 0.0);
1532        assert_eq!(rows[0].user_task_type.as_deref(), Some("retrieval"));
1533        assert_eq!(rows[0].ai_task_type, "simple");
1534    }
1535
1536    #[tokio::test]
1537    async fn embed_dimension_mismatch_is_bad_request() {
1538        let (gw, _store) = embed_gateway();
1539        let req = EmbeddingRequest {
1540            input: EmbeddingInput::Many(vec!["a".into()]),
1541            model: "default-embed".into(),
1542            dimensions: Some(8),
1543        };
1544        let err = gw.embed(req, RequestCtx::default()).await.unwrap_err();
1545        assert!(matches!(err, GatewayError::BadRequest(_)));
1546    }
1547
1548    #[tokio::test]
1549    async fn chat_blocks_when_route_policy_refuses() {
1550        use crate::guard::{GuardEngine, GuardrailsConfig};
1551        let routes = RouteTable::from_toml_str(
1552            r#"[routes."fast"]
1553               policy = "strict"
1554               legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1555        )
1556        .unwrap();
1557        let guard = GuardEngine::from_config(
1558            &GuardrailsConfig::from_toml_str(
1559                r#"[guardrails.strict]
1560                   scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#,
1561            )
1562            .unwrap(),
1563        )
1564        .unwrap();
1565        let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1566        let ledger = LedgerHandle::spawn(
1567            Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1568            16,
1569        );
1570        let gw = Gateway::builder()
1571            .routes(routes)
1572            .catalog(catalog)
1573            .pricing(PricingTable::default())
1574            .ledger(ledger)
1575            .guard(guard)
1576            .build()
1577            .unwrap();
1578        let req = serde_json::from_value(serde_json::json!({
1579            "model": "fast",
1580            "messages": [{ "role": "user", "content": "this is forbidden" }]
1581        }))
1582        .unwrap();
1583        let err = gw.chat(req, &RequestCtx::default()).await.unwrap_err();
1584        assert!(matches!(err, GatewayError::ContentBlocked { .. }));
1585    }
1586}