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