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