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;
14use crate::guard::GuardEngine;
15use crate::ledger::{LedgerHandle, UsageEntry};
16use crate::observability::GenAiSpan;
17use crate::pricing::PricingTable;
18use crate::providers::Catalog;
19use crate::routing::classify::{classify, Lane};
20use crate::routing::executor::{
21    execute_buffered_with_timeouts, CommittedStream, Completion, LegError, StreamTimeouts,
22};
23use crate::routing::request::ChatRequest;
24use crate::routing::stream::{Accumulator, FinishReason, StreamItem};
25use crate::routing::table::{ChainLeg, RouteTable};
26use crate::vertex_native::VertexNativeProvider;
27
28/// Embeddable gateway handle. Construct with [`Gateway::builder`].
29#[derive(Clone)]
30pub struct Gateway {
31    pub(crate) routes: Arc<RouteTable>,
32    pub(crate) catalog: Arc<Catalog>,
33    pub(crate) pricing: Arc<PricingTable>,
34    pub(crate) ledger: LedgerHandle,
35    pub(crate) vertex_native: Option<Arc<VertexNativeProvider>>,
36    pub(crate) timeouts: StreamTimeouts,
37    pub(crate) default_tenant: String,
38    pub(crate) embed_routes: Arc<crate::routing::embeddings::EmbeddingRouteTable>,
39    pub(crate) embedders:
40        std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>,
41    pub(crate) embed_default_input_per_mtok: f64,
42    pub(crate) guard: Arc<GuardEngine>,
43}
44
45/// Per-call identity for an in-process request (replaces HTTP headers).
46#[derive(Debug, Clone, Default)]
47pub struct RequestCtx {
48    pub tenant: Option<String>,
49    pub workspace: Option<String>,
50    /// End-user attribution within the tenant (`x-synapse-user` over HTTP).
51    pub user: Option<String>,
52    /// Conversation / agent thread (`x-synapse-thread` over HTTP).
53    pub thread: Option<String>,
54    /// Chat / work message within the thread (`x-synapse-message` over HTTP).
55    pub message: Option<String>,
56    /// Correlation id for the ledger row. When `None`, falls back to `message`
57    /// (so a chat turn correlates with usage), then a generated UUID.
58    /// Embedders (and the HTTP layer) can pass their own so the ledger row and
59    /// any client-facing response id share one value.
60    pub request_id: Option<String>,
61}
62
63impl RequestCtx {
64    /// Prefer an explicit `request_id`, else the chat `message` id, else `None`
65    /// (callers generate a UUID).
66    pub fn resolved_request_id(&self) -> Option<String> {
67        self.request_id
68            .clone()
69            .or_else(|| self.message.clone())
70            .filter(|s| !s.is_empty())
71    }
72}
73
74#[derive(Default)]
75pub struct GatewayBuilder {
76    routes: Option<RouteTable>,
77    catalog: Option<Catalog>,
78    pricing: Option<PricingTable>,
79    ledger: Option<LedgerHandle>,
80    vertex_native: Option<VertexNativeProvider>,
81    timeouts: Option<StreamTimeouts>,
82    default_tenant: Option<String>,
83    embed_routes: Option<crate::routing::embeddings::EmbeddingRouteTable>,
84    embedders:
85        Option<std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>>,
86    embed_default_input_per_mtok: Option<f64>,
87    guard: Option<GuardEngine>,
88}
89
90impl Gateway {
91    pub fn builder() -> GatewayBuilder {
92        GatewayBuilder::default()
93    }
94
95    /// Client-facing model aliases (for `/v1/models` and discovery).
96    pub fn model_aliases(&self) -> Vec<String> {
97        self.routes.aliases()
98    }
99
100    fn resolve_legs(&self, req: &ChatRequest) -> Result<Vec<ChainLeg>, GatewayError> {
101        self.routes
102            .legs(&req.model)
103            .ok_or_else(|| GatewayError::UnknownModel(req.model.clone()))
104            .map(<[ChainLeg]>::to_vec)
105    }
106
107    /// Run the route's guardrail policy over the request input. Falls back to
108    /// the `default` policy; a no-op when neither is configured.
109    fn guard_input(&self, req: &ChatRequest) -> Result<(), GatewayError> {
110        let policy = self.routes.policy_of(&req.model).unwrap_or("default");
111        self.guard.guard(policy, req)
112    }
113
114    fn tenant_of<'a>(&'a self, ctx: &'a RequestCtx) -> &'a str {
115        ctx.tenant.as_deref().unwrap_or(&self.default_tenant)
116    }
117
118    /// Fire cost + ledger + metrics for a completed call (buffered side-effects).
119    #[allow(clippy::too_many_arguments)]
120    fn record(
121        &self,
122        ctx: &RequestCtx,
123        route: &str,
124        lane: &'static str,
125        request_id: &str,
126        c: &Completion,
127        legs: u32,
128        started: Instant,
129    ) {
130        let cost = self
131            .pricing
132            .cost_usd(&c.provider, &c.model, c.input_tokens, c.output_tokens);
133        self.ledger.enqueue(UsageEntry {
134            ts: Utc::now(),
135            tenant: self.tenant_of(ctx).to_string(),
136            workspace: ctx.workspace.clone(),
137            user: ctx.user.clone(),
138            thread: ctx.thread.clone(),
139            message: ctx.message.clone(),
140            route: route.to_string(),
141            provider: c.provider.clone(),
142            model: c.model.clone(),
143            lane: lane.to_string(),
144            input_tokens: c.input_tokens,
145            output_tokens: c.output_tokens,
146            cost_usd: cost,
147            request_id: request_id.to_string(),
148            status: "ok".into(),
149            op: "chat".into(),
150        });
151        let span_lane = if lane == "native" {
152            Lane::NativeVertex
153        } else {
154            Lane::Standard
155        };
156        GenAiSpan::from_completion(
157            c,
158            span_lane,
159            route,
160            self.tenant_of(ctx),
161            ctx.workspace.as_deref(),
162            legs,
163            false,
164        )
165        .emit_metrics(started.elapsed().as_secs_f64());
166    }
167
168    /// Resolve the native Vertex committed stream (shared by chat/chat_stream).
169    ///
170    /// NOTE: the native lane is bounded only by the reqwest client timeout
171    /// (`config.request_timeout`); the first-chunk/idle `StreamTimeouts` are not
172    /// applied here yet — a tracked follow-up alongside native-lane fallback.
173    async fn native_committed(
174        &self,
175        req: &ChatRequest,
176        legs: &[ChainLeg],
177    ) -> Result<CommittedStream, GatewayError> {
178        let leg = legs
179            .iter()
180            .find(|l| l.provider == "vertex")
181            .ok_or_else(|| GatewayError::NativeFeatureUnsupported {
182                feature: "native-vertex".into(),
183                route: req.model.clone(),
184            })?;
185        let provider = self
186            .vertex_native
187            .as_ref()
188            .ok_or_else(|| GatewayError::BadRequest("native vertex lane not configured".into()))?;
189        provider
190            .stream_generate(&leg.model, req, leg.region.as_deref())
191            .await
192            .map(|stream| CommittedStream::single("vertex".into(), leg.model.clone(), stream))
193    }
194
195    /// Streaming in-process chat: commit a leg on the first item; returns a
196    /// `GuardedStream` that yields items and fires side-effects on completion/drop.
197    pub async fn chat_stream(
198        &self,
199        req: ChatRequest,
200        ctx: &RequestCtx,
201    ) -> Result<GuardedStream, GatewayError> {
202        use crate::routing::executor::execute_streaming_with_timeouts;
203        let started = Instant::now();
204        let legs = self.resolve_legs(&req)?;
205        self.guard_input(&req)?;
206        let request_id = ctx
207            .resolved_request_id()
208            .unwrap_or_else(|| Uuid::new_v4().to_string());
209        let (committed, lane_str, legs_attempted) = match classify(&req) {
210            Lane::Standard => (
211                execute_streaming_with_timeouts(
212                    &self.catalog,
213                    &req.model,
214                    &legs,
215                    &req,
216                    self.timeouts,
217                )
218                .await?,
219                "standard",
220                legs.len() as u32,
221            ),
222            Lane::NativeVertex => (self.native_committed(&req, &legs).await?, "native", 1),
223        };
224        let model = committed.model.clone();
225        let guard = StreamSideEffects::new(
226            self.ledger.clone(),
227            self.pricing.clone(),
228            req.model.clone(),
229            self.tenant_of(ctx).to_string(),
230            ctx.workspace.clone(),
231            ctx.user.clone(),
232            ctx.thread.clone(),
233            ctx.message.clone(),
234            committed.provider.clone(),
235            committed.model.clone(),
236            lane_str,
237            request_id,
238            legs_attempted,
239            started,
240        );
241        Ok(GuardedStream::new(committed.stream, model, guard))
242    }
243
244    /// Buffered in-process chat: stream the chain internally, aggregate, fire
245    /// side-effects, return the completion.
246    pub async fn chat(
247        &self,
248        req: ChatRequest,
249        ctx: &RequestCtx,
250    ) -> Result<Completion, GatewayError> {
251        let started = Instant::now();
252        let legs = self.resolve_legs(&req)?;
253        self.guard_input(&req)?;
254        let request_id = ctx
255            .resolved_request_id()
256            .unwrap_or_else(|| Uuid::new_v4().to_string());
257        let (completion, lane_str, legs_n) = match classify(&req) {
258            Lane::Standard => execute_buffered_with_timeouts(
259                &self.catalog,
260                &req.model,
261                &legs,
262                &req,
263                self.timeouts,
264            )
265            .await
266            .map(|c| (c, "standard", legs.len() as u32))?,
267            Lane::NativeVertex => {
268                let committed = self.native_committed(&req, &legs).await?;
269                (collect_committed(committed).await?, "native", 1)
270            }
271        };
272        self.record(
273            ctx,
274            &req.model,
275            lane_str,
276            &request_id,
277            &completion,
278            legs_n,
279            started,
280        );
281        Ok(completion)
282    }
283
284    /// Embed `req.input` against the embedding alias `req.model`, pinning output to
285    /// the alias's declared dimension. Tries each leg in order, falling through on
286    /// upstream error (simple ordered fallback). Records one ledger row on success.
287    ///
288    // TODO(embeddings): wrap legs in resilience breakers like the chat path.
289    pub async fn embed(
290        &self,
291        req: crate::embeddings::EmbeddingRequest,
292        ctx: RequestCtx,
293    ) -> Result<crate::embeddings::EmbeddingResponse, GatewayError> {
294        let alias = req.model.clone();
295        let dims = self
296            .embed_routes
297            .dimensions(&alias)
298            .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?;
299        if let Some(d) = req.dimensions {
300            if d != dims {
301                return Err(GatewayError::BadRequest(format!(
302                    "dimensions {d} does not match embedding alias '{alias}' dimension {dims}"
303                )));
304            }
305        }
306        let inputs = req.input.into_vec();
307        let legs = self
308            .embed_routes
309            .legs(&alias)
310            .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?
311            .to_vec();
312
313        let mut last_err: Option<GatewayError> = None;
314        for leg in &legs {
315            let Some(embedder) = self.embedders.get(&leg.provider) else {
316                continue;
317            };
318            let limit = match leg.provider.as_str() {
319                "vertex" => crate::embeddings::vertex::VERTEX_EMBED_BATCH,
320                _ => crate::embeddings::openai::OPENAI_EMBED_BATCH,
321            };
322            let started = std::time::Instant::now();
323            match self
324                .embed_all_batches(embedder.as_ref(), &leg.model, &inputs, dims, limit)
325                .await
326            {
327                Ok(out) => {
328                    metrics::counter!(
329                        "synapse_embeddings_total",
330                        "route" => alias.clone(),
331                        "model" => leg.model.clone(),
332                        "provider" => leg.provider.clone(),
333                    )
334                    .increment(1);
335                    metrics::histogram!(
336                        "synapse_embedding_duration_seconds",
337                        "route" => alias.clone(),
338                        "model" => leg.model.clone(),
339                        "provider" => leg.provider.clone(),
340                    )
341                    .record(started.elapsed().as_secs_f64());
342                    self.record_embed_usage(&ctx, &alias, leg, out.input_tokens);
343                    return Ok(crate::embeddings::build_response(alias, out));
344                }
345                Err(e) => last_err = Some(e),
346            }
347        }
348        Err(last_err.unwrap_or(GatewayError::AllLegsFailed {
349            route: alias,
350            failures: Vec::new(),
351        }))
352    }
353
354    /// Embed `inputs` in provider-sized batches, concatenating vectors (input order)
355    /// and summing input tokens into a single `EmbedOut`.
356    async fn embed_all_batches(
357        &self,
358        embedder: &dyn crate::embeddings::EmbeddingProvider,
359        model: &str,
360        inputs: &[String],
361        dims: u32,
362        limit: usize,
363    ) -> Result<crate::embeddings::EmbedOut, GatewayError> {
364        let mut vectors = Vec::with_capacity(inputs.len());
365        let mut input_tokens = 0u64;
366        for batch in crate::embeddings::split_batches(inputs, limit) {
367            let out = embedder.embed(model, batch, dims).await?;
368            input_tokens += out.input_tokens;
369            vectors.extend(out.vectors);
370        }
371        Ok(crate::embeddings::EmbedOut {
372            vectors,
373            input_tokens,
374        })
375    }
376
377    /// Fire one cost + ledger row for a completed embedding call (input tokens only).
378    fn record_embed_usage(&self, ctx: &RequestCtx, alias: &str, leg: &ChainLeg, input_tokens: u64) {
379        let cost = self.pricing.embedding_cost_usd(
380            &leg.provider,
381            &leg.model,
382            input_tokens,
383            self.embed_default_input_per_mtok,
384        );
385        self.ledger.enqueue(UsageEntry {
386            ts: Utc::now(),
387            tenant: self.tenant_of(ctx).to_string(),
388            workspace: ctx.workspace.clone(),
389            user: ctx.user.clone(),
390            thread: ctx.thread.clone(),
391            message: ctx.message.clone(),
392            route: alias.to_string(),
393            provider: leg.provider.clone(),
394            model: leg.model.clone(),
395            lane: "embedding".into(),
396            input_tokens,
397            output_tokens: 0,
398            cost_usd: cost,
399            request_id: ctx
400                .resolved_request_id()
401                .unwrap_or_else(|| Uuid::new_v4().to_string()),
402            status: "ok".into(),
403            op: "embedding".into(),
404        });
405    }
406}
407
408impl GatewayBuilder {
409    pub fn routes(mut self, routes: RouteTable) -> Self {
410        self.routes = Some(routes);
411        self
412    }
413    pub fn catalog(mut self, catalog: Catalog) -> Self {
414        self.catalog = Some(catalog);
415        self
416    }
417    pub fn pricing(mut self, pricing: PricingTable) -> Self {
418        self.pricing = Some(pricing);
419        self
420    }
421    pub fn ledger(mut self, ledger: LedgerHandle) -> Self {
422        self.ledger = Some(ledger);
423        self
424    }
425    pub fn vertex_native(mut self, v: Option<VertexNativeProvider>) -> Self {
426        self.vertex_native = v;
427        self
428    }
429    pub fn timeouts(mut self, t: StreamTimeouts) -> Self {
430        self.timeouts = Some(t);
431        self
432    }
433    pub fn default_tenant(mut self, t: impl Into<String>) -> Self {
434        self.default_tenant = Some(t.into());
435        self
436    }
437    pub fn embed_routes(mut self, t: crate::routing::embeddings::EmbeddingRouteTable) -> Self {
438        self.embed_routes = Some(t);
439        self
440    }
441    pub fn embedder(
442        mut self,
443        id: impl Into<String>,
444        e: Arc<dyn crate::embeddings::EmbeddingProvider>,
445    ) -> Self {
446        self.embedders
447            .get_or_insert_with(Default::default)
448            .insert(id.into(), e);
449        self
450    }
451    pub fn embed_default_input_per_mtok(mut self, v: f64) -> Self {
452        self.embed_default_input_per_mtok = Some(v);
453        self
454    }
455    pub fn guard(mut self, guard: GuardEngine) -> Self {
456        self.guard = Some(guard);
457        self
458    }
459
460    pub fn build(self) -> anyhow::Result<Gateway> {
461        Ok(Gateway {
462            routes: Arc::new(
463                self.routes
464                    .ok_or_else(|| anyhow::anyhow!("Gateway: routes required"))?,
465            ),
466            catalog: Arc::new(
467                self.catalog
468                    .ok_or_else(|| anyhow::anyhow!("Gateway: catalog required"))?,
469            ),
470            pricing: Arc::new(
471                self.pricing
472                    .ok_or_else(|| anyhow::anyhow!("Gateway: pricing required"))?,
473            ),
474            ledger: self
475                .ledger
476                .ok_or_else(|| anyhow::anyhow!("Gateway: ledger required"))?,
477            vertex_native: self.vertex_native.map(Arc::new),
478            timeouts: self.timeouts.unwrap_or_default(),
479            default_tenant: self.default_tenant.unwrap_or_else(|| "unattributed".into()),
480            embed_routes: Arc::new(self.embed_routes.unwrap_or_default()),
481            embedders: self.embedders.unwrap_or_default(),
482            embed_default_input_per_mtok: self.embed_default_input_per_mtok.unwrap_or(0.10),
483            guard: Arc::new(self.guard.unwrap_or_else(GuardEngine::empty)),
484        })
485    }
486}
487
488/// Accumulates running usage during a streamed response and fires cost, ledger,
489/// and metrics exactly once when dropped (normal end, error, or client disconnect).
490/// Emitting from `Drop` guarantees the side-effects run on every termination path,
491/// including client disconnect where the response future is cancelled.
492pub(crate) struct StreamSideEffects {
493    ledger: LedgerHandle,
494    pricing: Arc<PricingTable>,
495    route: String,
496    tenant: String,
497    workspace: Option<String>,
498    user: Option<String>,
499    thread: Option<String>,
500    message: Option<String>,
501    provider: String,
502    model: String,
503    lane: &'static str, // "standard" | "native"
504    request_id: String,
505    legs_attempted: u32,
506    started: Instant,
507    input_tokens: u64,
508    output_tokens: u64,
509    status: &'static str,
510    fired: bool,
511}
512
513impl StreamSideEffects {
514    #[allow(clippy::too_many_arguments)]
515    pub(crate) fn new(
516        ledger: LedgerHandle,
517        pricing: Arc<PricingTable>,
518        route: String,
519        tenant: String,
520        workspace: Option<String>,
521        user: Option<String>,
522        thread: Option<String>,
523        message: Option<String>,
524        provider: String,
525        model: String,
526        lane: &'static str,
527        request_id: String,
528        legs_attempted: u32,
529        started: Instant,
530    ) -> Self {
531        Self {
532            ledger,
533            pricing,
534            route,
535            tenant,
536            workspace,
537            user,
538            thread,
539            message,
540            provider,
541            model,
542            lane,
543            request_id,
544            legs_attempted,
545            started,
546            input_tokens: 0,
547            output_tokens: 0,
548            status: "ok",
549            fired: false,
550        }
551    }
552
553    /// Fold a streamed item into the running usage totals.
554    pub(crate) fn observe(&mut self, item: &StreamItem) {
555        if let StreamItem::Done {
556            input_tokens,
557            output_tokens,
558            ..
559        } = item
560        {
561            self.input_tokens = *input_tokens;
562            self.output_tokens = *output_tokens;
563        }
564    }
565
566    /// Mark the response as failed so the ledger row records `status = "error"`.
567    pub(crate) fn mark_error(&mut self) {
568        self.status = "error";
569    }
570}
571
572impl Drop for StreamSideEffects {
573    fn drop(&mut self) {
574        if self.fired {
575            return;
576        }
577        self.fired = true;
578
579        // Cost + ledger (fire-and-forget), mirroring the buffered handler.
580        let cost = self.pricing.cost_usd(
581            &self.provider,
582            &self.model,
583            self.input_tokens,
584            self.output_tokens,
585        );
586        self.ledger.enqueue(UsageEntry {
587            ts: Utc::now(),
588            tenant: self.tenant.clone(),
589            workspace: self.workspace.clone(),
590            user: self.user.clone(),
591            thread: self.thread.clone(),
592            message: self.message.clone(),
593            route: self.route.clone(),
594            provider: self.provider.clone(),
595            model: self.model.clone(),
596            lane: self.lane.to_string(),
597            input_tokens: self.input_tokens,
598            output_tokens: self.output_tokens,
599            cost_usd: cost,
600            request_id: self.request_id.clone(),
601            status: self.status.to_string(),
602            op: "chat".into(),
603        });
604
605        // Metrics via GenAiSpan for parity with the buffered path: build a minimal
606        // Completion from the running fields and emit with `stream: true`.
607        let completion = Completion {
608            provider: self.provider.clone(),
609            model: self.model.clone(),
610            content: String::new(),
611            tool_calls: Vec::new(),
612            finish_reason: FinishReason::Stop,
613            input_tokens: self.input_tokens,
614            output_tokens: self.output_tokens,
615        };
616        let lane = if self.lane == "native" {
617            Lane::NativeVertex
618        } else {
619            Lane::Standard
620        };
621        GenAiSpan::from_completion(
622            &completion,
623            lane,
624            &self.route,
625            &self.tenant,
626            self.workspace.as_deref(),
627            self.legs_attempted,
628            true,
629        )
630        .emit_metrics(self.started.elapsed().as_secs_f64());
631    }
632}
633
634/// A streaming response that yields normalized `StreamItem`s and fires
635/// cost/ledger/metrics exactly once when fully consumed OR dropped early
636/// (via the owned `StreamSideEffects` guard). Transport-independent.
637pub struct GuardedStream {
638    inner: BoxStream<'static, Result<StreamItem, LegError>>,
639    model: String,
640    guard: StreamSideEffects,
641}
642
643impl GuardedStream {
644    pub(crate) fn new(
645        inner: BoxStream<'static, Result<StreamItem, LegError>>,
646        model: String,
647        guard: StreamSideEffects,
648    ) -> Self {
649        Self {
650            inner,
651            model,
652            guard,
653        }
654    }
655
656    /// The model id of the committed leg (for the OpenAI chunk `model` field).
657    pub fn model(&self) -> &str {
658        &self.model
659    }
660}
661
662impl Stream for GuardedStream {
663    type Item = Result<StreamItem, LegError>;
664    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
665        // Both fields are Unpin (BoxStream is Pin<Box<..>>; StreamSideEffects is plain data).
666        let this = self.get_mut();
667        match this.inner.as_mut().poll_next(cx) {
668            Poll::Ready(Some(Ok(item))) => {
669                this.guard.observe(&item);
670                Poll::Ready(Some(Ok(item)))
671            }
672            Poll::Ready(Some(Err(e))) => {
673                this.guard.mark_error();
674                Poll::Ready(Some(Err(e)))
675            }
676            Poll::Ready(None) => Poll::Ready(None),
677            Poll::Pending => Poll::Pending,
678        }
679    }
680}
681
682/// Drain a committed stream into a single `Completion`. ROP: map the per-item
683/// error onto the failure track, `try_fold` items into an `Accumulator`, then map
684/// to a `Completion`.
685async fn collect_committed(committed: CommittedStream) -> Result<Completion, GatewayError> {
686    use futures::TryStreamExt;
687    let CommittedStream {
688        provider,
689        model,
690        stream,
691    } = committed;
692    stream
693        .map_err(|e: LegError| GatewayError::Upstream {
694            status: 502,
695            body: e.to_string(),
696        })
697        .try_fold(Accumulator::default(), |mut acc, item| async move {
698            acc.push(item);
699            Ok(acc)
700        })
701        .await
702        .map(|acc| Completion {
703            provider,
704            model,
705            content: acc.content,
706            tool_calls: acc.tool_calls,
707            finish_reason: acc.finish_reason,
708            input_tokens: acc.input_tokens,
709            output_tokens: acc.output_tokens,
710        })
711}
712
713#[cfg(test)]
714mod tests {
715    use super::*;
716    use crate::ledger::{InMemoryLedger, LedgerHandle, LedgerStore};
717
718    fn test_gateway() -> Gateway {
719        let routes = RouteTable::from_toml_str(
720            r#"[routes."fast"]
721               legs = [{ provider = "qwen", model = "qwen-max" }]"#,
722        )
723        .unwrap();
724        let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
725        let ledger = LedgerHandle::spawn(
726            Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
727            16,
728        );
729        Gateway::builder()
730            .routes(routes)
731            .catalog(catalog)
732            .pricing(PricingTable::default())
733            .ledger(ledger)
734            .default_tenant("acme")
735            .build()
736            .unwrap()
737    }
738
739    #[tokio::test]
740    async fn builder_builds_and_lists_aliases() {
741        let gw = test_gateway();
742        assert_eq!(gw.model_aliases(), vec!["fast".to_string()]);
743        assert_eq!(gw.default_tenant, "acme");
744    }
745
746    #[test]
747    fn builder_requires_components() {
748        assert!(Gateway::builder().build().is_err());
749    }
750
751    #[tokio::test]
752    async fn guard_records_usage_on_drop() {
753        let store = Arc::new(InMemoryLedger::default());
754        let ledger = LedgerHandle::spawn(store.clone(), 16);
755        let pricing = Arc::new(PricingTable::default());
756        {
757            let mut guard = StreamSideEffects::new(
758                ledger.clone(),
759                pricing,
760                "route".into(),
761                "tenant".into(),
762                None,
763                None,
764                None,
765                None,
766                "p".into(),
767                "m".into(),
768                "standard",
769                "rid".into(),
770                1,
771                Instant::now(),
772            );
773            guard.observe(&crate::routing::stream::StreamItem::Done {
774                input_tokens: 3,
775                output_tokens: 2,
776                finish_reason: crate::routing::stream::FinishReason::Stop,
777            });
778        } // drop -> records
779        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
780        let rows = store.entries();
781        assert_eq!(rows.len(), 1);
782        assert_eq!(rows[0].input_tokens, 3);
783        assert_eq!(rows[0].output_tokens, 2);
784        assert_eq!(rows[0].status, "ok");
785    }
786
787    #[tokio::test]
788    async fn guarded_stream_yields_items_and_records_on_drop() {
789        use crate::routing::stream::{FinishReason, StreamItem};
790        use futures::StreamExt;
791        let store = Arc::new(InMemoryLedger::default());
792        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
793        let inner = futures::stream::iter(vec![
794            Ok(StreamItem::Delta("hi".into())),
795            Ok(StreamItem::Done {
796                input_tokens: 4,
797                output_tokens: 2,
798                finish_reason: FinishReason::Stop,
799            }),
800        ])
801        .boxed();
802        let guard = StreamSideEffects::new(
803            ledger,
804            Arc::new(PricingTable::default()),
805            "route".into(),
806            "acme".into(),
807            None,
808            None,
809            None,
810            None,
811            "p".into(),
812            "m".into(),
813            "standard",
814            "rid".into(),
815            1,
816            std::time::Instant::now(),
817        );
818        {
819            let mut gs = GuardedStream::new(inner, "m".into(), guard);
820            let mut n = 0;
821            while let Some(item) = gs.next().await {
822                item.unwrap();
823                n += 1;
824            }
825            assert_eq!(n, 2);
826            assert_eq!(gs.model(), "m");
827        } // drop -> records
828        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
829        let rows = store.entries();
830        assert_eq!(rows.len(), 1);
831        assert_eq!(rows[0].input_tokens, 4);
832        assert_eq!(rows[0].tenant, "acme");
833    }
834
835    #[tokio::test]
836    async fn chat_returns_completion_and_records_ledger() {
837        use wiremock::matchers::{method, path};
838        use wiremock::{Mock, MockServer, ResponseTemplate};
839        let mock = MockServer::start().await;
840        Mock::given(method("POST")).and(path("/v1/chat/completions"))
841            .respond_with(ResponseTemplate::new(200)
842                .insert_header("content-type", "text/event-stream")
843                .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\n\
844                                  data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2}}\n\n\
845                                  data: [DONE]\n\n"))
846            .mount(&mock).await;
847        let routes = RouteTable::from_toml_str(
848            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
849        )
850        .unwrap();
851        let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
852        let store = Arc::new(InMemoryLedger::default());
853        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
854        let gw = Gateway::builder()
855            .routes(routes)
856            .catalog(catalog)
857            .pricing(PricingTable::default())
858            .ledger(ledger)
859            .default_tenant("def")
860            .build()
861            .unwrap();
862        let req = serde_json::from_value(serde_json::json!(
863            {"model":"fast","messages":[{"role":"user","content":"hi"}]}))
864        .unwrap();
865        let ctx = RequestCtx {
866            tenant: Some("acme".into()),
867            user: Some("user-42".into()),
868            thread: Some("thread-9".into()),
869            message: Some("msg-7".into()),
870            request_id: Some("corr-123".into()),
871            ..Default::default()
872        };
873        let c = gw.chat(req, &ctx).await.unwrap();
874        assert_eq!(c.content, "hi");
875        assert_eq!(c.input_tokens, 3);
876        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
877        let rows = store.entries();
878        assert_eq!(rows.len(), 1);
879        assert_eq!(rows[0].tenant, "acme");
880        assert_eq!(rows[0].user.as_deref(), Some("user-42"));
881        assert_eq!(rows[0].thread.as_deref(), Some("thread-9"));
882        assert_eq!(rows[0].message.as_deref(), Some("msg-7"));
883        // A caller-supplied request_id propagates to the ledger row.
884        assert_eq!(rows[0].request_id, "corr-123");
885    }
886
887    #[tokio::test]
888    async fn chat_stream_yields_items_and_records() {
889        use crate::routing::stream::StreamItem;
890        use futures::StreamExt;
891        use wiremock::matchers::{method, path};
892        use wiremock::{Mock, MockServer, ResponseTemplate};
893        let mock = MockServer::start().await;
894        Mock::given(method("POST")).and(path("/v1/chat/completions"))
895            .respond_with(ResponseTemplate::new(200)
896                .insert_header("content-type", "text/event-stream")
897                .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
898                                  data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
899                                  data: [DONE]\n\n"))
900            .mount(&mock).await;
901        let routes = RouteTable::from_toml_str(
902            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
903        )
904        .unwrap();
905        let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
906        let store = Arc::new(InMemoryLedger::default());
907        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
908        let gw = Gateway::builder()
909            .routes(routes)
910            .catalog(catalog)
911            .pricing(PricingTable::default())
912            .ledger(ledger)
913            .default_tenant("def")
914            .build()
915            .unwrap();
916        let req = serde_json::from_value(serde_json::json!(
917            {"model":"fast","stream":true,"messages":[{"role":"user","content":"hi"}]}))
918        .unwrap();
919        let mut stream = gw.chat_stream(req, &RequestCtx::default()).await.unwrap();
920        let mut got = false;
921        while let Some(i) = stream.next().await {
922            if matches!(i.unwrap(), StreamItem::Delta(ref t) if t == "go") {
923                got = true;
924            }
925        }
926        drop(stream);
927        assert!(got);
928        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
929        assert_eq!(store.entries().len(), 1);
930    }
931
932    // --- Embeddings ---------------------------------------------------------
933
934    use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
935    use async_trait::async_trait;
936
937    /// Always fails — exercises the ordered-fallback path.
938    struct FlakyEmbedder;
939    #[async_trait]
940    impl EmbeddingProvider for FlakyEmbedder {
941        async fn embed(
942            &self,
943            _model: &str,
944            _inputs: &[String],
945            _dims: u32,
946        ) -> Result<EmbedOut, GatewayError> {
947            Err(GatewayError::Upstream {
948                status: 500,
949                body: "x".into(),
950            })
951        }
952    }
953
954    /// Returns one zero-vector per input (length `dims`) and a fixed token count.
955    struct GoodEmbedder;
956    #[async_trait]
957    impl EmbeddingProvider for GoodEmbedder {
958        async fn embed(
959            &self,
960            _model: &str,
961            inputs: &[String],
962            dims: u32,
963        ) -> Result<EmbedOut, GatewayError> {
964            Ok(EmbedOut {
965                vectors: inputs.iter().map(|_| vec![0.0f32; dims as usize]).collect(),
966                input_tokens: 6,
967            })
968        }
969    }
970
971    fn embed_gateway() -> (Gateway, Arc<InMemoryLedger>) {
972        let embed_routes = crate::routing::embeddings::EmbeddingRouteTable::from_toml_str(
973            r#"
974            [embeddings."default-embed"]
975            dimensions = 4
976            legs = [
977              { provider = "flaky", model = "flaky-embed" },
978              { provider = "good", model = "good-embed" },
979            ]
980            "#,
981        )
982        .unwrap();
983        let routes = RouteTable::from_toml_str(
984            "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
985        )
986        .unwrap();
987        let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
988        let store = Arc::new(InMemoryLedger::default());
989        let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
990        let gw = Gateway::builder()
991            .routes(routes)
992            .catalog(catalog)
993            .pricing(PricingTable::default())
994            .ledger(ledger)
995            .default_tenant("acme")
996            .embed_routes(embed_routes)
997            .embedder(
998                "flaky",
999                Arc::new(FlakyEmbedder) as Arc<dyn EmbeddingProvider>,
1000            )
1001            .embedder("good", Arc::new(GoodEmbedder) as Arc<dyn EmbeddingProvider>)
1002            .build()
1003            .unwrap();
1004        (gw, store)
1005    }
1006
1007    #[tokio::test]
1008    async fn embed_falls_through_to_good_leg_and_records_usage() {
1009        let (gw, store) = embed_gateway();
1010        let req = EmbeddingRequest {
1011            input: EmbeddingInput::Many(vec!["a".into(), "b".into()]),
1012            model: "default-embed".into(),
1013            dimensions: None,
1014        };
1015        let resp = gw.embed(req, RequestCtx::default()).await.unwrap();
1016        assert_eq!(resp.data.len(), 2);
1017        assert!(resp.data.iter().all(|d| d.embedding.len() == 4));
1018        assert_eq!(resp.data[0].index, 0);
1019        assert_eq!(resp.data[1].index, 1);
1020        assert_eq!(resp.usage.prompt_tokens, 6);
1021
1022        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1023        let rows = store.entries();
1024        assert_eq!(rows.len(), 1);
1025        assert_eq!(rows[0].op, "embedding");
1026        assert_eq!(rows[0].lane, "embedding");
1027        assert_eq!(rows[0].output_tokens, 0);
1028        assert_eq!(rows[0].provider, "good");
1029        assert!(rows[0].cost_usd > 0.0);
1030    }
1031
1032    #[tokio::test]
1033    async fn embed_dimension_mismatch_is_bad_request() {
1034        let (gw, _store) = embed_gateway();
1035        let req = EmbeddingRequest {
1036            input: EmbeddingInput::Many(vec!["a".into()]),
1037            model: "default-embed".into(),
1038            dimensions: Some(8),
1039        };
1040        let err = gw.embed(req, RequestCtx::default()).await.unwrap_err();
1041        assert!(matches!(err, GatewayError::BadRequest(_)));
1042    }
1043
1044    #[tokio::test]
1045    async fn chat_blocks_when_route_policy_refuses() {
1046        use crate::guard::{GuardEngine, GuardrailsConfig};
1047        let routes = RouteTable::from_toml_str(
1048            r#"[routes."fast"]
1049               policy = "strict"
1050               legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1051        )
1052        .unwrap();
1053        let guard = GuardEngine::from_config(
1054            &GuardrailsConfig::from_toml_str(
1055                r#"[guardrails.strict]
1056                   scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#,
1057            )
1058            .unwrap(),
1059        )
1060        .unwrap();
1061        let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1062        let ledger = LedgerHandle::spawn(
1063            Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1064            16,
1065        );
1066        let gw = Gateway::builder()
1067            .routes(routes)
1068            .catalog(catalog)
1069            .pricing(PricingTable::default())
1070            .ledger(ledger)
1071            .guard(guard)
1072            .build()
1073            .unwrap();
1074        let req = serde_json::from_value(serde_json::json!({
1075            "model": "fast",
1076            "messages": [{ "role": "user", "content": "this is forbidden" }]
1077        }))
1078        .unwrap();
1079        let err = gw.chat(req, &RequestCtx::default()).await.unwrap_err();
1080        assert!(matches!(err, GatewayError::ContentBlocked { .. }));
1081    }
1082}