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