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