1use 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#[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#[derive(Debug, Clone, Default)]
47pub struct RequestCtx {
48 pub tenant: Option<String>,
49 pub workspace: Option<String>,
50 pub user: Option<String>,
52 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 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 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 #[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 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 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 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 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 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 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
469pub(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, 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 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 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 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 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
607pub 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 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 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
655async 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 } 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 } 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 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 use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
900 use async_trait::async_trait;
901
902 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 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}