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