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::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#[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#[derive(Debug, Clone, Default)]
45pub struct RequestCtx {
46 pub tenant: Option<String>,
47 pub workspace: Option<String>,
48 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 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 #[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 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 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 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 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 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 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
447pub(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, 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 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 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 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 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
581pub 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 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 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
629async 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 } 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 } 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 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 use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
870 use async_trait::async_trait;
871
872 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 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}