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