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, LegFailure};
14use crate::guard::GuardEngine;
15use crate::jev_native::JevNativeProvider;
16use crate::ledger::{LedgerHandle, UsageEntry};
17use crate::observability::GenAiSpan;
18use crate::pricing::PricingTable;
19use crate::providers::Catalog;
20use crate::routing::classify::{classify, vertex_triggers, Lane};
21use crate::routing::executor::{
22 execute_buffered_with_timeouts, CommittedStream, Completion, LegError, StreamTimeouts,
23};
24use crate::routing::request::ChatRequest;
25use crate::routing::stream::{Accumulator, FinishReason, StreamItem};
26use crate::routing::table::{ChainLeg, RouteTable};
27use crate::vertex_native::VertexNativeProvider;
28
29#[derive(Clone)]
31pub struct Gateway {
32 pub(crate) routes: Arc<RouteTable>,
33 pub(crate) catalog: Arc<Catalog>,
34 pub(crate) pricing: Arc<PricingTable>,
35 pub(crate) ledger: LedgerHandle,
36 pub(crate) vertex_native: Option<Arc<VertexNativeProvider>>,
37 pub(crate) jev_native: Option<Arc<JevNativeProvider>>,
38 pub(crate) timeouts: StreamTimeouts,
39 pub(crate) default_tenant: String,
40 pub(crate) embed_routes: Arc<crate::routing::embeddings::EmbeddingRouteTable>,
41 pub(crate) embedders:
42 std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>,
43 pub(crate) embed_default_input_per_mtok: f64,
44 pub(crate) guard: Arc<GuardEngine>,
45 pub(crate) ai_task_types: Arc<crate::ai_task_type::AiTaskTypeTable>,
46}
47
48#[derive(Debug, Clone, Default)]
50pub struct RequestCtx {
51 pub tenant: Option<String>,
52 pub workspace: Option<String>,
53 pub user: Option<String>,
55 pub thread: Option<String>,
57 pub message: Option<String>,
59 pub user_task_type: Option<String>,
63 pub ai_task_type: Option<String>,
67 pub request_id: Option<String>,
72}
73
74impl RequestCtx {
75 pub fn resolved_request_id(&self) -> Option<String> {
78 self.request_id
79 .clone()
80 .or_else(|| self.message.clone())
81 .filter(|s| !s.is_empty())
82 }
83}
84
85pub(crate) enum JevAttempt {
87 Decided {
90 completion: Completion,
91 answers: serde_json::Value,
92 },
93 Exhausted(Vec<LegFailure>),
96}
97
98#[derive(Debug)]
101pub enum ChatOutcome {
102 Plain(Completion),
103 Hybrid(HybridOutcome),
104}
105
106#[derive(Debug)]
109pub struct HybridOutcome {
110 pub model: String,
112 pub answers: serde_json::Value,
114 pub survivors: Vec<String>,
116 pub degraded: bool,
118 pub extraction_ran: bool,
120 pub extractions: serde_json::Map<String, serde_json::Value>,
122 pub input_tokens: u64,
124 pub output_tokens: u64,
125}
126
127#[derive(Debug, Clone)]
133pub struct Attribution {
134 pub workspace: Option<String>,
135 pub user: Option<String>,
136 pub thread: Option<String>,
137 pub message: Option<String>,
138 pub user_task_type: Option<String>,
139 pub ai_task_type: String,
141}
142
143impl Default for Attribution {
144 fn default() -> Self {
145 Self {
146 workspace: None,
147 user: None,
148 thread: None,
149 message: None,
150 user_task_type: None,
151 ai_task_type: crate::ai_task_type::DEFAULT_AI_TASK_TYPE.to_string(),
152 }
153 }
154}
155
156#[derive(Default)]
157pub struct GatewayBuilder {
158 routes: Option<RouteTable>,
159 catalog: Option<Catalog>,
160 pricing: Option<PricingTable>,
161 ledger: Option<LedgerHandle>,
162 vertex_native: Option<VertexNativeProvider>,
163 jev_native: Option<JevNativeProvider>,
164 timeouts: Option<StreamTimeouts>,
165 default_tenant: Option<String>,
166 embed_routes: Option<crate::routing::embeddings::EmbeddingRouteTable>,
167 embedders:
168 Option<std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>>,
169 embed_default_input_per_mtok: Option<f64>,
170 guard: Option<GuardEngine>,
171 ai_task_types: Option<crate::ai_task_type::AiTaskTypeTable>,
172}
173
174impl Gateway {
175 pub fn builder() -> GatewayBuilder {
176 GatewayBuilder::default()
177 }
178
179 pub fn model_aliases(&self) -> Vec<String> {
181 self.routes.aliases()
182 }
183
184 fn resolve_legs(&self, req: &ChatRequest) -> Result<Vec<ChainLeg>, GatewayError> {
185 self.routes
186 .legs(&req.model)
187 .ok_or_else(|| GatewayError::UnknownModel(req.model.clone()))
188 .map(<[ChainLeg]>::to_vec)
189 }
190
191 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 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(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(&leg.model, req, leg.region.as_deref())
310 .await
311 {
312 Ok(stream) => {
313 return Ok(CommittedStream::single(
314 "vertex".into(),
315 leg.model.clone(),
316 stream,
317 ));
318 }
319 Err(e) if native_start_retryable(&e) => {
320 failures.push(LegFailure {
321 provider: leg.provider.clone(),
322 model: leg.model.clone(),
323 message: e.to_string(),
324 });
325 last_retryable = Some(e);
326 }
327 Err(e) => return Err(e),
328 }
329 }
330 Err(
331 last_retryable.unwrap_or_else(|| GatewayError::AllLegsFailed {
332 route: req.model.clone(),
333 failures,
334 }),
335 )
336 }
337
338 pub(crate) async fn jev_attempt(
345 &self,
346 req: &ChatRequest,
347 legs: &[ChainLeg],
348 ) -> Result<JevAttempt, GatewayError> {
349 let jev_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "typesafe").collect();
350 if jev_legs.is_empty() {
351 return Ok(JevAttempt::Exhausted(Vec::new()));
352 }
353 let ext = req
354 .jev
355 .as_ref()
356 .filter(|j| !j.questions.is_empty())
357 .ok_or_else(|| {
358 GatewayError::BadRequest(format!(
359 "route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
360 req.model
361 ))
362 })?;
363 let provider = self.jev_native.as_ref().ok_or_else(|| {
364 GatewayError::BadRequest(
365 "typesafe legs require TYPESAFE_API_KEY to be configured".into(),
366 )
367 })?;
368
369 let state = ext
370 .state
371 .clone()
372 .or_else(|| serde_json::to_string(&req.messages).ok())
373 .unwrap_or_default();
374 let mut body = serde_json::json!({
375 "state": state,
376 "questions": ext.questions.clone(),
377 });
378
379 let mut failures: Vec<LegFailure> = Vec::new();
380 for leg in jev_legs {
381 let model = if leg.model.is_empty() {
382 crate::jev_native::DEFAULT_MODEL.to_string()
383 } else {
384 leg.model.clone()
385 };
386 body["model"] = serde_json::Value::String(model.clone());
387
388 let resp = match provider.evaluate(body.clone()).await {
389 Ok(r) => r,
390 Err(e) if native_start_retryable(&e) => {
391 failures.push(LegFailure {
392 provider: leg.provider.clone(),
393 model,
394 message: e.to_string(),
395 });
396 continue;
397 }
398 Err(e) => return Err(e),
399 };
400
401 let status = resp.status();
402 let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
403 status: 502,
404 body: e.to_string(),
405 })?;
406 if status.is_success() {
407 let value: serde_json::Value =
408 serde_json::from_slice(&bytes).map_err(|e| GatewayError::Upstream {
409 status: 502,
410 body: format!("jev response is not JSON: {e}"),
411 })?;
412 return Ok(JevAttempt::Decided {
413 answers: value["answers"].clone(),
414 completion: Completion {
415 provider: "typesafe".into(),
416 model: value["model"].as_str().unwrap_or(&model).to_string(),
417 content: serde_json::to_string(&value["answers"])
418 .unwrap_or_else(|_| "{}".into()),
419 tool_calls: Vec::new(),
420 finish_reason: FinishReason::Stop,
421 input_tokens: value["usage"]["input_tokens"].as_u64().unwrap_or(0),
422 output_tokens: value["usage"]["output_tokens"].as_u64().unwrap_or(0),
423 },
424 });
425 }
426
427 let message = String::from_utf8_lossy(&bytes).into_owned();
428 if status.is_server_error()
429 || status == reqwest::StatusCode::TOO_MANY_REQUESTS
430 || status == reqwest::StatusCode::REQUEST_TIMEOUT
431 {
432 failures.push(LegFailure {
433 provider: leg.provider.clone(),
434 model,
435 message: format!("{status}: {message}"),
436 });
437 continue;
438 }
439 return Err(GatewayError::BadRequest(format!(
440 "typesafe {}: {message}",
441 status.as_u16()
442 )));
443 }
444 Ok(JevAttempt::Exhausted(failures))
445 }
446
447 fn jev_committed(c: Completion) -> CommittedStream {
450 let items: Vec<Result<StreamItem, LegError>> = vec![
451 Ok(StreamItem::Delta(c.content.clone())),
452 Ok(StreamItem::Done {
453 input_tokens: c.input_tokens,
454 output_tokens: c.output_tokens,
455 finish_reason: c.finish_reason,
456 }),
457 ];
458 CommittedStream::single("typesafe".into(), c.model, futures::stream::iter(items))
459 }
460
461 pub async fn chat_stream(
464 &self,
465 req: ChatRequest,
466 ctx: &RequestCtx,
467 ) -> Result<GuardedStream, GatewayError> {
468 use crate::routing::executor::execute_streaming_with_timeouts;
469 let started = Instant::now();
470 let legs = self.resolve_legs(&req)?;
471 self.guard_input(&req)?;
472 let request_id = ctx
473 .resolved_request_id()
474 .unwrap_or_else(|| Uuid::new_v4().to_string());
475 let vertex_leg_count = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
476 require_jev_block(&req, &legs)?;
477 if let Err(message) = crate::routing::jev_extract::validate_extract(
478 &req,
479 legs.iter().any(|l| l.provider != "typesafe"),
480 ) {
481 return Err(GatewayError::BadRequest(message));
482 }
483 let (committed, lane_str, legs_attempted) = match classify(&req) {
484 Lane::Standard => (
485 execute_streaming_with_timeouts(
486 &self.catalog,
487 &req.model,
488 &legs,
489 &req,
490 self.timeouts,
491 )
492 .await?,
493 "standard",
494 legs.len() as u32,
495 ),
496 Lane::NativeVertex => (
497 self.native_committed(&req, &legs).await?,
498 "native",
499 vertex_leg_count.max(1),
500 ),
501 Lane::Jev => match self.jev_attempt(&req, &legs).await? {
502 JevAttempt::Decided { completion, .. } => (
503 Self::jev_committed(completion),
504 "jev",
505 legs.iter()
506 .filter(|l| l.provider == "typesafe")
507 .count()
508 .max(1) as u32,
509 ),
510 JevAttempt::Exhausted(failures) => {
511 let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
512 if rest.is_empty() {
513 return Err(GatewayError::AllLegsFailed {
514 route: req.model.clone(),
515 failures,
516 });
517 }
518 if vertex_triggers(&req) {
519 let vertex_n =
520 rest.iter().filter(|l| l.provider == "vertex").count() as u32;
521 (
522 self.native_committed(&req, &rest).await?,
523 "native",
524 vertex_n.max(1),
525 )
526 } else {
527 (
528 execute_streaming_with_timeouts(
529 &self.catalog,
530 &req.model,
531 &rest,
532 &req,
533 self.timeouts,
534 )
535 .await?,
536 "standard",
537 rest.len() as u32,
538 )
539 }
540 }
541 },
542 };
543 let model = committed.model.clone();
544 let guard = StreamSideEffects::new(
545 self.ledger.clone(),
546 self.pricing.clone(),
547 req.model.clone(),
548 self.tenant_of(ctx).to_string(),
549 self.attribution_of(ctx, &req.model),
550 committed.provider.clone(),
551 committed.model.clone(),
552 lane_str,
553 request_id,
554 legs_attempted,
555 started,
556 );
557 Ok(GuardedStream::new(committed.stream, model, guard))
558 }
559
560 pub async fn chat(
563 &self,
564 req: ChatRequest,
565 ctx: &RequestCtx,
566 ) -> Result<ChatOutcome, GatewayError> {
567 let started = Instant::now();
568 let legs = self.resolve_legs(&req)?;
569 self.guard_input(&req)?;
570 let request_id = ctx
571 .resolved_request_id()
572 .unwrap_or_else(|| Uuid::new_v4().to_string());
573 require_jev_block(&req, &legs)?;
574 if let Err(message) = crate::routing::jev_extract::validate_extract(
575 &req,
576 legs.iter().any(|l| l.provider != "typesafe"),
577 ) {
578 return Err(GatewayError::BadRequest(message));
579 }
580 if req.jev.as_ref().and_then(|j| j.extract.as_ref()).is_some() {
581 let outcome = self
582 .chat_hybrid(req, ctx, &legs, started, &request_id)
583 .await?;
584 return Ok(ChatOutcome::Hybrid(outcome));
585 }
586 let (completion, lane_str, legs_n) = match classify(&req) {
587 Lane::Standard => execute_buffered_with_timeouts(
588 &self.catalog,
589 &req.model,
590 &legs,
591 &req,
592 self.timeouts,
593 )
594 .await
595 .map(|c| (c, "standard", legs.len() as u32))?,
596 Lane::NativeVertex => {
597 let vertex_n = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
598 let committed = self.native_committed(&req, &legs).await?;
599 (
600 collect_committed(committed).await?,
601 "native",
602 vertex_n.max(1),
603 )
604 }
605 Lane::Jev => match self.jev_attempt(&req, &legs).await? {
606 JevAttempt::Decided { completion, .. } => (
607 completion,
608 "jev",
609 legs.iter()
610 .filter(|l| l.provider == "typesafe")
611 .count()
612 .max(1) as u32,
613 ),
614 JevAttempt::Exhausted(failures) => {
615 let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
616 if rest.is_empty() {
617 return Err(GatewayError::AllLegsFailed {
618 route: req.model.clone(),
619 failures,
620 });
621 }
622 if vertex_triggers(&req) {
623 let committed = self.native_committed(&req, &rest).await?;
624 let vertex_n =
625 rest.iter().filter(|l| l.provider == "vertex").count() as u32;
626 (
627 collect_committed(committed).await?,
628 "native",
629 vertex_n.max(1),
630 )
631 } else {
632 execute_buffered_with_timeouts(
633 &self.catalog,
634 &req.model,
635 &rest,
636 &req,
637 self.timeouts,
638 )
639 .await
640 .map(|c| (c, "standard", rest.len() as u32))?
641 }
642 }
643 },
644 };
645 self.record(
646 ctx,
647 &req.model,
648 lane_str,
649 &request_id,
650 &completion,
651 legs_n,
652 started,
653 );
654 Ok(ChatOutcome::Plain(completion))
655 }
656
657 pub async fn embed(
663 &self,
664 req: crate::embeddings::EmbeddingRequest,
665 ctx: RequestCtx,
666 ) -> Result<crate::embeddings::EmbeddingResponse, GatewayError> {
667 let alias = req.model.clone();
668 let dims = self
669 .embed_routes
670 .dimensions(&alias)
671 .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?;
672 if let Some(d) = req.dimensions {
673 if d != dims {
674 return Err(GatewayError::BadRequest(format!(
675 "dimensions {d} does not match embedding alias '{alias}' dimension {dims}"
676 )));
677 }
678 }
679 let inputs = req.input.into_vec();
680 let legs = self
681 .embed_routes
682 .legs(&alias)
683 .ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?
684 .to_vec();
685
686 let mut last_err: Option<GatewayError> = None;
687 for leg in &legs {
688 let Some(embedder) = self.embedders.get(&leg.provider) else {
689 continue;
690 };
691 let limit = match leg.provider.as_str() {
692 "vertex" => crate::embeddings::vertex::VERTEX_EMBED_BATCH,
693 _ => crate::embeddings::openai::OPENAI_EMBED_BATCH,
694 };
695 let started = std::time::Instant::now();
696 match self
697 .embed_all_batches(embedder.as_ref(), &leg.model, &inputs, dims, limit)
698 .await
699 {
700 Ok(out) => {
701 metrics::counter!(
702 "synapse_embeddings_total",
703 "route" => alias.clone(),
704 "model" => leg.model.clone(),
705 "provider" => leg.provider.clone(),
706 )
707 .increment(1);
708 metrics::histogram!(
709 "synapse_embedding_duration_seconds",
710 "route" => alias.clone(),
711 "model" => leg.model.clone(),
712 "provider" => leg.provider.clone(),
713 )
714 .record(started.elapsed().as_secs_f64());
715 self.record_embed_usage(&ctx, &alias, leg, out.input_tokens);
716 return Ok(crate::embeddings::build_response(alias, out));
717 }
718 Err(e) => last_err = Some(e),
719 }
720 }
721 Err(last_err.unwrap_or(GatewayError::AllLegsFailed {
722 route: alias,
723 failures: Vec::new(),
724 }))
725 }
726
727 async fn embed_all_batches(
730 &self,
731 embedder: &dyn crate::embeddings::EmbeddingProvider,
732 model: &str,
733 inputs: &[String],
734 dims: u32,
735 limit: usize,
736 ) -> Result<crate::embeddings::EmbedOut, GatewayError> {
737 let mut vectors = Vec::with_capacity(inputs.len());
738 let mut input_tokens = 0u64;
739 for batch in crate::embeddings::split_batches(inputs, limit) {
740 let out = embedder.embed(model, batch, dims).await?;
741 input_tokens += out.input_tokens;
742 vectors.extend(out.vectors);
743 }
744 Ok(crate::embeddings::EmbedOut {
745 vectors,
746 input_tokens,
747 })
748 }
749
750 fn record_embed_usage(&self, ctx: &RequestCtx, alias: &str, leg: &ChainLeg, input_tokens: u64) {
752 let cost = self.pricing.embedding_cost_usd(
753 &leg.provider,
754 &leg.model,
755 input_tokens,
756 self.embed_default_input_per_mtok,
757 );
758 let attr = self.attribution_of(ctx, alias);
759 self.ledger.enqueue(UsageEntry {
760 ts: Utc::now(),
761 tenant: self.tenant_of(ctx).to_string(),
762 workspace: attr.workspace,
763 user: attr.user,
764 thread: attr.thread,
765 message: attr.message,
766 route: alias.to_string(),
767 provider: leg.provider.clone(),
768 model: leg.model.clone(),
769 lane: "embedding".into(),
770 input_tokens,
771 output_tokens: 0,
772 cost_usd: cost,
773 request_id: ctx
774 .resolved_request_id()
775 .unwrap_or_else(|| Uuid::new_v4().to_string()),
776 status: "ok".into(),
777 op: "embedding".into(),
778 user_task_type: attr.user_task_type,
779 ai_task_type: attr.ai_task_type,
780 });
781 }
782}
783
784fn native_start_retryable(e: &GatewayError) -> bool {
787 match e {
788 GatewayError::Upstream { status, .. } => {
789 *status >= 500
790 || *status == reqwest::StatusCode::TOO_MANY_REQUESTS.as_u16()
791 || *status == reqwest::StatusCode::REQUEST_TIMEOUT.as_u16()
792 }
793 GatewayError::UpstreamTimeout => true,
794 _ => false,
795 }
796}
797
798fn require_jev_block(req: &ChatRequest, legs: &[ChainLeg]) -> Result<(), GatewayError> {
802 let needs_block = legs.iter().any(|l| l.provider == "typesafe");
803 let has_block = req.jev.as_ref().is_some_and(|j| !j.questions.is_empty());
804 match (needs_block, has_block) {
805 (true, false) => Err(GatewayError::BadRequest(format!(
806 "route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
807 req.model
808 ))),
809 _ => Ok(()),
810 }
811}
812
813fn non_typesafe_legs(legs: &[ChainLeg]) -> Vec<ChainLeg> {
816 legs.iter()
817 .filter(|l| l.provider != "typesafe")
818 .cloned()
819 .collect()
820}
821
822impl GatewayBuilder {
823 pub fn routes(mut self, routes: RouteTable) -> Self {
824 self.routes = Some(routes);
825 self
826 }
827 pub fn catalog(mut self, catalog: Catalog) -> Self {
828 self.catalog = Some(catalog);
829 self
830 }
831 pub fn pricing(mut self, pricing: PricingTable) -> Self {
832 self.pricing = Some(pricing);
833 self
834 }
835 pub fn ledger(mut self, ledger: LedgerHandle) -> Self {
836 self.ledger = Some(ledger);
837 self
838 }
839 pub fn vertex_native(mut self, v: Option<VertexNativeProvider>) -> Self {
840 self.vertex_native = v;
841 self
842 }
843 pub fn jev_native(mut self, v: Option<JevNativeProvider>) -> Self {
844 self.jev_native = v;
845 self
846 }
847 pub fn timeouts(mut self, t: StreamTimeouts) -> Self {
848 self.timeouts = Some(t);
849 self
850 }
851 pub fn default_tenant(mut self, t: impl Into<String>) -> Self {
852 self.default_tenant = Some(t.into());
853 self
854 }
855 pub fn embed_routes(mut self, t: crate::routing::embeddings::EmbeddingRouteTable) -> Self {
856 self.embed_routes = Some(t);
857 self
858 }
859 pub fn embedder(
860 mut self,
861 id: impl Into<String>,
862 e: Arc<dyn crate::embeddings::EmbeddingProvider>,
863 ) -> Self {
864 self.embedders
865 .get_or_insert_with(Default::default)
866 .insert(id.into(), e);
867 self
868 }
869 pub fn embed_default_input_per_mtok(mut self, v: f64) -> Self {
870 self.embed_default_input_per_mtok = Some(v);
871 self
872 }
873 pub fn guard(mut self, guard: GuardEngine) -> Self {
874 self.guard = Some(guard);
875 self
876 }
877 pub fn ai_task_types(mut self, t: crate::ai_task_type::AiTaskTypeTable) -> Self {
878 self.ai_task_types = Some(t);
879 self
880 }
881
882 pub fn build(self) -> anyhow::Result<Gateway> {
883 Ok(Gateway {
884 routes: Arc::new(
885 self.routes
886 .ok_or_else(|| anyhow::anyhow!("Gateway: routes required"))?,
887 ),
888 catalog: Arc::new(
889 self.catalog
890 .ok_or_else(|| anyhow::anyhow!("Gateway: catalog required"))?,
891 ),
892 pricing: Arc::new(
893 self.pricing
894 .ok_or_else(|| anyhow::anyhow!("Gateway: pricing required"))?,
895 ),
896 ledger: self
897 .ledger
898 .ok_or_else(|| anyhow::anyhow!("Gateway: ledger required"))?,
899 vertex_native: self.vertex_native.map(Arc::new),
900 jev_native: self.jev_native.map(Arc::new),
901 timeouts: self.timeouts.unwrap_or_default(),
902 default_tenant: self.default_tenant.unwrap_or_else(|| "unattributed".into()),
903 embed_routes: Arc::new(self.embed_routes.unwrap_or_default()),
904 embedders: self.embedders.unwrap_or_default(),
905 embed_default_input_per_mtok: self.embed_default_input_per_mtok.unwrap_or(0.10),
906 guard: Arc::new(self.guard.unwrap_or_else(GuardEngine::empty)),
907 ai_task_types: Arc::new(self.ai_task_types.unwrap_or_default()),
908 })
909 }
910}
911
912pub(crate) struct StreamSideEffects {
917 ledger: LedgerHandle,
918 pricing: Arc<PricingTable>,
919 route: String,
920 tenant: String,
921 attribution: Attribution,
922 provider: String,
923 model: String,
924 lane: &'static str, request_id: String,
926 legs_attempted: u32,
927 started: Instant,
928 input_tokens: u64,
929 output_tokens: u64,
930 status: &'static str,
931 fired: bool,
932}
933
934impl StreamSideEffects {
935 #[allow(clippy::too_many_arguments)]
936 pub(crate) fn new(
937 ledger: LedgerHandle,
938 pricing: Arc<PricingTable>,
939 route: String,
940 tenant: String,
941 attribution: Attribution,
942 provider: String,
943 model: String,
944 lane: &'static str,
945 request_id: String,
946 legs_attempted: u32,
947 started: Instant,
948 ) -> Self {
949 Self {
950 ledger,
951 pricing,
952 route,
953 tenant,
954 attribution,
955 provider,
956 model,
957 lane,
958 request_id,
959 legs_attempted,
960 started,
961 input_tokens: 0,
962 output_tokens: 0,
963 status: "ok",
964 fired: false,
965 }
966 }
967
968 pub(crate) fn observe(&mut self, item: &StreamItem) {
970 if let StreamItem::Done {
971 input_tokens,
972 output_tokens,
973 ..
974 } = item
975 {
976 self.input_tokens = *input_tokens;
977 self.output_tokens = *output_tokens;
978 }
979 }
980
981 pub(crate) fn mark_error(&mut self) {
983 self.status = "error";
984 }
985}
986
987impl Drop for StreamSideEffects {
988 fn drop(&mut self) {
989 if self.fired {
990 return;
991 }
992 self.fired = true;
993
994 let cost = self.pricing.cost_usd(
996 &self.provider,
997 &self.model,
998 self.input_tokens,
999 self.output_tokens,
1000 );
1001 self.ledger.enqueue(UsageEntry {
1002 ts: Utc::now(),
1003 tenant: self.tenant.clone(),
1004 workspace: self.attribution.workspace.clone(),
1005 user: self.attribution.user.clone(),
1006 thread: self.attribution.thread.clone(),
1007 message: self.attribution.message.clone(),
1008 route: self.route.clone(),
1009 provider: self.provider.clone(),
1010 model: self.model.clone(),
1011 lane: self.lane.to_string(),
1012 input_tokens: self.input_tokens,
1013 output_tokens: self.output_tokens,
1014 cost_usd: cost,
1015 request_id: self.request_id.clone(),
1016 status: self.status.to_string(),
1017 op: "chat".into(),
1018 user_task_type: self.attribution.user_task_type.clone(),
1019 ai_task_type: self.attribution.ai_task_type.clone(),
1020 });
1021
1022 let completion = Completion {
1025 provider: self.provider.clone(),
1026 model: self.model.clone(),
1027 content: String::new(),
1028 tool_calls: Vec::new(),
1029 finish_reason: FinishReason::Stop,
1030 input_tokens: self.input_tokens,
1031 output_tokens: self.output_tokens,
1032 };
1033 let lane = if self.lane == "native" {
1034 Lane::NativeVertex
1035 } else {
1036 Lane::Standard
1037 };
1038 GenAiSpan::from_completion(
1039 &completion,
1040 lane,
1041 &self.route,
1042 &self.tenant,
1043 self.attribution.workspace.as_deref(),
1044 self.legs_attempted,
1045 true,
1046 )
1047 .emit_metrics(self.started.elapsed().as_secs_f64());
1048 }
1049}
1050
1051pub struct GuardedStream {
1055 inner: BoxStream<'static, Result<StreamItem, LegError>>,
1056 model: String,
1057 guard: StreamSideEffects,
1058}
1059
1060impl GuardedStream {
1061 pub(crate) fn new(
1062 inner: BoxStream<'static, Result<StreamItem, LegError>>,
1063 model: String,
1064 guard: StreamSideEffects,
1065 ) -> Self {
1066 Self {
1067 inner,
1068 model,
1069 guard,
1070 }
1071 }
1072
1073 pub fn model(&self) -> &str {
1075 &self.model
1076 }
1077}
1078
1079impl Stream for GuardedStream {
1080 type Item = Result<StreamItem, LegError>;
1081 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1082 let this = self.get_mut();
1084 match this.inner.as_mut().poll_next(cx) {
1085 Poll::Ready(Some(Ok(item))) => {
1086 this.guard.observe(&item);
1087 Poll::Ready(Some(Ok(item)))
1088 }
1089 Poll::Ready(Some(Err(e))) => {
1090 this.guard.mark_error();
1091 Poll::Ready(Some(Err(e)))
1092 }
1093 Poll::Ready(None) => Poll::Ready(None),
1094 Poll::Pending => Poll::Pending,
1095 }
1096 }
1097}
1098
1099pub(crate) async fn collect_committed(
1103 committed: CommittedStream,
1104) -> Result<Completion, GatewayError> {
1105 use futures::TryStreamExt;
1106 let CommittedStream {
1107 provider,
1108 model,
1109 stream,
1110 } = committed;
1111 stream
1112 .map_err(|e: LegError| GatewayError::Upstream {
1113 status: 502,
1114 body: e.to_string(),
1115 })
1116 .try_fold(Accumulator::default(), |mut acc, item| async move {
1117 acc.push(item);
1118 Ok(acc)
1119 })
1120 .await
1121 .map(|acc| Completion {
1122 provider,
1123 model,
1124 content: acc.content,
1125 tool_calls: acc.tool_calls,
1126 finish_reason: acc.finish_reason,
1127 input_tokens: acc.input_tokens,
1128 output_tokens: acc.output_tokens,
1129 })
1130}
1131
1132#[cfg(test)]
1133mod tests {
1134 use super::*;
1135 use crate::ledger::{InMemoryLedger, LedgerHandle, LedgerStore};
1136
1137 fn test_gateway() -> Gateway {
1138 let routes = RouteTable::from_toml_str(
1139 r#"[routes."fast"]
1140 legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1141 )
1142 .unwrap();
1143 let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1144 let ledger = LedgerHandle::spawn(
1145 Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1146 16,
1147 );
1148 Gateway::builder()
1149 .routes(routes)
1150 .catalog(catalog)
1151 .pricing(PricingTable::default())
1152 .ledger(ledger)
1153 .default_tenant("acme")
1154 .build()
1155 .unwrap()
1156 }
1157
1158 #[tokio::test]
1159 async fn builder_builds_and_lists_aliases() {
1160 let gw = test_gateway();
1161 assert_eq!(gw.model_aliases(), vec!["fast".to_string()]);
1162 assert_eq!(gw.default_tenant, "acme");
1163 }
1164
1165 #[test]
1166 fn builder_requires_components() {
1167 assert!(Gateway::builder().build().is_err());
1168 }
1169
1170 #[tokio::test]
1171 async fn guard_records_usage_on_drop() {
1172 let store = Arc::new(InMemoryLedger::default());
1173 let ledger = LedgerHandle::spawn(store.clone(), 16);
1174 let pricing = Arc::new(PricingTable::default());
1175 {
1176 let mut guard = StreamSideEffects::new(
1177 ledger.clone(),
1178 pricing,
1179 "route".into(),
1180 "tenant".into(),
1181 Attribution::default(),
1182 "p".into(),
1183 "m".into(),
1184 "standard",
1185 "rid".into(),
1186 1,
1187 Instant::now(),
1188 );
1189 guard.observe(&crate::routing::stream::StreamItem::Done {
1190 input_tokens: 3,
1191 output_tokens: 2,
1192 finish_reason: crate::routing::stream::FinishReason::Stop,
1193 });
1194 } tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1196 let rows = store.entries();
1197 assert_eq!(rows.len(), 1);
1198 assert_eq!(rows[0].input_tokens, 3);
1199 assert_eq!(rows[0].output_tokens, 2);
1200 assert_eq!(rows[0].status, "ok");
1201 }
1202
1203 #[tokio::test]
1204 async fn guarded_stream_yields_items_and_records_on_drop() {
1205 use crate::routing::stream::{FinishReason, StreamItem};
1206 use futures::StreamExt;
1207 let store = Arc::new(InMemoryLedger::default());
1208 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1209 let inner = futures::stream::iter(vec![
1210 Ok(StreamItem::Delta("hi".into())),
1211 Ok(StreamItem::Done {
1212 input_tokens: 4,
1213 output_tokens: 2,
1214 finish_reason: FinishReason::Stop,
1215 }),
1216 ])
1217 .boxed();
1218 let guard = StreamSideEffects::new(
1219 ledger,
1220 Arc::new(PricingTable::default()),
1221 "route".into(),
1222 "acme".into(),
1223 Attribution::default(),
1224 "p".into(),
1225 "m".into(),
1226 "standard",
1227 "rid".into(),
1228 1,
1229 std::time::Instant::now(),
1230 );
1231 {
1232 let mut gs = GuardedStream::new(inner, "m".into(), guard);
1233 let mut n = 0;
1234 while let Some(item) = gs.next().await {
1235 item.unwrap();
1236 n += 1;
1237 }
1238 assert_eq!(n, 2);
1239 assert_eq!(gs.model(), "m");
1240 } tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1242 let rows = store.entries();
1243 assert_eq!(rows.len(), 1);
1244 assert_eq!(rows[0].input_tokens, 4);
1245 assert_eq!(rows[0].tenant, "acme");
1246 }
1247
1248 fn resolver_gateway(ai_task_types: &str) -> Gateway {
1250 let routes = RouteTable::from_toml_str(
1251 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1252 )
1253 .unwrap();
1254 Gateway::builder()
1255 .routes(routes)
1256 .catalog(Catalog::for_test(vec![(
1257 "qwen",
1258 "http://127.0.0.1:1/v1".into(),
1259 )]))
1260 .pricing(PricingTable::default())
1261 .ledger(LedgerHandle::spawn(
1262 Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1263 4,
1264 ))
1265 .ai_task_types(
1266 crate::ai_task_type::AiTaskTypeTable::from_toml_str(ai_task_types).unwrap(),
1267 )
1268 .build()
1269 .unwrap()
1270 }
1271
1272 #[tokio::test]
1273 async fn ai_task_type_prefers_the_caller_header_over_the_alias_mapping() {
1274 let gw = resolver_gateway("conversation = [\"fast\"]");
1275 let ctx = RequestCtx {
1276 ai_task_type: Some("caller-supplied".into()),
1277 ..Default::default()
1278 };
1279 assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "caller-supplied");
1280 }
1281
1282 #[tokio::test]
1283 async fn ai_task_type_is_inferred_from_the_route_alias_when_no_header() {
1284 let gw = resolver_gateway("conversation = [\"fast\", \"planning\"]");
1285 let ctx = RequestCtx::default();
1286 assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
1287 assert_eq!(gw.ai_task_type_of(&ctx, "planning"), "conversation");
1288 }
1289
1290 #[tokio::test]
1291 async fn ai_task_type_defaults_to_simple_for_an_unmapped_alias() {
1292 let gw = resolver_gateway("conversation = [\"fast\"]");
1293 assert_eq!(
1294 gw.ai_task_type_of(&RequestCtx::default(), "graph-llm"),
1295 "simple"
1296 );
1297 }
1298
1299 #[tokio::test]
1300 async fn an_empty_ai_task_type_header_falls_back_to_inference() {
1301 let gw = resolver_gateway("conversation = [\"fast\"]");
1302 let ctx = RequestCtx {
1303 ai_task_type: Some(String::new()),
1304 ..Default::default()
1305 };
1306 assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
1307 }
1308
1309 #[tokio::test]
1310 async fn chat_returns_completion_and_records_ledger() {
1311 use wiremock::matchers::{method, path};
1312 use wiremock::{Mock, MockServer, ResponseTemplate};
1313 let mock = MockServer::start().await;
1314 Mock::given(method("POST")).and(path("/v1/chat/completions"))
1315 .respond_with(ResponseTemplate::new(200)
1316 .insert_header("content-type", "text/event-stream")
1317 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\n\
1318 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2}}\n\n\
1319 data: [DONE]\n\n"))
1320 .mount(&mock).await;
1321 let routes = RouteTable::from_toml_str(
1322 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1323 )
1324 .unwrap();
1325 let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
1326 let store = Arc::new(InMemoryLedger::default());
1327 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1328 let gw = Gateway::builder()
1329 .routes(routes)
1330 .catalog(catalog)
1331 .pricing(PricingTable::default())
1332 .ledger(ledger)
1333 .default_tenant("def")
1334 .ai_task_types(
1335 crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
1336 .unwrap(),
1337 )
1338 .build()
1339 .unwrap();
1340 let req = serde_json::from_value(serde_json::json!(
1341 {"model":"fast","messages":[{"role":"user","content":"hi"}]}))
1342 .unwrap();
1343 let ctx = RequestCtx {
1344 tenant: Some("acme".into()),
1345 workspace: Some("ws-9".into()),
1346 user: Some("user-42".into()),
1347 thread: Some("thread-9".into()),
1348 message: Some("msg-7".into()),
1349 user_task_type: Some("summarisation".into()),
1350 ai_task_type: None,
1352 request_id: Some("corr-123".into()),
1353 };
1354 let c = match gw.chat(req, &ctx).await.unwrap() {
1355 ChatOutcome::Plain(c) => c,
1356 ChatOutcome::Hybrid(_) => panic!("expected plain completion"),
1357 };
1358 assert_eq!(c.content, "hi");
1359 assert_eq!(c.input_tokens, 3);
1360 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1361 let rows = store.entries();
1362 assert_eq!(rows.len(), 1);
1363 assert_eq!(rows[0].tenant, "acme");
1364 assert_eq!(rows[0].user.as_deref(), Some("user-42"));
1365 assert_eq!(rows[0].thread.as_deref(), Some("thread-9"));
1366 assert_eq!(rows[0].message.as_deref(), Some("msg-7"));
1367 assert_eq!(rows[0].user_task_type.as_deref(), Some("summarisation"));
1368 assert_eq!(rows[0].ai_task_type, "conversation");
1370 assert_eq!(rows[0].request_id, "corr-123");
1372 }
1373
1374 #[tokio::test]
1375 async fn chat_stream_yields_items_and_records() {
1376 use crate::routing::stream::StreamItem;
1377 use futures::StreamExt;
1378 use wiremock::matchers::{method, path};
1379 use wiremock::{Mock, MockServer, ResponseTemplate};
1380 let mock = MockServer::start().await;
1381 Mock::given(method("POST")).and(path("/v1/chat/completions"))
1382 .respond_with(ResponseTemplate::new(200)
1383 .insert_header("content-type", "text/event-stream")
1384 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
1385 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
1386 data: [DONE]\n\n"))
1387 .mount(&mock).await;
1388 let routes = RouteTable::from_toml_str(
1389 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1390 )
1391 .unwrap();
1392 let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
1393 let store = Arc::new(InMemoryLedger::default());
1394 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1395 let gw = Gateway::builder()
1396 .routes(routes)
1397 .catalog(catalog)
1398 .pricing(PricingTable::default())
1399 .ledger(ledger)
1400 .default_tenant("def")
1401 .ai_task_types(
1402 crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
1403 .unwrap(),
1404 )
1405 .build()
1406 .unwrap();
1407 let req = serde_json::from_value(serde_json::json!(
1408 {"model":"fast","stream":true,"messages":[{"role":"user","content":"hi"}]}))
1409 .unwrap();
1410 let ctx = RequestCtx {
1411 user_task_type: Some("code-review".into()),
1412 ..Default::default()
1413 };
1414 let mut stream = gw.chat_stream(req, &ctx).await.unwrap();
1415 let mut got = false;
1416 while let Some(i) = stream.next().await {
1417 if matches!(i.unwrap(), StreamItem::Delta(ref t) if t == "go") {
1418 got = true;
1419 }
1420 }
1421 drop(stream);
1422 assert!(got);
1423 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1424 let rows = store.entries();
1425 assert_eq!(rows.len(), 1);
1426 assert_eq!(rows[0].user_task_type.as_deref(), Some("code-review"));
1427 assert_eq!(rows[0].ai_task_type, "conversation");
1428 }
1429
1430 use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
1433 use async_trait::async_trait;
1434
1435 struct FlakyEmbedder;
1437 #[async_trait]
1438 impl EmbeddingProvider for FlakyEmbedder {
1439 async fn embed(
1440 &self,
1441 _model: &str,
1442 _inputs: &[String],
1443 _dims: u32,
1444 ) -> Result<EmbedOut, GatewayError> {
1445 Err(GatewayError::Upstream {
1446 status: 500,
1447 body: "x".into(),
1448 })
1449 }
1450 }
1451
1452 struct GoodEmbedder;
1454 #[async_trait]
1455 impl EmbeddingProvider for GoodEmbedder {
1456 async fn embed(
1457 &self,
1458 _model: &str,
1459 inputs: &[String],
1460 dims: u32,
1461 ) -> Result<EmbedOut, GatewayError> {
1462 Ok(EmbedOut {
1463 vectors: inputs.iter().map(|_| vec![0.0f32; dims as usize]).collect(),
1464 input_tokens: 6,
1465 })
1466 }
1467 }
1468
1469 fn embed_gateway() -> (Gateway, Arc<InMemoryLedger>) {
1470 let embed_routes = crate::routing::embeddings::EmbeddingRouteTable::from_toml_str(
1471 r#"
1472 [embeddings."default-embed"]
1473 dimensions = 4
1474 legs = [
1475 { provider = "flaky", model = "flaky-embed" },
1476 { provider = "good", model = "good-embed" },
1477 ]
1478 "#,
1479 )
1480 .unwrap();
1481 let routes = RouteTable::from_toml_str(
1482 "[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
1483 )
1484 .unwrap();
1485 let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1486 let store = Arc::new(InMemoryLedger::default());
1487 let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
1488 let gw = Gateway::builder()
1489 .routes(routes)
1490 .catalog(catalog)
1491 .pricing(PricingTable::default())
1492 .ledger(ledger)
1493 .default_tenant("acme")
1494 .embed_routes(embed_routes)
1495 .embedder(
1496 "flaky",
1497 Arc::new(FlakyEmbedder) as Arc<dyn EmbeddingProvider>,
1498 )
1499 .embedder("good", Arc::new(GoodEmbedder) as Arc<dyn EmbeddingProvider>)
1500 .build()
1501 .unwrap();
1502 (gw, store)
1503 }
1504
1505 #[tokio::test]
1506 async fn embed_falls_through_to_good_leg_and_records_usage() {
1507 let (gw, store) = embed_gateway();
1508 let req = EmbeddingRequest {
1509 input: EmbeddingInput::Many(vec!["a".into(), "b".into()]),
1510 model: "default-embed".into(),
1511 dimensions: None,
1512 };
1513 let ctx = RequestCtx {
1514 user_task_type: Some("retrieval".into()),
1515 ..Default::default()
1516 };
1517 let resp = gw.embed(req, ctx).await.unwrap();
1518 assert_eq!(resp.data.len(), 2);
1519 assert!(resp.data.iter().all(|d| d.embedding.len() == 4));
1520 assert_eq!(resp.data[0].index, 0);
1521 assert_eq!(resp.data[1].index, 1);
1522 assert_eq!(resp.usage.prompt_tokens, 6);
1523
1524 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1525 let rows = store.entries();
1526 assert_eq!(rows.len(), 1);
1527 assert_eq!(rows[0].op, "embedding");
1528 assert_eq!(rows[0].lane, "embedding");
1529 assert_eq!(rows[0].output_tokens, 0);
1530 assert_eq!(rows[0].provider, "good");
1531 assert!(rows[0].cost_usd > 0.0);
1532 assert_eq!(rows[0].user_task_type.as_deref(), Some("retrieval"));
1533 assert_eq!(rows[0].ai_task_type, "simple");
1534 }
1535
1536 #[tokio::test]
1537 async fn embed_dimension_mismatch_is_bad_request() {
1538 let (gw, _store) = embed_gateway();
1539 let req = EmbeddingRequest {
1540 input: EmbeddingInput::Many(vec!["a".into()]),
1541 model: "default-embed".into(),
1542 dimensions: Some(8),
1543 };
1544 let err = gw.embed(req, RequestCtx::default()).await.unwrap_err();
1545 assert!(matches!(err, GatewayError::BadRequest(_)));
1546 }
1547
1548 #[tokio::test]
1549 async fn chat_blocks_when_route_policy_refuses() {
1550 use crate::guard::{GuardEngine, GuardrailsConfig};
1551 let routes = RouteTable::from_toml_str(
1552 r#"[routes."fast"]
1553 policy = "strict"
1554 legs = [{ provider = "qwen", model = "qwen-max" }]"#,
1555 )
1556 .unwrap();
1557 let guard = GuardEngine::from_config(
1558 &GuardrailsConfig::from_toml_str(
1559 r#"[guardrails.strict]
1560 scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#,
1561 )
1562 .unwrap(),
1563 )
1564 .unwrap();
1565 let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
1566 let ledger = LedgerHandle::spawn(
1567 Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
1568 16,
1569 );
1570 let gw = Gateway::builder()
1571 .routes(routes)
1572 .catalog(catalog)
1573 .pricing(PricingTable::default())
1574 .ledger(ledger)
1575 .guard(guard)
1576 .build()
1577 .unwrap();
1578 let req = serde_json::from_value(serde_json::json!({
1579 "model": "fast",
1580 "messages": [{ "role": "user", "content": "this is forbidden" }]
1581 }))
1582 .unwrap();
1583 let err = gw.chat(req, &RequestCtx::default()).await.unwrap_err();
1584 assert!(matches!(err, GatewayError::ContentBlocked { .. }));
1585 }
1586}