1use std::borrow::Cow;
5use std::pin::Pin;
6use std::sync::Arc;
7use std::task::{Context, Poll};
8use std::time::Instant;
9
10use chrono::Utc;
11use futures::stream::{BoxStream, Stream};
12use tap::Tap;
13use uuid::Uuid;
14
15use crate::error::{GatewayError, LegFailure};
16use crate::guard::GuardEngine;
17use crate::jev_native::JevNativeProvider;
18use crate::ledger::{LedgerHandle, UsageEntry};
19use crate::observability::GenAiSpan;
20use crate::pricing::PricingTable;
21use crate::providers::Catalog;
22use crate::routing::classify::{classify, vertex_triggers, Lane};
23use crate::routing::effort::Effort;
24use crate::routing::executor::{
25 execute_buffered_with_timeouts, CommittedStream, Completion, LegError, StreamTimeouts,
26};
27use crate::routing::jev_router::RoutingReport;
28use crate::routing::request::{ChatRequest, VertexExt};
29use crate::routing::stream::{Accumulator, FinishReason, StreamItem};
30use crate::routing::table::{ChainLeg, RouteTable};
31use crate::telemetry::GatewayMetrics;
32use crate::vertex_native::VertexNativeProvider;
33
34#[derive(Clone)]
36pub struct Gateway {
37 pub(crate) routes: Arc<RouteTable>,
38 pub(crate) catalog: Arc<Catalog>,
39 pub(crate) pricing: Arc<PricingTable>,
40 pub(crate) ledger: LedgerHandle,
41 pub(crate) vertex_native: Option<Arc<VertexNativeProvider>>,
42 pub(crate) jev_native: Option<Arc<JevNativeProvider>>,
43 pub(crate) timeouts: StreamTimeouts,
44 pub(crate) default_tenant: String,
45 pub(crate) embed_routes: Arc<crate::routing::embeddings::EmbeddingRouteTable>,
46 pub(crate) embedders:
47 std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>,
48 pub(crate) embed_default_input_per_mtok: f64,
49 pub(crate) guard: Arc<GuardEngine>,
50 pub(crate) ai_task_types: Arc<crate::ai_task_type::AiTaskTypeTable>,
51 pub(crate) metrics: Arc<GatewayMetrics>,
52}
53
54#[derive(Debug, Clone, Default)]
56pub struct RequestCtx {
57 pub tenant: Option<String>,
58 pub workspace: Option<String>,
59 pub user: Option<String>,
61 pub thread: Option<String>,
63 pub message: Option<String>,
65 pub user_task_type: Option<String>,
69 pub ai_task_type: Option<String>,
73 pub request_id: Option<String>,
78}
79
80impl RequestCtx {
81 pub fn resolved_request_id(&self) -> Option<String> {
84 self.request_id
85 .clone()
86 .or_else(|| self.message.clone())
87 .filter(|s| !s.is_empty())
88 }
89}
90
91pub(crate) enum JevAttempt {
93 Decided {
96 completion: Completion,
97 answers: serde_json::Value,
98 },
99 Exhausted(Vec<LegFailure>),
102}
103
104#[derive(Debug)]
107pub enum ChatOutcome {
108 Plain(Completion),
109 Hybrid(HybridOutcome),
110}
111
112#[derive(Debug)]
115pub struct HybridOutcome {
116 pub model: String,
118 pub answers: serde_json::Value,
120 pub survivors: Vec<String>,
122 pub degraded: bool,
124 pub extraction_ran: bool,
126 pub extractions: serde_json::Map<String, serde_json::Value>,
128 pub input_tokens: u64,
130 pub output_tokens: u64,
131}
132
133#[derive(Debug, Clone)]
139pub struct Attribution {
140 pub workspace: Option<String>,
141 pub user: Option<String>,
142 pub thread: Option<String>,
143 pub message: Option<String>,
144 pub user_task_type: Option<String>,
145 pub ai_task_type: String,
147}
148
149impl Default for Attribution {
150 fn default() -> Self {
151 Self {
152 workspace: None,
153 user: None,
154 thread: None,
155 message: None,
156 user_task_type: None,
157 ai_task_type: crate::ai_task_type::DEFAULT_AI_TASK_TYPE.to_string(),
158 }
159 }
160}
161
162#[derive(Default)]
163pub struct GatewayBuilder {
164 routes: Option<RouteTable>,
165 catalog: Option<Catalog>,
166 pricing: Option<PricingTable>,
167 ledger: Option<LedgerHandle>,
168 vertex_native: Option<VertexNativeProvider>,
169 jev_native: Option<JevNativeProvider>,
170 timeouts: Option<StreamTimeouts>,
171 default_tenant: Option<String>,
172 embed_routes: Option<crate::routing::embeddings::EmbeddingRouteTable>,
173 embedders:
174 Option<std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>>,
175 embed_default_input_per_mtok: Option<f64>,
176 guard: Option<GuardEngine>,
177 ai_task_types: Option<crate::ai_task_type::AiTaskTypeTable>,
178 metrics: Option<Arc<GatewayMetrics>>,
179}
180
181impl Gateway {
182 pub fn builder() -> GatewayBuilder {
183 GatewayBuilder::default()
184 }
185
186 pub fn model_aliases(&self) -> Vec<String> {
188 self.routes.aliases()
189 }
190
191 pub(crate) fn guard_input(&self, req: &ChatRequest) -> Result<(), GatewayError> {
194 let policy = self.routes.policy_of(&req.model).unwrap_or("default");
195 self.guard.guard(policy, req)
196 }
197
198 pub(crate) fn tenant_of<'a>(&'a self, ctx: &'a RequestCtx) -> &'a str {
199 ctx.tenant.as_deref().unwrap_or(&self.default_tenant)
200 }
201
202 pub(crate) fn ai_task_type_of(&self, ctx: &RequestCtx, alias: &str) -> String {
205 ctx.ai_task_type
206 .as_deref()
207 .filter(|s| !s.is_empty())
208 .unwrap_or_else(|| self.ai_task_types.resolve(alias))
209 .to_string()
210 }
211
212 pub(crate) fn attribution_of(&self, ctx: &RequestCtx, alias: &str) -> Attribution {
215 Attribution {
216 workspace: ctx.workspace.clone(),
217 user: ctx.user.clone(),
218 thread: ctx.thread.clone(),
219 message: ctx.message.clone(),
220 user_task_type: ctx.user_task_type.clone(),
221 ai_task_type: self.ai_task_type_of(ctx, alias),
222 }
223 }
224
225 #[allow(clippy::too_many_arguments)]
227 pub(crate) fn record(
228 &self,
229 ctx: &RequestCtx,
230 route: &str,
231 lane: &'static str,
232 request_id: &str,
233 c: &Completion,
234 legs: u32,
235 started: Instant,
236 ) {
237 let cost = self
238 .pricing
239 .cost_usd(&c.provider, &c.model, c.input_tokens, c.output_tokens);
240 let attr = self.attribution_of(ctx, route);
241 self.ledger.enqueue(UsageEntry {
242 ts: Utc::now(),
243 tenant: self.tenant_of(ctx).to_string(),
244 workspace: attr.workspace,
245 user: attr.user,
246 thread: attr.thread,
247 message: attr.message,
248 route: route.to_string(),
249 provider: c.provider.clone(),
250 model: c.model.clone(),
251 lane: lane.to_string(),
252 input_tokens: c.input_tokens,
253 output_tokens: c.output_tokens,
254 cost_usd: cost,
255 request_id: request_id.to_string(),
256 status: "ok".into(),
257 op: "chat".into(),
258 user_task_type: attr.user_task_type,
259 ai_task_type: attr.ai_task_type,
260 });
261 let span_lane = match lane {
262 "native" => Lane::NativeVertex,
263 "jev" => Lane::Jev,
264 _ => Lane::Standard,
265 };
266 GenAiSpan::from_completion(
267 c,
268 span_lane,
269 route,
270 self.tenant_of(ctx),
271 ctx.workspace.as_deref(),
272 legs,
273 false,
274 )
275 .emit_metrics(&self.metrics, started.elapsed().as_secs_f64());
276 }
277
278 pub(crate) async fn native_committed(
288 &self,
289 req: &ChatRequest,
290 legs: &[ChainLeg],
291 ) -> Result<CommittedStream, GatewayError> {
292 let provider = self
293 .vertex_native
294 .as_ref()
295 .ok_or_else(|| GatewayError::BadRequest("native vertex lane not configured".into()))?;
296
297 let vertex_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "vertex").collect();
298 if vertex_legs.is_empty() {
299 return Err(GatewayError::NativeFeatureUnsupported {
300 feature: "native-vertex".into(),
301 route: req.model.clone(),
302 });
303 }
304
305 let mut failures: Vec<LegFailure> = Vec::new();
306 let mut last_retryable: Option<GatewayError> = None;
307 for leg in &vertex_legs {
308 match provider
309 .stream_generate(
310 &leg.model,
311 &with_leg_thinking(req, leg),
312 leg.region.as_deref(),
313 )
314 .await
315 {
316 Ok(stream) => {
317 return Ok(CommittedStream::single(
318 "vertex".into(),
319 leg.model.clone(),
320 stream,
321 ));
322 }
323 Err(e) if native_start_retryable(&e) => {
324 failures.push(LegFailure {
325 provider: leg.provider.clone(),
326 model: leg.model.clone(),
327 message: e.to_string(),
328 });
329 last_retryable = Some(e);
330 }
331 Err(e) => return Err(e),
332 }
333 }
334 Err(
335 last_retryable.unwrap_or_else(|| GatewayError::AllLegsFailed {
336 route: req.model.clone(),
337 failures,
338 }),
339 )
340 }
341
342 pub(crate) async fn jev_attempt(
349 &self,
350 req: &ChatRequest,
351 legs: &[ChainLeg],
352 ) -> Result<JevAttempt, GatewayError> {
353 let jev_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "typesafe").collect();
354 if jev_legs.is_empty() {
355 return Ok(JevAttempt::Exhausted(Vec::new()));
356 }
357 let ext = req
358 .jev
359 .as_ref()
360 .filter(|j| !j.questions.is_empty())
361 .ok_or_else(|| {
362 GatewayError::BadRequest(format!(
363 "route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
364 req.model
365 ))
366 })?;
367 let provider = self.jev_native.as_ref().ok_or_else(|| {
368 GatewayError::BadRequest(
369 "typesafe legs require TYPESAFE_API_KEY to be configured".into(),
370 )
371 })?;
372
373 let state = ext
374 .state
375 .clone()
376 .or_else(|| serde_json::to_string(&req.messages).ok())
377 .unwrap_or_default();
378 let mut body = serde_json::json!({
379 "state": state,
380 "questions": ext.questions.clone(),
381 });
382
383 let mut failures: Vec<LegFailure> = Vec::new();
384 for leg in jev_legs {
385 let model = if leg.model.is_empty() {
386 crate::jev_native::DEFAULT_MODEL.to_string()
387 } else {
388 leg.model.clone()
389 };
390 body["model"] = serde_json::Value::String(model.clone());
391
392 let resp = match provider.evaluate(body.clone()).await {
393 Ok(r) => r,
394 Err(e) if native_start_retryable(&e) => {
395 failures.push(LegFailure {
396 provider: leg.provider.clone(),
397 model,
398 message: e.to_string(),
399 });
400 continue;
401 }
402 Err(e) => return Err(e),
403 };
404
405 let status = resp.status();
406 let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
407 status: 502,
408 body: e.to_string(),
409 })?;
410 if status.is_success() {
411 let value: serde_json::Value =
412 serde_json::from_slice(&bytes).map_err(|e| GatewayError::Upstream {
413 status: 502,
414 body: format!("jev response is not JSON: {e}"),
415 })?;
416 return Ok(JevAttempt::Decided {
417 answers: value["answers"].clone(),
418 completion: Completion {
419 provider: "typesafe".into(),
420 model: value["model"].as_str().unwrap_or(&model).to_string(),
421 content: serde_json::to_string(&value["answers"])
422 .unwrap_or_else(|_| "{}".into()),
423 tool_calls: Vec::new(),
424 finish_reason: FinishReason::Stop,
425 input_tokens: value["usage"]["input_tokens"].as_u64().unwrap_or(0),
426 output_tokens: value["usage"]["output_tokens"].as_u64().unwrap_or(0),
427 },
428 });
429 }
430
431 let message = String::from_utf8_lossy(&bytes).into_owned();
432 if status.is_server_error()
433 || status == reqwest::StatusCode::TOO_MANY_REQUESTS
434 || status == reqwest::StatusCode::REQUEST_TIMEOUT
435 {
436 failures.push(LegFailure {
437 provider: leg.provider.clone(),
438 model,
439 message: format!("{status}: {message}"),
440 });
441 continue;
442 }
443 return Err(GatewayError::BadRequest(format!(
444 "typesafe {}: {message}",
445 status.as_u16()
446 )));
447 }
448 Ok(JevAttempt::Exhausted(failures))
449 }
450
451 fn jev_committed(c: Completion) -> CommittedStream {
454 let items: Vec<Result<StreamItem, LegError>> = vec![
455 Ok(StreamItem::Delta(c.content.clone())),
456 Ok(StreamItem::Done {
457 input_tokens: c.input_tokens,
458 output_tokens: c.output_tokens,
459 finish_reason: c.finish_reason,
460 }),
461 ];
462 CommittedStream::single("typesafe".into(), c.model, futures::stream::iter(items))
463 }
464
465 pub async fn chat_stream(
468 &self,
469 req: ChatRequest,
470 ctx: &RequestCtx,
471 ) -> Result<GuardedStream, GatewayError> {
472 use crate::routing::executor::execute_streaming_with_timeouts;
473 let started = Instant::now();
474 let request_id = ctx
475 .resolved_request_id()
476 .unwrap_or_else(|| Uuid::new_v4().to_string());
477 let plan = self.plan_route(&req, ctx, &request_id).await?;
478 let legs = plan.chain();
479 let vertex_leg_count = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
480 require_jev_block(&req, &legs)?;
481 if let Err(message) = crate::routing::jev_extract::validate_extract(
482 &req,
483 legs.iter().any(|l| l.provider != "typesafe"),
484 legs.iter().any(|l| l.provider == "vertex"),
485 ) {
486 return Err(GatewayError::BadRequest(message));
487 }
488 let (committed, lane_str, legs_attempted) = match classify(&req) {
489 Lane::Standard => (
490 execute_streaming_with_timeouts(
491 &self.catalog,
492 &req.model,
493 &legs,
494 &req,
495 self.timeouts,
496 )
497 .await?,
498 "standard",
499 legs.len() as u32,
500 ),
501 Lane::NativeVertex => (
502 self.native_committed(&req, &legs).await?,
503 "native",
504 vertex_leg_count.max(1),
505 ),
506 Lane::Jev => match self.jev_attempt(&req, &legs).await? {
507 JevAttempt::Decided { completion, .. } => (
508 Self::jev_committed(completion),
509 "jev",
510 legs.iter()
511 .filter(|l| l.provider == "typesafe")
512 .count()
513 .max(1) as u32,
514 ),
515 JevAttempt::Exhausted(failures) => {
516 let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
517 if rest.is_empty() {
518 return Err(GatewayError::AllLegsFailed {
519 route: req.model.clone(),
520 failures,
521 });
522 }
523 if vertex_triggers(&req) {
524 let vertex_n =
525 rest.iter().filter(|l| l.provider == "vertex").count() as u32;
526 (
527 self.native_committed(&req, &rest).await?,
528 "native",
529 vertex_n.max(1),
530 )
531 } else {
532 (
533 execute_streaming_with_timeouts(
534 &self.catalog,
535 &req.model,
536 &rest,
537 &req,
538 self.timeouts,
539 )
540 .await?,
541 "standard",
542 rest.len() as u32,
543 )
544 }
545 }
546 },
547 };
548 let model = committed.model.clone();
549 let routing = plan.report_for(Some((
550 committed.provider.as_str(),
551 committed.model.as_str(),
552 )));
553 let guard = StreamSideEffects::new(
554 self.ledger.clone(),
555 self.pricing.clone(),
556 self.metrics.clone(),
557 req.model.clone(),
558 self.tenant_of(ctx).to_string(),
559 self.attribution_of(ctx, &req.model),
560 committed.provider.clone(),
561 committed.model.clone(),
562 lane_str,
563 request_id,
564 legs_attempted,
565 started,
566 );
567 Ok(GuardedStream::new(committed.stream, model, guard, routing))
568 }
569
570 pub async fn chat(
573 &self,
574 req: ChatRequest,
575 ctx: &RequestCtx,
576 ) -> Result<ChatOutcome, GatewayError> {
577 self.chat_routed(req, ctx).await.map(|(outcome, _)| outcome)
578 }
579
580 pub async fn chat_routed(
583 &self,
584 req: ChatRequest,
585 ctx: &RequestCtx,
586 ) -> Result<(ChatOutcome, RoutingReport), GatewayError> {
587 let started = Instant::now();
588 let request_id = ctx
589 .resolved_request_id()
590 .unwrap_or_else(|| Uuid::new_v4().to_string());
591 let plan = self.plan_route(&req, ctx, &request_id).await?;
592 let legs = plan.chain();
593 require_jev_block(&req, &legs)?;
594 if let Err(message) = crate::routing::jev_extract::validate_extract(
595 &req,
596 legs.iter().any(|l| l.provider != "typesafe"),
597 legs.iter().any(|l| l.provider == "vertex"),
598 ) {
599 return Err(GatewayError::BadRequest(message));
600 }
601 if req.jev.as_ref().and_then(|j| j.extract.as_ref()).is_some() {
602 let outcome = self
603 .chat_hybrid(req, ctx, &legs, started, &request_id)
604 .await?;
605 return Ok((ChatOutcome::Hybrid(outcome), plan.report_for(None)));
606 }
607 let (completion, lane_str, legs_n) = match classify(&req) {
608 Lane::Standard => execute_buffered_with_timeouts(
609 &self.catalog,
610 &req.model,
611 &legs,
612 &req,
613 self.timeouts,
614 )
615 .await
616 .map(|c| (c, "standard", legs.len() as u32))?,
617 Lane::NativeVertex => {
618 let vertex_n = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
619 let committed = self.native_committed(&req, &legs).await?;
620 (
621 collect_committed(committed).await?,
622 "native",
623 vertex_n.max(1),
624 )
625 }
626 Lane::Jev => match self.jev_attempt(&req, &legs).await? {
627 JevAttempt::Decided { completion, .. } => (
628 completion,
629 "jev",
630 legs.iter()
631 .filter(|l| l.provider == "typesafe")
632 .count()
633 .max(1) as u32,
634 ),
635 JevAttempt::Exhausted(failures) => {
636 let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
637 if rest.is_empty() {
638 return Err(GatewayError::AllLegsFailed {
639 route: req.model.clone(),
640 failures,
641 });
642 }
643 if vertex_triggers(&req) {
644 let committed = self.native_committed(&req, &rest).await?;
645 let vertex_n =
646 rest.iter().filter(|l| l.provider == "vertex").count() as u32;
647 (
648 collect_committed(committed).await?,
649 "native",
650 vertex_n.max(1),
651 )
652 } else {
653 execute_buffered_with_timeouts(
654 &self.catalog,
655 &req.model,
656 &rest,
657 &req,
658 self.timeouts,
659 )
660 .await
661 .map(|c| (c, "standard", rest.len() as u32))?
662 }
663 }
664 },
665 };
666 self.record(
667 ctx,
668 &req.model,
669 lane_str,
670 &request_id,
671 &completion,
672 legs_n,
673 started,
674 );
675 let routing = plan.report_for(Some((
676 completion.provider.as_str(),
677 completion.model.as_str(),
678 )));
679 Ok((ChatOutcome::Plain(completion), routing))
680 }
681
682 pub async fn embed(
688 &self,
689 req: crate::embeddings::EmbeddingRequest,
690 ctx: RequestCtx,
691 ) -> Result<crate::embeddings::EmbeddingResponse, GatewayError> {
692 let alias = req.model.clone();
693 let dims = self
694 .embed_routes
695 .dimensions(&alias)
696 .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?;
697 if let Some(d) = req.dimensions {
698 if d != dims {
699 return Err(GatewayError::BadRequest(format!(
700 "dimensions {d} does not match embedding alias '{alias}' dimension {dims}"
701 )));
702 }
703 }
704 let inputs = req.input.into_vec();
705 let legs = self
706 .embed_routes
707 .legs(&alias)
708 .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?
709 .to_vec();
710
711 let mut last_err: Option<GatewayError> = None;
712 for leg in &legs {
713 let Some(embedder) = self.embedders.get(&leg.provider) else {
714 continue;
715 };
716 let limit = match leg.provider.as_str() {
717 "vertex" => crate::embeddings::vertex::VERTEX_EMBED_BATCH,
718 _ => crate::embeddings::openai::OPENAI_EMBED_BATCH,
719 };
720 let started = std::time::Instant::now();
721 match self
722 .embed_all_batches(embedder.as_ref(), &leg.model, &inputs, dims, limit)
723 .await
724 {
725 Ok(out) => {
726 self.metrics.embedding(
727 &alias,
728 &leg.model,
729 &leg.provider,
730 started.elapsed().as_secs_f64(),
731 );
732 self.record_embed_usage(&ctx, &alias, leg, out.input_tokens);
733 return Ok(crate::embeddings::build_response(alias, out));
734 }
735 Err(e) => last_err = Some(e),
736 }
737 }
738 Err(last_err.unwrap_or(GatewayError::AllLegsFailed {
739 route: alias,
740 failures: Vec::new(),
741 }))
742 }
743
744 async fn embed_all_batches(
747 &self,
748 embedder: &dyn crate::embeddings::EmbeddingProvider,
749 model: &str,
750 inputs: &[String],
751 dims: u32,
752 limit: usize,
753 ) -> Result<crate::embeddings::EmbedOut, GatewayError> {
754 let mut vectors = Vec::with_capacity(inputs.len());
755 let mut input_tokens = 0u64;
756 for batch in crate::embeddings::split_batches(inputs, limit) {
757 let out = embedder.embed(model, batch, dims).await?;
758 input_tokens += out.input_tokens;
759 vectors.extend(out.vectors);
760 }
761 Ok(crate::embeddings::EmbedOut {
762 vectors,
763 input_tokens,
764 })
765 }
766
767 fn record_embed_usage(&self, ctx: &RequestCtx, alias: &str, leg: &ChainLeg, input_tokens: u64) {
769 let cost = self.pricing.embedding_cost_usd(
770 &leg.provider,
771 &leg.model,
772 input_tokens,
773 self.embed_default_input_per_mtok,
774 );
775 let attr = self.attribution_of(ctx, alias);
776 self.ledger.enqueue(UsageEntry {
777 ts: Utc::now(),
778 tenant: self.tenant_of(ctx).to_string(),
779 workspace: attr.workspace,
780 user: attr.user,
781 thread: attr.thread,
782 message: attr.message,
783 route: alias.to_string(),
784 provider: leg.provider.clone(),
785 model: leg.model.clone(),
786 lane: "embedding".into(),
787 input_tokens,
788 output_tokens: 0,
789 cost_usd: cost,
790 request_id: ctx
791 .resolved_request_id()
792 .unwrap_or_else(|| Uuid::new_v4().to_string()),
793 status: "ok".into(),
794 op: "embedding".into(),
795 user_task_type: attr.user_task_type,
796 ai_task_type: attr.ai_task_type,
797 });
798 }
799}
800
801fn native_start_retryable(e: &GatewayError) -> bool {
804 match e {
805 GatewayError::Upstream { status, .. } => {
806 *status >= 500
807 || *status == reqwest::StatusCode::TOO_MANY_REQUESTS.as_u16()
808 || *status == reqwest::StatusCode::REQUEST_TIMEOUT.as_u16()
809 }
810 GatewayError::UpstreamTimeout => true,
811 _ => false,
812 }
813}
814
815fn require_jev_block(req: &ChatRequest, legs: &[ChainLeg]) -> Result<(), GatewayError> {
819 let needs_block = legs.iter().any(|l| l.provider == "typesafe");
820 let has_block = req.jev.as_ref().is_some_and(|j| !j.questions.is_empty());
821 match (needs_block, has_block) {
822 (true, false) => Err(GatewayError::BadRequest(format!(
823 "route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
824 req.model
825 ))),
826 _ => Ok(()),
827 }
828}
829
830fn non_typesafe_legs(legs: &[ChainLeg]) -> Vec<ChainLeg> {
833 legs.iter()
834 .filter(|l| l.provider != "typesafe")
835 .cloned()
836 .collect()
837}
838
839fn with_leg_thinking<'a>(req: &'a ChatRequest, leg: &ChainLeg) -> Cow<'a, ChatRequest> {
842 let client_set = req
843 .vertex
844 .as_ref()
845 .is_some_and(|v| v.thinking_config.is_some());
846 match (client_set, leg.effort.and_then(Effort::thinking_budget)) {
847 (false, Some(budget)) => Cow::Owned(ChatRequest {
848 vertex: Some(VertexExt {
849 thinking_config: Some(serde_json::json!({ "thinkingBudget": budget })),
850 ..req.vertex.clone().unwrap_or_default()
851 }),
852 ..req.clone()
853 }),
854 _ => Cow::Borrowed(req),
855 }
856}
857
858impl GatewayBuilder {
859 pub fn routes(mut self, routes: RouteTable) -> Self {
860 self.routes = Some(routes);
861 self
862 }
863 pub fn catalog(mut self, catalog: Catalog) -> Self {
864 self.catalog = Some(catalog);
865 self
866 }
867 pub fn pricing(mut self, pricing: PricingTable) -> Self {
868 self.pricing = Some(pricing);
869 self
870 }
871 pub fn ledger(mut self, ledger: LedgerHandle) -> Self {
872 self.ledger = Some(ledger);
873 self
874 }
875 pub fn vertex_native(mut self, v: Option<VertexNativeProvider>) -> Self {
876 self.vertex_native = v;
877 self
878 }
879 pub fn jev_native(mut self, v: Option<JevNativeProvider>) -> Self {
880 self.jev_native = v;
881 self
882 }
883 pub fn timeouts(mut self, t: StreamTimeouts) -> Self {
884 self.timeouts = Some(t);
885 self
886 }
887 pub fn default_tenant(mut self, t: impl Into<String>) -> Self {
888 self.default_tenant = Some(t.into());
889 self
890 }
891 pub fn embed_routes(mut self, t: crate::routing::embeddings::EmbeddingRouteTable) -> Self {
892 self.embed_routes = Some(t);
893 self
894 }
895 pub fn embedder(
896 mut self,
897 id: impl Into<String>,
898 e: Arc<dyn crate::embeddings::EmbeddingProvider>,
899 ) -> Self {
900 self.embedders
901 .get_or_insert_with(Default::default)
902 .insert(id.into(), e);
903 self
904 }
905 pub fn embed_default_input_per_mtok(mut self, v: f64) -> Self {
906 self.embed_default_input_per_mtok = Some(v);
907 self
908 }
909 pub fn guard(mut self, guard: GuardEngine) -> Self {
910 self.guard = Some(guard);
911 self
912 }
913 pub fn ai_task_types(mut self, t: crate::ai_task_type::AiTaskTypeTable) -> Self {
914 self.ai_task_types = Some(t);
915 self
916 }
917 pub fn metrics(mut self, metrics: Arc<GatewayMetrics>) -> Self {
920 self.metrics = Some(metrics);
921 self
922 }
923
924 pub fn build(self) -> anyhow::Result<Gateway> {
925 let metrics = self.metrics.unwrap_or_else(GatewayMetrics::noop);
926 Ok(Gateway {
927 routes: Arc::new(
928 self.routes
929 .ok_or_else(|| anyhow::anyhow!("Gateway: routes required"))?,
930 ),
931 catalog: Arc::new(
932 self.catalog
933 .ok_or_else(|| anyhow::anyhow!("Gateway: catalog required"))?
934 .tap(|c| c.attach_metrics(&metrics)),
935 ),
936 pricing: Arc::new(
937 self.pricing
938 .ok_or_else(|| anyhow::anyhow!("Gateway: pricing required"))?,
939 ),
940 ledger: self
941 .ledger
942 .ok_or_else(|| anyhow::anyhow!("Gateway: ledger required"))?,
943 vertex_native: self.vertex_native.map(Arc::new),
944 jev_native: self.jev_native.map(Arc::new),
945 timeouts: self.timeouts.unwrap_or_default(),
946 default_tenant: self.default_tenant.unwrap_or_else(|| "unattributed".into()),
947 embed_routes: Arc::new(self.embed_routes.unwrap_or_default()),
948 embedders: self.embedders.unwrap_or_default(),
949 embed_default_input_per_mtok: self.embed_default_input_per_mtok.unwrap_or(0.10),
950 guard: Arc::new(
951 self.guard
952 .unwrap_or_else(GuardEngine::empty)
953 .with_metrics(metrics.clone()),
954 ),
955 ai_task_types: Arc::new(self.ai_task_types.unwrap_or_default()),
956 metrics,
957 })
958 }
959}
960
961pub(crate) struct StreamSideEffects {
966 ledger: LedgerHandle,
967 pricing: Arc<PricingTable>,
968 metrics: Arc<GatewayMetrics>,
969 route: String,
970 tenant: String,
971 attribution: Attribution,
972 provider: String,
973 model: String,
974 lane: &'static str, request_id: String,
976 legs_attempted: u32,
977 started: Instant,
978 input_tokens: u64,
979 output_tokens: u64,
980 status: &'static str,
981 fired: bool,
982}
983
984impl StreamSideEffects {
985 #[allow(clippy::too_many_arguments)]
986 pub(crate) fn new(
987 ledger: LedgerHandle,
988 pricing: Arc<PricingTable>,
989 metrics: Arc<GatewayMetrics>,
990 route: String,
991 tenant: String,
992 attribution: Attribution,
993 provider: String,
994 model: String,
995 lane: &'static str,
996 request_id: String,
997 legs_attempted: u32,
998 started: Instant,
999 ) -> Self {
1000 Self {
1001 ledger,
1002 pricing,
1003 metrics,
1004 route,
1005 tenant,
1006 attribution,
1007 provider,
1008 model,
1009 lane,
1010 request_id,
1011 legs_attempted,
1012 started,
1013 input_tokens: 0,
1014 output_tokens: 0,
1015 status: "ok",
1016 fired: false,
1017 }
1018 }
1019
1020 pub(crate) fn observe(&mut self, item: &StreamItem) {
1022 if let StreamItem::Done {
1023 input_tokens,
1024 output_tokens,
1025 ..
1026 } = item
1027 {
1028 self.input_tokens = *input_tokens;
1029 self.output_tokens = *output_tokens;
1030 }
1031 }
1032
1033 pub(crate) fn mark_error(&mut self) {
1035 self.status = "error";
1036 }
1037}
1038
1039impl Drop for StreamSideEffects {
1040 fn drop(&mut self) {
1041 if self.fired {
1042 return;
1043 }
1044 self.fired = true;
1045
1046 let cost = self.pricing.cost_usd(
1048 &self.provider,
1049 &self.model,
1050 self.input_tokens,
1051 self.output_tokens,
1052 );
1053 self.ledger.enqueue(UsageEntry {
1054 ts: Utc::now(),
1055 tenant: self.tenant.clone(),
1056 workspace: self.attribution.workspace.clone(),
1057 user: self.attribution.user.clone(),
1058 thread: self.attribution.thread.clone(),
1059 message: self.attribution.message.clone(),
1060 route: self.route.clone(),
1061 provider: self.provider.clone(),
1062 model: self.model.clone(),
1063 lane: self.lane.to_string(),
1064 input_tokens: self.input_tokens,
1065 output_tokens: self.output_tokens,
1066 cost_usd: cost,
1067 request_id: self.request_id.clone(),
1068 status: self.status.to_string(),
1069 op: "chat".into(),
1070 user_task_type: self.attribution.user_task_type.clone(),
1071 ai_task_type: self.attribution.ai_task_type.clone(),
1072 });
1073
1074 let completion = Completion {
1077 provider: self.provider.clone(),
1078 model: self.model.clone(),
1079 content: String::new(),
1080 tool_calls: Vec::new(),
1081 finish_reason: FinishReason::Stop,
1082 input_tokens: self.input_tokens,
1083 output_tokens: self.output_tokens,
1084 };
1085 let lane = if self.lane == "native" {
1086 Lane::NativeVertex
1087 } else {
1088 Lane::Standard
1089 };
1090 GenAiSpan::from_completion(
1091 &completion,
1092 lane,
1093 &self.route,
1094 &self.tenant,
1095 self.attribution.workspace.as_deref(),
1096 self.legs_attempted,
1097 true,
1098 )
1099 .emit_metrics(&self.metrics, self.started.elapsed().as_secs_f64());
1100 }
1101}
1102
1103pub struct GuardedStream {
1107 inner: BoxStream<'static, Result<StreamItem, LegError>>,
1108 model: String,
1109 guard: StreamSideEffects,
1110 routing: RoutingReport,
1111}
1112
1113impl GuardedStream {
1114 pub(crate) fn new(
1115 inner: BoxStream<'static, Result<StreamItem, LegError>>,
1116 model: String,
1117 guard: StreamSideEffects,
1118 routing: RoutingReport,
1119 ) -> Self {
1120 Self {
1121 inner,
1122 model,
1123 guard,
1124 routing,
1125 }
1126 }
1127
1128 pub fn model(&self) -> &str {
1130 &self.model
1131 }
1132
1133 pub fn routing(&self) -> &RoutingReport {
1135 &self.routing
1136 }
1137}
1138
1139impl Stream for GuardedStream {
1140 type Item = Result<StreamItem, LegError>;
1141 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1142 let this = self.get_mut();
1144 match this.inner.as_mut().poll_next(cx) {
1145 Poll::Ready(Some(Ok(item))) => {
1146 this.guard.observe(&item);
1147 Poll::Ready(Some(Ok(item)))
1148 }
1149 Poll::Ready(Some(Err(e))) => {
1150 this.guard.mark_error();
1151 Poll::Ready(Some(Err(e)))
1152 }
1153 Poll::Ready(None) => Poll::Ready(None),
1154 Poll::Pending => Poll::Pending,
1155 }
1156 }
1157}
1158
1159pub(crate) async fn collect_committed(
1163 committed: CommittedStream,
1164) -> Result<Completion, GatewayError> {
1165 use futures::TryStreamExt;
1166 let CommittedStream {
1167 provider,
1168 model,
1169 stream,
1170 } = committed;
1171 stream
1172 .map_err(|e: LegError| GatewayError::Upstream {
1173 status: 502,
1174 body: e.to_string(),
1175 })
1176 .try_fold(Accumulator::default(), |mut acc, item| async move {
1177 acc.push(item);
1178 Ok(acc)
1179 })
1180 .await
1181 .map(|acc| Completion {
1182 provider,
1183 model,
1184 content: acc.content,
1185 tool_calls: acc.tool_calls,
1186 finish_reason: acc.finish_reason,
1187 input_tokens: acc.input_tokens,
1188 output_tokens: acc.output_tokens,
1189 })
1190}
1191
1192#[cfg(test)]
1193mod tests {
1194 use super::*;
1195 use crate::ledger::{InMemoryLedger, LedgerHandle, LedgerStore};
1196
1197 fn test_gateway() -> Gateway {
1198 let routes = RouteTable::from_toml_str(
1199 r#"[routes."fast"]
1200 legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1201 )
1202 .unwrap();
1203 let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1204 let ledger = LedgerHandle::spawn(
1205 Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1206 16,
1207 );
1208 Gateway::builder()
1209 .routes(routes)
1210 .catalog(catalog)
1211 .pricing(PricingTable::default())
1212 .ledger(ledger)
1213 .default_tenant("acme")
1214 .build()
1215 .unwrap()
1216 }
1217
1218 #[tokio::test]
1219 async fn builder_builds_and_lists_aliases() {
1220 let gw = test_gateway();
1221 assert_eq!(gw.model_aliases(), vec!["fast".to_string()]);
1222 assert_eq!(gw.default_tenant, "acme");
1223 }
1224
1225 #[test]
1226 fn builder_requires_components() {
1227 assert!(Gateway::builder().build().is_err());
1228 }
1229
1230 #[tokio::test]
1231 async fn guard_records_usage_on_drop() {
1232 let store = Arc::new(InMemoryLedger::default());
1233 let ledger = LedgerHandle::spawn(store.clone(), 16);
1234 let pricing = Arc::new(PricingTable::default());
1235 {
1236 let mut guard = StreamSideEffects::new(
1237 ledger.clone(),
1238 pricing,
1239 GatewayMetrics::noop(),
1240 "route".into(),
1241 "tenant".into(),
1242 Attribution::default(),
1243 "p".into(),
1244 "m".into(),
1245 "standard",
1246 "rid".into(),
1247 1,
1248 Instant::now(),
1249 );
1250 guard.observe(&crate::routing::stream::StreamItem::Done {
1251 input_tokens: 3,
1252 output_tokens: 2,
1253 finish_reason: crate::routing::stream::FinishReason::Stop,
1254 });
1255 } tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1257 let rows = store.entries();
1258 assert_eq!(rows.len(), 1);
1259 assert_eq!(rows[0].input_tokens, 3);
1260 assert_eq!(rows[0].output_tokens, 2);
1261 assert_eq!(rows[0].status, "ok");
1262 }
1263
1264 #[tokio::test]
1265 async fn guarded_stream_yields_items_and_records_on_drop() {
1266 use crate::routing::stream::{FinishReason, StreamItem};
1267 use futures::StreamExt;
1268 let store = Arc::new(InMemoryLedger::default());
1269 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1270 let inner = futures::stream::iter(vec![
1271 Ok(StreamItem::Delta("hi".into())),
1272 Ok(StreamItem::Done {
1273 input_tokens: 4,
1274 output_tokens: 2,
1275 finish_reason: FinishReason::Stop,
1276 }),
1277 ])
1278 .boxed();
1279 let guard = StreamSideEffects::new(
1280 ledger,
1281 Arc::new(PricingTable::default()),
1282 GatewayMetrics::noop(),
1283 "route".into(),
1284 "acme".into(),
1285 Attribution::default(),
1286 "p".into(),
1287 "m".into(),
1288 "standard",
1289 "rid".into(),
1290 1,
1291 std::time::Instant::now(),
1292 );
1293 {
1294 let mut gs = GuardedStream::new(inner, "m".into(), guard, RoutingReport::default());
1295 let mut n = 0;
1296 while let Some(item) = gs.next().await {
1297 item.unwrap();
1298 n += 1;
1299 }
1300 assert_eq!(n, 2);
1301 assert_eq!(gs.model(), "m");
1302 } tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1304 let rows = store.entries();
1305 assert_eq!(rows.len(), 1);
1306 assert_eq!(rows[0].input_tokens, 4);
1307 assert_eq!(rows[0].tenant, "acme");
1308 }
1309
1310 fn resolver_gateway(ai_task_types: &str) -> Gateway {
1312 let routes = RouteTable::from_toml_str(
1313 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1314 )
1315 .unwrap();
1316 Gateway::builder()
1317 .routes(routes)
1318 .catalog(Catalog::for_test(vec![(
1319 "qwen",
1320 "http://127.0.0.1:1/v1".into(),
1321 )]))
1322 .pricing(PricingTable::default())
1323 .ledger(LedgerHandle::spawn(
1324 Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1325 4,
1326 ))
1327 .ai_task_types(
1328 crate::ai_task_type::AiTaskTypeTable::from_toml_str(ai_task_types).unwrap(),
1329 )
1330 .build()
1331 .unwrap()
1332 }
1333
1334 #[tokio::test]
1335 async fn ai_task_type_prefers_the_caller_header_over_the_alias_mapping() {
1336 let gw = resolver_gateway("conversation = [\"fast\"]");
1337 let ctx = RequestCtx {
1338 ai_task_type: Some("caller-supplied".into()),
1339 ..Default::default()
1340 };
1341 assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "caller-supplied");
1342 }
1343
1344 #[tokio::test]
1345 async fn ai_task_type_is_inferred_from_the_route_alias_when_no_header() {
1346 let gw = resolver_gateway("conversation = [\"fast\", \"planning\"]");
1347 let ctx = RequestCtx::default();
1348 assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
1349 assert_eq!(gw.ai_task_type_of(&ctx, "planning"), "conversation");
1350 }
1351
1352 #[tokio::test]
1353 async fn ai_task_type_defaults_to_simple_for_an_unmapped_alias() {
1354 let gw = resolver_gateway("conversation = [\"fast\"]");
1355 assert_eq!(
1356 gw.ai_task_type_of(&RequestCtx::default(), "graph-llm"),
1357 "simple"
1358 );
1359 }
1360
1361 #[tokio::test]
1362 async fn an_empty_ai_task_type_header_falls_back_to_inference() {
1363 let gw = resolver_gateway("conversation = [\"fast\"]");
1364 let ctx = RequestCtx {
1365 ai_task_type: Some(String::new()),
1366 ..Default::default()
1367 };
1368 assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
1369 }
1370
1371 #[tokio::test]
1372 async fn chat_returns_completion_and_records_ledger() {
1373 use wiremock::matchers::{method, path};
1374 use wiremock::{Mock, MockServer, ResponseTemplate};
1375 let mock = MockServer::start().await;
1376 Mock::given(method("POST")).and(path("/v1/chat/completions"))
1377 .respond_with(ResponseTemplate::new(200)
1378 .insert_header("content-type", "text/event-stream")
1379 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\n\
1380 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2}}\n\n\
1381 data: [DONE]\n\n"))
1382 .mount(&mock).await;
1383 let routes = RouteTable::from_toml_str(
1384 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1385 )
1386 .unwrap();
1387 let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
1388 let store = Arc::new(InMemoryLedger::default());
1389 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1390 let gw = Gateway::builder()
1391 .routes(routes)
1392 .catalog(catalog)
1393 .pricing(PricingTable::default())
1394 .ledger(ledger)
1395 .default_tenant("def")
1396 .ai_task_types(
1397 crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
1398 .unwrap(),
1399 )
1400 .build()
1401 .unwrap();
1402 let req = serde_json::from_value(serde_json::json!(
1403 {"model":"fast","messages":[{"role":"user","content":"hi"}]}))
1404 .unwrap();
1405 let ctx = RequestCtx {
1406 tenant: Some("acme".into()),
1407 workspace: Some("ws-9".into()),
1408 user: Some("user-42".into()),
1409 thread: Some("thread-9".into()),
1410 message: Some("msg-7".into()),
1411 user_task_type: Some("summarisation".into()),
1412 ai_task_type: None,
1414 request_id: Some("corr-123".into()),
1415 };
1416 let c = match gw.chat(req, &ctx).await.unwrap() {
1417 ChatOutcome::Plain(c) => c,
1418 ChatOutcome::Hybrid(_) => panic!("expected plain completion"),
1419 };
1420 assert_eq!(c.content, "hi");
1421 assert_eq!(c.input_tokens, 3);
1422 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1423 let rows = store.entries();
1424 assert_eq!(rows.len(), 1);
1425 assert_eq!(rows[0].tenant, "acme");
1426 assert_eq!(rows[0].user.as_deref(), Some("user-42"));
1427 assert_eq!(rows[0].thread.as_deref(), Some("thread-9"));
1428 assert_eq!(rows[0].message.as_deref(), Some("msg-7"));
1429 assert_eq!(rows[0].user_task_type.as_deref(), Some("summarisation"));
1430 assert_eq!(rows[0].ai_task_type, "conversation");
1432 assert_eq!(rows[0].request_id, "corr-123");
1434 }
1435
1436 #[tokio::test]
1437 async fn chat_stream_yields_items_and_records() {
1438 use crate::routing::stream::StreamItem;
1439 use futures::StreamExt;
1440 use wiremock::matchers::{method, path};
1441 use wiremock::{Mock, MockServer, ResponseTemplate};
1442 let mock = MockServer::start().await;
1443 Mock::given(method("POST")).and(path("/v1/chat/completions"))
1444 .respond_with(ResponseTemplate::new(200)
1445 .insert_header("content-type", "text/event-stream")
1446 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
1447 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
1448 data: [DONE]\n\n"))
1449 .mount(&mock).await;
1450 let routes = RouteTable::from_toml_str(
1451 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1452 )
1453 .unwrap();
1454 let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
1455 let store = Arc::new(InMemoryLedger::default());
1456 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1457 let gw = Gateway::builder()
1458 .routes(routes)
1459 .catalog(catalog)
1460 .pricing(PricingTable::default())
1461 .ledger(ledger)
1462 .default_tenant("def")
1463 .ai_task_types(
1464 crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
1465 .unwrap(),
1466 )
1467 .build()
1468 .unwrap();
1469 let req = serde_json::from_value(serde_json::json!(
1470 {"model":"fast","stream":true,"messages":[{"role":"user","content":"hi"}]}))
1471 .unwrap();
1472 let ctx = RequestCtx {
1473 user_task_type: Some("code-review".into()),
1474 ..Default::default()
1475 };
1476 let mut stream = gw.chat_stream(req, &ctx).await.unwrap();
1477 let mut got = false;
1478 while let Some(i) = stream.next().await {
1479 if matches!(i.unwrap(), StreamItem::Delta(ref t) if t == "go") {
1480 got = true;
1481 }
1482 }
1483 drop(stream);
1484 assert!(got);
1485 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1486 let rows = store.entries();
1487 assert_eq!(rows.len(), 1);
1488 assert_eq!(rows[0].user_task_type.as_deref(), Some("code-review"));
1489 assert_eq!(rows[0].ai_task_type, "conversation");
1490 }
1491
1492 use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
1495 use async_trait::async_trait;
1496
1497 struct FlakyEmbedder;
1499 #[async_trait]
1500 impl EmbeddingProvider for FlakyEmbedder {
1501 async fn embed(
1502 &self,
1503 _model: &str,
1504 _inputs: &[String],
1505 _dims: u32,
1506 ) -> Result<EmbedOut, GatewayError> {
1507 Err(GatewayError::Upstream {
1508 status: 500,
1509 body: "x".into(),
1510 })
1511 }
1512 }
1513
1514 struct GoodEmbedder;
1516 #[async_trait]
1517 impl EmbeddingProvider for GoodEmbedder {
1518 async fn embed(
1519 &self,
1520 _model: &str,
1521 inputs: &[String],
1522 dims: u32,
1523 ) -> Result<EmbedOut, GatewayError> {
1524 Ok(EmbedOut {
1525 vectors: inputs.iter().map(|_| vec![0.0f32; dims as usize]).collect(),
1526 input_tokens: 6,
1527 })
1528 }
1529 }
1530
1531 fn embed_gateway() -> (Gateway, Arc<InMemoryLedger>) {
1532 let embed_routes = crate::routing::embeddings::EmbeddingRouteTable::from_toml_str(
1533 r#"
1534 [embeddings."default-embed"]
1535 dimensions = 4
1536 legs = [
1537 { provider = "flaky", model = "flaky-embed" },
1538 { provider = "good", model = "good-embed" },
1539 ]
1540 "#,
1541 )
1542 .unwrap();
1543 let routes = RouteTable::from_toml_str(
1544 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1545 )
1546 .unwrap();
1547 let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1548 let store = Arc::new(InMemoryLedger::default());
1549 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1550 let gw = Gateway::builder()
1551 .routes(routes)
1552 .catalog(catalog)
1553 .pricing(PricingTable::default())
1554 .ledger(ledger)
1555 .default_tenant("acme")
1556 .embed_routes(embed_routes)
1557 .embedder(
1558 "flaky",
1559 Arc::new(FlakyEmbedder) as Arc<dyn EmbeddingProvider>,
1560 )
1561 .embedder("good", Arc::new(GoodEmbedder) as Arc<dyn EmbeddingProvider>)
1562 .build()
1563 .unwrap();
1564 (gw, store)
1565 }
1566
1567 #[tokio::test]
1568 async fn embed_falls_through_to_good_leg_and_records_usage() {
1569 let (gw, store) = embed_gateway();
1570 let req = EmbeddingRequest {
1571 input: EmbeddingInput::Many(vec!["a".into(), "b".into()]),
1572 model: "default-embed".into(),
1573 dimensions: None,
1574 };
1575 let ctx = RequestCtx {
1576 user_task_type: Some("retrieval".into()),
1577 ..Default::default()
1578 };
1579 let resp = gw.embed(req, ctx).await.unwrap();
1580 assert_eq!(resp.data.len(), 2);
1581 assert!(resp.data.iter().all(|d| d.embedding.len() == 4));
1582 assert_eq!(resp.data[0].index, 0);
1583 assert_eq!(resp.data[1].index, 1);
1584 assert_eq!(resp.usage.prompt_tokens, 6);
1585
1586 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1587 let rows = store.entries();
1588 assert_eq!(rows.len(), 1);
1589 assert_eq!(rows[0].op, "embedding");
1590 assert_eq!(rows[0].lane, "embedding");
1591 assert_eq!(rows[0].output_tokens, 0);
1592 assert_eq!(rows[0].provider, "good");
1593 assert!(rows[0].cost_usd > 0.0);
1594 assert_eq!(rows[0].user_task_type.as_deref(), Some("retrieval"));
1595 assert_eq!(rows[0].ai_task_type, "simple");
1596 }
1597
1598 #[tokio::test]
1599 async fn embed_dimension_mismatch_is_bad_request() {
1600 let (gw, _store) = embed_gateway();
1601 let req = EmbeddingRequest {
1602 input: EmbeddingInput::Many(vec!["a".into()]),
1603 model: "default-embed".into(),
1604 dimensions: Some(8),
1605 };
1606 let err = gw.embed(req, RequestCtx::default()).await.unwrap_err();
1607 assert!(matches!(err, GatewayError::BadRequest(_)));
1608 }
1609
1610 #[tokio::test]
1611 async fn chat_blocks_when_route_policy_refuses() {
1612 use crate::guard::{GuardEngine, GuardrailsConfig};
1613 let routes = RouteTable::from_toml_str(
1614 r#"[routes."fast"]
1615 policy = "strict"
1616 legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1617 )
1618 .unwrap();
1619 let guard = GuardEngine::from_config(
1620 &GuardrailsConfig::from_toml_str(
1621 r#"[guardrails.strict]
1622 scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#,
1623 )
1624 .unwrap(),
1625 )
1626 .unwrap();
1627 let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1628 let ledger = LedgerHandle::spawn(
1629 Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1630 16,
1631 );
1632 let gw = Gateway::builder()
1633 .routes(routes)
1634 .catalog(catalog)
1635 .pricing(PricingTable::default())
1636 .ledger(ledger)
1637 .guard(guard)
1638 .build()
1639 .unwrap();
1640 let req = serde_json::from_value(serde_json::json!({
1641 "model": "fast",
1642 "messages": [{ "role": "user", "content": "this is forbidden" }]
1643 }))
1644 .unwrap();
1645 let err = gw.chat(req, &RequestCtx::default()).await.unwrap_err();
1646 assert!(matches!(err, GatewayError::ContentBlocked { .. }));
1647 }
1648
1649 fn vertex_req(vertex: serde_json::Value) -> ChatRequest {
1650 serde_json::from_value(serde_json::json!({
1651 "model": "auto",
1652 "messages": [{"role": "user", "content": "hi"}],
1653 "vertex": vertex
1654 }))
1655 .unwrap()
1656 }
1657
1658 fn vertex_leg(effort: Option<crate::routing::effort::Effort>) -> ChainLeg {
1659 ChainLeg {
1660 provider: "vertex".into(),
1661 model: "gemini-2.5-pro".into(),
1662 effort,
1663 ..Default::default()
1664 }
1665 }
1666
1667 #[test]
1668 fn leg_effort_becomes_a_thinking_budget_and_keeps_the_vertex_block() {
1669 let req = vertex_req(serde_json::json!({"response_schema": {"type": "object"}}));
1670 let sent = with_leg_thinking(
1671 &req,
1672 &vertex_leg(Some(crate::routing::effort::Effort::High)),
1673 );
1674 let v = sent.vertex.as_ref().unwrap();
1675 assert_eq!(
1676 v.thinking_config,
1677 Some(serde_json::json!({"thinkingBudget": 8192}))
1678 );
1679 assert_eq!(
1680 v.response_schema,
1681 Some(serde_json::json!({"type": "object"}))
1682 );
1683 }
1684
1685 #[test]
1686 fn client_thinking_config_wins_over_leg_effort() {
1687 let req = vertex_req(serde_json::json!({"thinking_config": {"thinkingLevel": "low"}}));
1688 let sent = with_leg_thinking(&req, &vertex_leg(Some(crate::routing::effort::Effort::Max)));
1689 assert!(matches!(sent, std::borrow::Cow::Borrowed(_)));
1690 }
1691
1692 #[test]
1693 fn effort_none_or_absent_leaves_the_request_untouched() {
1694 let req = vertex_req(serde_json::json!({"response_schema": {"type": "object"}}));
1695 assert!(matches!(
1696 with_leg_thinking(
1697 &req,
1698 &vertex_leg(Some(crate::routing::effort::Effort::None))
1699 ),
1700 std::borrow::Cow::Borrowed(_)
1701 ));
1702 assert!(matches!(
1703 with_leg_thinking(&req, &vertex_leg(None)),
1704 std::borrow::Cow::Borrowed(_)
1705 ));
1706 }
1707
1708 #[test]
1709 fn null_thinking_config_or_missing_vertex_block_gets_the_leg_budget() {
1710 let budget = Some(serde_json::json!({"thinkingBudget": 1024}));
1711 let leg = vertex_leg(Some(crate::routing::effort::Effort::Low));
1712 let null_config = vertex_req(serde_json::json!({"thinking_config": null}));
1713 let no_vertex: ChatRequest = serde_json::from_value(serde_json::json!({
1714 "model": "auto",
1715 "messages": [{"role": "user", "content": "hi"}]
1716 }))
1717 .unwrap();
1718 [null_config, no_vertex].iter().for_each(|req| {
1719 assert_eq!(
1720 with_leg_thinking(req, &leg)
1721 .vertex
1722 .as_ref()
1723 .and_then(|v| v.thinking_config.clone()),
1724 budget
1725 )
1726 });
1727 }
1728}