use std::borrow::Cow;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Instant;
use chrono::Utc;
use futures::stream::{BoxStream, Stream};
use tap::Tap;
use uuid::Uuid;
use crate::error::{GatewayError, LegFailure};
use crate::guard::GuardEngine;
use crate::jev_native::JevNativeProvider;
use crate::ledger::{LedgerHandle, UsageEntry};
use crate::observability::GenAiSpan;
use crate::pricing::PricingTable;
use crate::providers::Catalog;
use crate::routing::classify::{classify, vertex_triggers, Lane};
use crate::routing::effort::Effort;
use crate::routing::executor::{
execute_buffered_with_timeouts, CommittedStream, Completion, LegError, StreamTimeouts,
};
use crate::routing::jev_router::RoutingReport;
use crate::routing::request::{ChatRequest, VertexExt};
use crate::routing::stream::{Accumulator, FinishReason, StreamItem};
use crate::routing::table::{ChainLeg, RouteTable};
use crate::telemetry::GatewayMetrics;
use crate::vertex_native::VertexNativeProvider;
#[derive(Clone)]
pub struct Gateway {
pub(crate) routes: Arc<RouteTable>,
pub(crate) catalog: Arc<Catalog>,
pub(crate) pricing: Arc<PricingTable>,
pub(crate) ledger: LedgerHandle,
pub(crate) vertex_native: Option<Arc<VertexNativeProvider>>,
pub(crate) jev_native: Option<Arc<JevNativeProvider>>,
pub(crate) timeouts: StreamTimeouts,
pub(crate) default_tenant: String,
pub(crate) embed_routes: Arc<crate::routing::embeddings::EmbeddingRouteTable>,
pub(crate) embedders:
std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>,
pub(crate) embed_default_input_per_mtok: f64,
pub(crate) guard: Arc<GuardEngine>,
pub(crate) ai_task_types: Arc<crate::ai_task_type::AiTaskTypeTable>,
pub(crate) metrics: Arc<GatewayMetrics>,
}
#[derive(Debug, Clone, Default)]
pub struct RequestCtx {
pub tenant: Option<String>,
pub workspace: Option<String>,
pub user: Option<String>,
pub thread: Option<String>,
pub message: Option<String>,
pub user_task_type: Option<String>,
pub ai_task_type: Option<String>,
pub request_id: Option<String>,
}
impl RequestCtx {
pub fn resolved_request_id(&self) -> Option<String> {
self.request_id
.clone()
.or_else(|| self.message.clone())
.filter(|s| !s.is_empty())
}
}
pub(crate) enum JevAttempt {
Decided {
completion: Completion,
answers: serde_json::Value,
},
Exhausted(Vec<LegFailure>),
}
#[derive(Debug)]
pub enum ChatOutcome {
Plain(Completion),
Hybrid(HybridOutcome),
}
#[derive(Debug)]
pub struct HybridOutcome {
pub model: String,
pub answers: serde_json::Value,
pub survivors: Vec<String>,
pub degraded: bool,
pub extraction_ran: bool,
pub extractions: serde_json::Map<String, serde_json::Value>,
pub input_tokens: u64,
pub output_tokens: u64,
}
#[derive(Debug, Clone)]
pub struct Attribution {
pub workspace: Option<String>,
pub user: Option<String>,
pub thread: Option<String>,
pub message: Option<String>,
pub user_task_type: Option<String>,
pub ai_task_type: String,
}
impl Default for Attribution {
fn default() -> Self {
Self {
workspace: None,
user: None,
thread: None,
message: None,
user_task_type: None,
ai_task_type: crate::ai_task_type::DEFAULT_AI_TASK_TYPE.to_string(),
}
}
}
#[derive(Default)]
pub struct GatewayBuilder {
routes: Option<RouteTable>,
catalog: Option<Catalog>,
pricing: Option<PricingTable>,
ledger: Option<LedgerHandle>,
vertex_native: Option<VertexNativeProvider>,
jev_native: Option<JevNativeProvider>,
timeouts: Option<StreamTimeouts>,
default_tenant: Option<String>,
embed_routes: Option<crate::routing::embeddings::EmbeddingRouteTable>,
embedders:
Option<std::collections::HashMap<String, Arc<dyn crate::embeddings::EmbeddingProvider>>>,
embed_default_input_per_mtok: Option<f64>,
guard: Option<GuardEngine>,
ai_task_types: Option<crate::ai_task_type::AiTaskTypeTable>,
metrics: Option<Arc<GatewayMetrics>>,
}
impl Gateway {
pub fn builder() -> GatewayBuilder {
GatewayBuilder::default()
}
pub fn model_aliases(&self) -> Vec<String> {
self.routes.aliases()
}
pub(crate) fn guard_input(&self, req: &ChatRequest) -> Result<(), GatewayError> {
let policy = self.routes.policy_of(&req.model).unwrap_or("default");
self.guard.guard(policy, req)
}
pub(crate) fn tenant_of<'a>(&'a self, ctx: &'a RequestCtx) -> &'a str {
ctx.tenant.as_deref().unwrap_or(&self.default_tenant)
}
pub(crate) fn ai_task_type_of(&self, ctx: &RequestCtx, alias: &str) -> String {
ctx.ai_task_type
.as_deref()
.filter(|s| !s.is_empty())
.unwrap_or_else(|| self.ai_task_types.resolve(alias))
.to_string()
}
pub(crate) fn attribution_of(&self, ctx: &RequestCtx, alias: &str) -> Attribution {
Attribution {
workspace: ctx.workspace.clone(),
user: ctx.user.clone(),
thread: ctx.thread.clone(),
message: ctx.message.clone(),
user_task_type: ctx.user_task_type.clone(),
ai_task_type: self.ai_task_type_of(ctx, alias),
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn record(
&self,
ctx: &RequestCtx,
route: &str,
lane: &'static str,
request_id: &str,
c: &Completion,
legs: u32,
started: Instant,
) {
let cost = self
.pricing
.cost_usd(&c.provider, &c.model, c.input_tokens, c.output_tokens);
let attr = self.attribution_of(ctx, route);
self.ledger.enqueue(UsageEntry {
ts: Utc::now(),
tenant: self.tenant_of(ctx).to_string(),
workspace: attr.workspace,
user: attr.user,
thread: attr.thread,
message: attr.message,
route: route.to_string(),
provider: c.provider.clone(),
model: c.model.clone(),
lane: lane.to_string(),
input_tokens: c.input_tokens,
output_tokens: c.output_tokens,
cost_usd: cost,
request_id: request_id.to_string(),
status: "ok".into(),
op: "chat".into(),
user_task_type: attr.user_task_type,
ai_task_type: attr.ai_task_type,
});
let span_lane = match lane {
"native" => Lane::NativeVertex,
"jev" => Lane::Jev,
_ => Lane::Standard,
};
GenAiSpan::from_completion(
c,
span_lane,
route,
self.tenant_of(ctx),
ctx.workspace.as_deref(),
legs,
false,
)
.emit_metrics(&self.metrics, started.elapsed().as_secs_f64());
}
pub(crate) async fn native_committed(
&self,
req: &ChatRequest,
legs: &[ChainLeg],
) -> Result<CommittedStream, GatewayError> {
let provider = self
.vertex_native
.as_ref()
.ok_or_else(|| GatewayError::BadRequest("native vertex lane not configured".into()))?;
let vertex_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "vertex").collect();
if vertex_legs.is_empty() {
return Err(GatewayError::NativeFeatureUnsupported {
feature: "native-vertex".into(),
route: req.model.clone(),
});
}
let mut failures: Vec<LegFailure> = Vec::new();
let mut last_retryable: Option<GatewayError> = None;
for leg in &vertex_legs {
match provider
.stream_generate(
&leg.model,
&with_leg_thinking(req, leg),
leg.region.as_deref(),
)
.await
{
Ok(stream) => {
return Ok(CommittedStream::single(
"vertex".into(),
leg.model.clone(),
stream,
));
}
Err(e) if native_start_retryable(&e) => {
failures.push(LegFailure {
provider: leg.provider.clone(),
model: leg.model.clone(),
message: e.to_string(),
});
last_retryable = Some(e);
}
Err(e) => return Err(e),
}
}
Err(
last_retryable.unwrap_or_else(|| GatewayError::AllLegsFailed {
route: req.model.clone(),
failures,
}),
)
}
pub(crate) async fn jev_attempt(
&self,
req: &ChatRequest,
legs: &[ChainLeg],
) -> Result<JevAttempt, GatewayError> {
let jev_legs: Vec<&ChainLeg> = legs.iter().filter(|l| l.provider == "typesafe").collect();
if jev_legs.is_empty() {
return Ok(JevAttempt::Exhausted(Vec::new()));
}
let ext = req
.jev
.as_ref()
.filter(|j| !j.questions.is_empty())
.ok_or_else(|| {
GatewayError::BadRequest(format!(
"route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
req.model
))
})?;
let provider = self.jev_native.as_ref().ok_or_else(|| {
GatewayError::BadRequest(
"typesafe legs require TYPESAFE_API_KEY to be configured".into(),
)
})?;
let state = ext
.state
.clone()
.or_else(|| serde_json::to_string(&req.messages).ok())
.unwrap_or_default();
let mut body = serde_json::json!({
"state": state,
"questions": ext.questions.clone(),
});
let mut failures: Vec<LegFailure> = Vec::new();
for leg in jev_legs {
let model = if leg.model.is_empty() {
crate::jev_native::DEFAULT_MODEL.to_string()
} else {
leg.model.clone()
};
body["model"] = serde_json::Value::String(model.clone());
let resp = match provider.evaluate(body.clone()).await {
Ok(r) => r,
Err(e) if native_start_retryable(&e) => {
failures.push(LegFailure {
provider: leg.provider.clone(),
model,
message: e.to_string(),
});
continue;
}
Err(e) => return Err(e),
};
let status = resp.status();
let bytes = resp.bytes().await.map_err(|e| GatewayError::Upstream {
status: 502,
body: e.to_string(),
})?;
if status.is_success() {
let value: serde_json::Value =
serde_json::from_slice(&bytes).map_err(|e| GatewayError::Upstream {
status: 502,
body: format!("jev response is not JSON: {e}"),
})?;
return Ok(JevAttempt::Decided {
answers: value["answers"].clone(),
completion: Completion {
provider: "typesafe".into(),
model: value["model"].as_str().unwrap_or(&model).to_string(),
content: serde_json::to_string(&value["answers"])
.unwrap_or_else(|_| "{}".into()),
tool_calls: Vec::new(),
finish_reason: FinishReason::Stop,
input_tokens: value["usage"]["input_tokens"].as_u64().unwrap_or(0),
output_tokens: value["usage"]["output_tokens"].as_u64().unwrap_or(0),
},
});
}
let message = String::from_utf8_lossy(&bytes).into_owned();
if status.is_server_error()
|| status == reqwest::StatusCode::TOO_MANY_REQUESTS
|| status == reqwest::StatusCode::REQUEST_TIMEOUT
{
failures.push(LegFailure {
provider: leg.provider.clone(),
model,
message: format!("{status}: {message}"),
});
continue;
}
return Err(GatewayError::BadRequest(format!(
"typesafe {}: {message}",
status.as_u16()
)));
}
Ok(JevAttempt::Exhausted(failures))
}
fn jev_committed(c: Completion) -> CommittedStream {
let items: Vec<Result<StreamItem, LegError>> = vec![
Ok(StreamItem::Delta(c.content.clone())),
Ok(StreamItem::Done {
input_tokens: c.input_tokens,
output_tokens: c.output_tokens,
finish_reason: c.finish_reason,
}),
];
CommittedStream::single("typesafe".into(), c.model, futures::stream::iter(items))
}
pub async fn chat_stream(
&self,
req: ChatRequest,
ctx: &RequestCtx,
) -> Result<GuardedStream, GatewayError> {
use crate::routing::executor::execute_streaming_with_timeouts;
let started = Instant::now();
let request_id = ctx
.resolved_request_id()
.unwrap_or_else(|| Uuid::new_v4().to_string());
let plan = self.plan_route(&req, ctx, &request_id).await?;
let legs = plan.chain();
let vertex_leg_count = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
require_jev_block(&req, &legs)?;
if let Err(message) = crate::routing::jev_extract::validate_extract(
&req,
legs.iter().any(|l| l.provider != "typesafe"),
legs.iter().any(|l| l.provider == "vertex"),
) {
return Err(GatewayError::BadRequest(message));
}
let (committed, lane_str, legs_attempted) = match classify(&req) {
Lane::Standard => (
execute_streaming_with_timeouts(
&self.catalog,
&req.model,
&legs,
&req,
self.timeouts,
)
.await?,
"standard",
legs.len() as u32,
),
Lane::NativeVertex => (
self.native_committed(&req, &legs).await?,
"native",
vertex_leg_count.max(1),
),
Lane::Jev => match self.jev_attempt(&req, &legs).await? {
JevAttempt::Decided { completion, .. } => (
Self::jev_committed(completion),
"jev",
legs.iter()
.filter(|l| l.provider == "typesafe")
.count()
.max(1) as u32,
),
JevAttempt::Exhausted(failures) => {
let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
if rest.is_empty() {
return Err(GatewayError::AllLegsFailed {
route: req.model.clone(),
failures,
});
}
if vertex_triggers(&req) {
let vertex_n =
rest.iter().filter(|l| l.provider == "vertex").count() as u32;
(
self.native_committed(&req, &rest).await?,
"native",
vertex_n.max(1),
)
} else {
(
execute_streaming_with_timeouts(
&self.catalog,
&req.model,
&rest,
&req,
self.timeouts,
)
.await?,
"standard",
rest.len() as u32,
)
}
}
},
};
let model = committed.model.clone();
let routing = plan.report_for(Some((
committed.provider.as_str(),
committed.model.as_str(),
)));
let guard = StreamSideEffects::new(
self.ledger.clone(),
self.pricing.clone(),
self.metrics.clone(),
req.model.clone(),
self.tenant_of(ctx).to_string(),
self.attribution_of(ctx, &req.model),
committed.provider.clone(),
committed.model.clone(),
lane_str,
request_id,
legs_attempted,
started,
);
Ok(GuardedStream::new(committed.stream, model, guard, routing))
}
pub async fn chat(
&self,
req: ChatRequest,
ctx: &RequestCtx,
) -> Result<ChatOutcome, GatewayError> {
self.chat_routed(req, ctx).await.map(|(outcome, _)| outcome)
}
pub async fn chat_routed(
&self,
req: ChatRequest,
ctx: &RequestCtx,
) -> Result<(ChatOutcome, RoutingReport), GatewayError> {
let started = Instant::now();
let request_id = ctx
.resolved_request_id()
.unwrap_or_else(|| Uuid::new_v4().to_string());
let plan = self.plan_route(&req, ctx, &request_id).await?;
let legs = plan.chain();
require_jev_block(&req, &legs)?;
if let Err(message) = crate::routing::jev_extract::validate_extract(
&req,
legs.iter().any(|l| l.provider != "typesafe"),
legs.iter().any(|l| l.provider == "vertex"),
) {
return Err(GatewayError::BadRequest(message));
}
if req.jev.as_ref().and_then(|j| j.extract.as_ref()).is_some() {
let outcome = self
.chat_hybrid(req, ctx, &legs, started, &request_id)
.await?;
return Ok((ChatOutcome::Hybrid(outcome), plan.report_for(None)));
}
let (completion, lane_str, legs_n) = match classify(&req) {
Lane::Standard => execute_buffered_with_timeouts(
&self.catalog,
&req.model,
&legs,
&req,
self.timeouts,
)
.await
.map(|c| (c, "standard", legs.len() as u32))?,
Lane::NativeVertex => {
let vertex_n = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
let committed = self.native_committed(&req, &legs).await?;
(
collect_committed(committed).await?,
"native",
vertex_n.max(1),
)
}
Lane::Jev => match self.jev_attempt(&req, &legs).await? {
JevAttempt::Decided { completion, .. } => (
completion,
"jev",
legs.iter()
.filter(|l| l.provider == "typesafe")
.count()
.max(1) as u32,
),
JevAttempt::Exhausted(failures) => {
let rest: Vec<ChainLeg> = non_typesafe_legs(&legs);
if rest.is_empty() {
return Err(GatewayError::AllLegsFailed {
route: req.model.clone(),
failures,
});
}
if vertex_triggers(&req) {
let committed = self.native_committed(&req, &rest).await?;
let vertex_n =
rest.iter().filter(|l| l.provider == "vertex").count() as u32;
(
collect_committed(committed).await?,
"native",
vertex_n.max(1),
)
} else {
execute_buffered_with_timeouts(
&self.catalog,
&req.model,
&rest,
&req,
self.timeouts,
)
.await
.map(|c| (c, "standard", rest.len() as u32))?
}
}
},
};
self.record(
ctx,
&req.model,
lane_str,
&request_id,
&completion,
legs_n,
started,
);
let routing = plan.report_for(Some((
completion.provider.as_str(),
completion.model.as_str(),
)));
Ok((ChatOutcome::Plain(completion), routing))
}
pub async fn embed(
&self,
req: crate::embeddings::EmbeddingRequest,
ctx: RequestCtx,
) -> Result<crate::embeddings::EmbeddingResponse, GatewayError> {
let alias = req.model.clone();
let dims = self
.embed_routes
.dimensions(&alias)
.ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?;
if let Some(d) = req.dimensions {
if d != dims {
return Err(GatewayError::BadRequest(format!(
"dimensions {d} does not match embedding alias '{alias}' dimension {dims}"
)));
}
}
let inputs = req.input.into_vec();
let legs = self
.embed_routes
.legs(&alias)
.ok_or_else(|| GatewayError::UnknownModel(alias.clone()))?
.to_vec();
let mut last_err: Option<GatewayError> = None;
for leg in &legs {
let Some(embedder) = self.embedders.get(&leg.provider) else {
continue;
};
let limit = match leg.provider.as_str() {
"vertex" => crate::embeddings::vertex::VERTEX_EMBED_BATCH,
_ => crate::embeddings::openai::OPENAI_EMBED_BATCH,
};
let started = std::time::Instant::now();
match self
.embed_all_batches(embedder.as_ref(), &leg.model, &inputs, dims, limit)
.await
{
Ok(out) => {
self.metrics.embedding(
&alias,
&leg.model,
&leg.provider,
started.elapsed().as_secs_f64(),
);
self.record_embed_usage(&ctx, &alias, leg, out.input_tokens);
return Ok(crate::embeddings::build_response(alias, out));
}
Err(e) => last_err = Some(e),
}
}
Err(last_err.unwrap_or(GatewayError::AllLegsFailed {
route: alias,
failures: Vec::new(),
}))
}
async fn embed_all_batches(
&self,
embedder: &dyn crate::embeddings::EmbeddingProvider,
model: &str,
inputs: &[String],
dims: u32,
limit: usize,
) -> Result<crate::embeddings::EmbedOut, GatewayError> {
let mut vectors = Vec::with_capacity(inputs.len());
let mut input_tokens = 0u64;
for batch in crate::embeddings::split_batches(inputs, limit) {
let out = embedder.embed(model, batch, dims).await?;
input_tokens += out.input_tokens;
vectors.extend(out.vectors);
}
Ok(crate::embeddings::EmbedOut {
vectors,
input_tokens,
})
}
fn record_embed_usage(&self, ctx: &RequestCtx, alias: &str, leg: &ChainLeg, input_tokens: u64) {
let cost = self.pricing.embedding_cost_usd(
&leg.provider,
&leg.model,
input_tokens,
self.embed_default_input_per_mtok,
);
let attr = self.attribution_of(ctx, alias);
self.ledger.enqueue(UsageEntry {
ts: Utc::now(),
tenant: self.tenant_of(ctx).to_string(),
workspace: attr.workspace,
user: attr.user,
thread: attr.thread,
message: attr.message,
route: alias.to_string(),
provider: leg.provider.clone(),
model: leg.model.clone(),
lane: "embedding".into(),
input_tokens,
output_tokens: 0,
cost_usd: cost,
request_id: ctx
.resolved_request_id()
.unwrap_or_else(|| Uuid::new_v4().to_string()),
status: "ok".into(),
op: "embedding".into(),
user_task_type: attr.user_task_type,
ai_task_type: attr.ai_task_type,
});
}
}
fn native_start_retryable(e: &GatewayError) -> bool {
match e {
GatewayError::Upstream { status, .. } => {
*status >= 500
|| *status == reqwest::StatusCode::TOO_MANY_REQUESTS.as_u16()
|| *status == reqwest::StatusCode::REQUEST_TIMEOUT.as_u16()
}
GatewayError::UpstreamTimeout => true,
_ => false,
}
}
fn require_jev_block(req: &ChatRequest, legs: &[ChainLeg]) -> Result<(), GatewayError> {
let needs_block = legs.iter().any(|l| l.provider == "typesafe");
let has_block = req.jev.as_ref().is_some_and(|j| !j.questions.is_empty());
match (needs_block, has_block) {
(true, false) => Err(GatewayError::BadRequest(format!(
"route '{}' uses provider 'typesafe'; send a 'jev' extension block with questions",
req.model
))),
_ => Ok(()),
}
}
fn non_typesafe_legs(legs: &[ChainLeg]) -> Vec<ChainLeg> {
legs.iter()
.filter(|l| l.provider != "typesafe")
.cloned()
.collect()
}
fn with_leg_thinking<'a>(req: &'a ChatRequest, leg: &ChainLeg) -> Cow<'a, ChatRequest> {
let client_set = req
.vertex
.as_ref()
.is_some_and(|v| v.thinking_config.is_some());
match (client_set, leg.effort.and_then(Effort::thinking_budget)) {
(false, Some(budget)) => Cow::Owned(ChatRequest {
vertex: Some(VertexExt {
thinking_config: Some(serde_json::json!({ "thinkingBudget": budget })),
..req.vertex.clone().unwrap_or_default()
}),
..req.clone()
}),
_ => Cow::Borrowed(req),
}
}
impl GatewayBuilder {
pub fn routes(mut self, routes: RouteTable) -> Self {
self.routes = Some(routes);
self
}
pub fn catalog(mut self, catalog: Catalog) -> Self {
self.catalog = Some(catalog);
self
}
pub fn pricing(mut self, pricing: PricingTable) -> Self {
self.pricing = Some(pricing);
self
}
pub fn ledger(mut self, ledger: LedgerHandle) -> Self {
self.ledger = Some(ledger);
self
}
pub fn vertex_native(mut self, v: Option<VertexNativeProvider>) -> Self {
self.vertex_native = v;
self
}
pub fn jev_native(mut self, v: Option<JevNativeProvider>) -> Self {
self.jev_native = v;
self
}
pub fn timeouts(mut self, t: StreamTimeouts) -> Self {
self.timeouts = Some(t);
self
}
pub fn default_tenant(mut self, t: impl Into<String>) -> Self {
self.default_tenant = Some(t.into());
self
}
pub fn embed_routes(mut self, t: crate::routing::embeddings::EmbeddingRouteTable) -> Self {
self.embed_routes = Some(t);
self
}
pub fn embedder(
mut self,
id: impl Into<String>,
e: Arc<dyn crate::embeddings::EmbeddingProvider>,
) -> Self {
self.embedders
.get_or_insert_with(Default::default)
.insert(id.into(), e);
self
}
pub fn embed_default_input_per_mtok(mut self, v: f64) -> Self {
self.embed_default_input_per_mtok = Some(v);
self
}
pub fn guard(mut self, guard: GuardEngine) -> Self {
self.guard = Some(guard);
self
}
pub fn ai_task_types(mut self, t: crate::ai_task_type::AiTaskTypeTable) -> Self {
self.ai_task_types = Some(t);
self
}
pub fn metrics(mut self, metrics: Arc<GatewayMetrics>) -> Self {
self.metrics = Some(metrics);
self
}
pub fn build(self) -> anyhow::Result<Gateway> {
let metrics = self.metrics.unwrap_or_else(GatewayMetrics::noop);
Ok(Gateway {
routes: Arc::new(
self.routes
.ok_or_else(|| anyhow::anyhow!("Gateway: routes required"))?,
),
catalog: Arc::new(
self.catalog
.ok_or_else(|| anyhow::anyhow!("Gateway: catalog required"))?
.tap(|c| c.attach_metrics(&metrics)),
),
pricing: Arc::new(
self.pricing
.ok_or_else(|| anyhow::anyhow!("Gateway: pricing required"))?,
),
ledger: self
.ledger
.ok_or_else(|| anyhow::anyhow!("Gateway: ledger required"))?,
vertex_native: self.vertex_native.map(Arc::new),
jev_native: self.jev_native.map(Arc::new),
timeouts: self.timeouts.unwrap_or_default(),
default_tenant: self.default_tenant.unwrap_or_else(|| "unattributed".into()),
embed_routes: Arc::new(self.embed_routes.unwrap_or_default()),
embedders: self.embedders.unwrap_or_default(),
embed_default_input_per_mtok: self.embed_default_input_per_mtok.unwrap_or(0.10),
guard: Arc::new(
self.guard
.unwrap_or_else(GuardEngine::empty)
.with_metrics(metrics.clone()),
),
ai_task_types: Arc::new(self.ai_task_types.unwrap_or_default()),
metrics,
})
}
}
pub(crate) struct StreamSideEffects {
ledger: LedgerHandle,
pricing: Arc<PricingTable>,
metrics: Arc<GatewayMetrics>,
route: String,
tenant: String,
attribution: Attribution,
provider: String,
model: String,
lane: &'static str, request_id: String,
legs_attempted: u32,
started: Instant,
input_tokens: u64,
output_tokens: u64,
status: &'static str,
fired: bool,
}
impl StreamSideEffects {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
ledger: LedgerHandle,
pricing: Arc<PricingTable>,
metrics: Arc<GatewayMetrics>,
route: String,
tenant: String,
attribution: Attribution,
provider: String,
model: String,
lane: &'static str,
request_id: String,
legs_attempted: u32,
started: Instant,
) -> Self {
Self {
ledger,
pricing,
metrics,
route,
tenant,
attribution,
provider,
model,
lane,
request_id,
legs_attempted,
started,
input_tokens: 0,
output_tokens: 0,
status: "ok",
fired: false,
}
}
pub(crate) fn observe(&mut self, item: &StreamItem) {
if let StreamItem::Done {
input_tokens,
output_tokens,
..
} = item
{
self.input_tokens = *input_tokens;
self.output_tokens = *output_tokens;
}
}
pub(crate) fn mark_error(&mut self) {
self.status = "error";
}
}
impl Drop for StreamSideEffects {
fn drop(&mut self) {
if self.fired {
return;
}
self.fired = true;
let cost = self.pricing.cost_usd(
&self.provider,
&self.model,
self.input_tokens,
self.output_tokens,
);
self.ledger.enqueue(UsageEntry {
ts: Utc::now(),
tenant: self.tenant.clone(),
workspace: self.attribution.workspace.clone(),
user: self.attribution.user.clone(),
thread: self.attribution.thread.clone(),
message: self.attribution.message.clone(),
route: self.route.clone(),
provider: self.provider.clone(),
model: self.model.clone(),
lane: self.lane.to_string(),
input_tokens: self.input_tokens,
output_tokens: self.output_tokens,
cost_usd: cost,
request_id: self.request_id.clone(),
status: self.status.to_string(),
op: "chat".into(),
user_task_type: self.attribution.user_task_type.clone(),
ai_task_type: self.attribution.ai_task_type.clone(),
});
let completion = Completion {
provider: self.provider.clone(),
model: self.model.clone(),
content: String::new(),
tool_calls: Vec::new(),
finish_reason: FinishReason::Stop,
input_tokens: self.input_tokens,
output_tokens: self.output_tokens,
};
let lane = if self.lane == "native" {
Lane::NativeVertex
} else {
Lane::Standard
};
GenAiSpan::from_completion(
&completion,
lane,
&self.route,
&self.tenant,
self.attribution.workspace.as_deref(),
self.legs_attempted,
true,
)
.emit_metrics(&self.metrics, self.started.elapsed().as_secs_f64());
}
}
pub struct GuardedStream {
inner: BoxStream<'static, Result<StreamItem, LegError>>,
model: String,
guard: StreamSideEffects,
routing: RoutingReport,
}
impl GuardedStream {
pub(crate) fn new(
inner: BoxStream<'static, Result<StreamItem, LegError>>,
model: String,
guard: StreamSideEffects,
routing: RoutingReport,
) -> Self {
Self {
inner,
model,
guard,
routing,
}
}
pub fn model(&self) -> &str {
&self.model
}
pub fn routing(&self) -> &RoutingReport {
&self.routing
}
}
impl Stream for GuardedStream {
type Item = Result<StreamItem, LegError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
match this.inner.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok(item))) => {
this.guard.observe(&item);
Poll::Ready(Some(Ok(item)))
}
Poll::Ready(Some(Err(e))) => {
this.guard.mark_error();
Poll::Ready(Some(Err(e)))
}
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
pub(crate) async fn collect_committed(
committed: CommittedStream,
) -> Result<Completion, GatewayError> {
use futures::TryStreamExt;
let CommittedStream {
provider,
model,
stream,
} = committed;
stream
.map_err(|e: LegError| GatewayError::Upstream {
status: 502,
body: e.to_string(),
})
.try_fold(Accumulator::default(), |mut acc, item| async move {
acc.push(item);
Ok(acc)
})
.await
.map(|acc| Completion {
provider,
model,
content: acc.content,
tool_calls: acc.tool_calls,
finish_reason: acc.finish_reason,
input_tokens: acc.input_tokens,
output_tokens: acc.output_tokens,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ledger::{InMemoryLedger, LedgerHandle, LedgerStore};
fn test_gateway() -> Gateway {
let routes = RouteTable::from_toml_str(
r#"[routes."fast"]
legs = [{ provider = "qwen", model = "qwen-max" }]"#,
)
.unwrap();
let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
let ledger = LedgerHandle::spawn(
Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
16,
);
Gateway::builder()
.routes(routes)
.catalog(catalog)
.pricing(PricingTable::default())
.ledger(ledger)
.default_tenant("acme")
.build()
.unwrap()
}
#[tokio::test]
async fn builder_builds_and_lists_aliases() {
let gw = test_gateway();
assert_eq!(gw.model_aliases(), vec!["fast".to_string()]);
assert_eq!(gw.default_tenant, "acme");
}
#[test]
fn builder_requires_components() {
assert!(Gateway::builder().build().is_err());
}
#[tokio::test]
async fn guard_records_usage_on_drop() {
let store = Arc::new(InMemoryLedger::default());
let ledger = LedgerHandle::spawn(store.clone(), 16);
let pricing = Arc::new(PricingTable::default());
{
let mut guard = StreamSideEffects::new(
ledger.clone(),
pricing,
GatewayMetrics::noop(),
"route".into(),
"tenant".into(),
Attribution::default(),
"p".into(),
"m".into(),
"standard",
"rid".into(),
1,
Instant::now(),
);
guard.observe(&crate::routing::stream::StreamItem::Done {
input_tokens: 3,
output_tokens: 2,
finish_reason: crate::routing::stream::FinishReason::Stop,
});
} tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let rows = store.entries();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].input_tokens, 3);
assert_eq!(rows[0].output_tokens, 2);
assert_eq!(rows[0].status, "ok");
}
#[tokio::test]
async fn guarded_stream_yields_items_and_records_on_drop() {
use crate::routing::stream::{FinishReason, StreamItem};
use futures::StreamExt;
let store = Arc::new(InMemoryLedger::default());
let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
let inner = futures::stream::iter(vec![
Ok(StreamItem::Delta("hi".into())),
Ok(StreamItem::Done {
input_tokens: 4,
output_tokens: 2,
finish_reason: FinishReason::Stop,
}),
])
.boxed();
let guard = StreamSideEffects::new(
ledger,
Arc::new(PricingTable::default()),
GatewayMetrics::noop(),
"route".into(),
"acme".into(),
Attribution::default(),
"p".into(),
"m".into(),
"standard",
"rid".into(),
1,
std::time::Instant::now(),
);
{
let mut gs = GuardedStream::new(inner, "m".into(), guard, RoutingReport::default());
let mut n = 0;
while let Some(item) = gs.next().await {
item.unwrap();
n += 1;
}
assert_eq!(n, 2);
assert_eq!(gs.model(), "m");
} tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let rows = store.entries();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].input_tokens, 4);
assert_eq!(rows[0].tenant, "acme");
}
fn resolver_gateway(ai_task_types: &str) -> Gateway {
let routes = RouteTable::from_toml_str(
"[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
)
.unwrap();
Gateway::builder()
.routes(routes)
.catalog(Catalog::for_test(vec![(
"qwen",
"http://127.0.0.1:1/v1".into(),
)]))
.pricing(PricingTable::default())
.ledger(LedgerHandle::spawn(
Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
4,
))
.ai_task_types(
crate::ai_task_type::AiTaskTypeTable::from_toml_str(ai_task_types).unwrap(),
)
.build()
.unwrap()
}
#[tokio::test]
async fn ai_task_type_prefers_the_caller_header_over_the_alias_mapping() {
let gw = resolver_gateway("conversation = [\"fast\"]");
let ctx = RequestCtx {
ai_task_type: Some("caller-supplied".into()),
..Default::default()
};
assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "caller-supplied");
}
#[tokio::test]
async fn ai_task_type_is_inferred_from_the_route_alias_when_no_header() {
let gw = resolver_gateway("conversation = [\"fast\", \"planning\"]");
let ctx = RequestCtx::default();
assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
assert_eq!(gw.ai_task_type_of(&ctx, "planning"), "conversation");
}
#[tokio::test]
async fn ai_task_type_defaults_to_simple_for_an_unmapped_alias() {
let gw = resolver_gateway("conversation = [\"fast\"]");
assert_eq!(
gw.ai_task_type_of(&RequestCtx::default(), "graph-llm"),
"simple"
);
}
#[tokio::test]
async fn an_empty_ai_task_type_header_falls_back_to_inference() {
let gw = resolver_gateway("conversation = [\"fast\"]");
let ctx = RequestCtx {
ai_task_type: Some(String::new()),
..Default::default()
};
assert_eq!(gw.ai_task_type_of(&ctx, "fast"), "conversation");
}
#[tokio::test]
async fn chat_returns_completion_and_records_ledger() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock = MockServer::start().await;
Mock::given(method("POST")).and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200)
.insert_header("content-type", "text/event-stream")
.set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\n\
data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2}}\n\n\
data: [DONE]\n\n"))
.mount(&mock).await;
let routes = RouteTable::from_toml_str(
"[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
)
.unwrap();
let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
let store = Arc::new(InMemoryLedger::default());
let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
let gw = Gateway::builder()
.routes(routes)
.catalog(catalog)
.pricing(PricingTable::default())
.ledger(ledger)
.default_tenant("def")
.ai_task_types(
crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
.unwrap(),
)
.build()
.unwrap();
let req = serde_json::from_value(serde_json::json!(
{"model":"fast","messages":[{"role":"user","content":"hi"}]}))
.unwrap();
let ctx = RequestCtx {
tenant: Some("acme".into()),
workspace: Some("ws-9".into()),
user: Some("user-42".into()),
thread: Some("thread-9".into()),
message: Some("msg-7".into()),
user_task_type: Some("summarisation".into()),
ai_task_type: None,
request_id: Some("corr-123".into()),
};
let c = match gw.chat(req, &ctx).await.unwrap() {
ChatOutcome::Plain(c) => c,
ChatOutcome::Hybrid(_) => panic!("expected plain completion"),
};
assert_eq!(c.content, "hi");
assert_eq!(c.input_tokens, 3);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let rows = store.entries();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].tenant, "acme");
assert_eq!(rows[0].user.as_deref(), Some("user-42"));
assert_eq!(rows[0].thread.as_deref(), Some("thread-9"));
assert_eq!(rows[0].message.as_deref(), Some("msg-7"));
assert_eq!(rows[0].user_task_type.as_deref(), Some("summarisation"));
assert_eq!(rows[0].ai_task_type, "conversation");
assert_eq!(rows[0].request_id, "corr-123");
}
#[tokio::test]
async fn chat_stream_yields_items_and_records() {
use crate::routing::stream::StreamItem;
use futures::StreamExt;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock = MockServer::start().await;
Mock::given(method("POST")).and(path("/v1/chat/completions"))
.respond_with(ResponseTemplate::new(200)
.insert_header("content-type", "text/event-stream")
.set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
data: [DONE]\n\n"))
.mount(&mock).await;
let routes = RouteTable::from_toml_str(
"[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
)
.unwrap();
let catalog = Catalog::for_test(vec![("qwen", format!("{}/v1", mock.uri()))]);
let store = Arc::new(InMemoryLedger::default());
let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
let gw = Gateway::builder()
.routes(routes)
.catalog(catalog)
.pricing(PricingTable::default())
.ledger(ledger)
.default_tenant("def")
.ai_task_types(
crate::ai_task_type::AiTaskTypeTable::from_toml_str("conversation = [\"fast\"]")
.unwrap(),
)
.build()
.unwrap();
let req = serde_json::from_value(serde_json::json!(
{"model":"fast","stream":true,"messages":[{"role":"user","content":"hi"}]}))
.unwrap();
let ctx = RequestCtx {
user_task_type: Some("code-review".into()),
..Default::default()
};
let mut stream = gw.chat_stream(req, &ctx).await.unwrap();
let mut got = false;
while let Some(i) = stream.next().await {
if matches!(i.unwrap(), StreamItem::Delta(ref t) if t == "go") {
got = true;
}
}
drop(stream);
assert!(got);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let rows = store.entries();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].user_task_type.as_deref(), Some("code-review"));
assert_eq!(rows[0].ai_task_type, "conversation");
}
use crate::embeddings::{EmbedOut, EmbeddingInput, EmbeddingProvider, EmbeddingRequest};
use async_trait::async_trait;
struct FlakyEmbedder;
#[async_trait]
impl EmbeddingProvider for FlakyEmbedder {
async fn embed(
&self,
_model: &str,
_inputs: &[String],
_dims: u32,
) -> Result<EmbedOut, GatewayError> {
Err(GatewayError::Upstream {
status: 500,
body: "x".into(),
})
}
}
struct GoodEmbedder;
#[async_trait]
impl EmbeddingProvider for GoodEmbedder {
async fn embed(
&self,
_model: &str,
inputs: &[String],
dims: u32,
) -> Result<EmbedOut, GatewayError> {
Ok(EmbedOut {
vectors: inputs.iter().map(|_| vec![0.0f32; dims as usize]).collect(),
input_tokens: 6,
})
}
}
fn embed_gateway() -> (Gateway, Arc<InMemoryLedger>) {
let embed_routes = crate::routing::embeddings::EmbeddingRouteTable::from_toml_str(
r#"
[embeddings."default-embed"]
dimensions = 4
legs = [
{ provider = "flaky", model = "flaky-embed" },
{ provider = "good", model = "good-embed" },
]
"#,
)
.unwrap();
let routes = RouteTable::from_toml_str(
"[routes.\"fast\"]\nlegs = [{ provider = \"qwen\", model = \"qwen-max\" }]",
)
.unwrap();
let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
let store = Arc::new(InMemoryLedger::default());
let ledger = LedgerHandle::spawn(store.clone() as Arc<dyn LedgerStore>, 16);
let gw = Gateway::builder()
.routes(routes)
.catalog(catalog)
.pricing(PricingTable::default())
.ledger(ledger)
.default_tenant("acme")
.embed_routes(embed_routes)
.embedder(
"flaky",
Arc::new(FlakyEmbedder) as Arc<dyn EmbeddingProvider>,
)
.embedder("good", Arc::new(GoodEmbedder) as Arc<dyn EmbeddingProvider>)
.build()
.unwrap();
(gw, store)
}
#[tokio::test]
async fn embed_falls_through_to_good_leg_and_records_usage() {
let (gw, store) = embed_gateway();
let req = EmbeddingRequest {
input: EmbeddingInput::Many(vec!["a".into(), "b".into()]),
model: "default-embed".into(),
dimensions: None,
};
let ctx = RequestCtx {
user_task_type: Some("retrieval".into()),
..Default::default()
};
let resp = gw.embed(req, ctx).await.unwrap();
assert_eq!(resp.data.len(), 2);
assert!(resp.data.iter().all(|d| d.embedding.len() == 4));
assert_eq!(resp.data[0].index, 0);
assert_eq!(resp.data[1].index, 1);
assert_eq!(resp.usage.prompt_tokens, 6);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let rows = store.entries();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].op, "embedding");
assert_eq!(rows[0].lane, "embedding");
assert_eq!(rows[0].output_tokens, 0);
assert_eq!(rows[0].provider, "good");
assert!(rows[0].cost_usd > 0.0);
assert_eq!(rows[0].user_task_type.as_deref(), Some("retrieval"));
assert_eq!(rows[0].ai_task_type, "simple");
}
#[tokio::test]
async fn embed_dimension_mismatch_is_bad_request() {
let (gw, _store) = embed_gateway();
let req = EmbeddingRequest {
input: EmbeddingInput::Many(vec!["a".into()]),
model: "default-embed".into(),
dimensions: Some(8),
};
let err = gw.embed(req, RequestCtx::default()).await.unwrap_err();
assert!(matches!(err, GatewayError::BadRequest(_)));
}
#[tokio::test]
async fn chat_blocks_when_route_policy_refuses() {
use crate::guard::{GuardEngine, GuardrailsConfig};
let routes = RouteTable::from_toml_str(
r#"[routes."fast"]
policy = "strict"
legs = [{ provider = "qwen", model = "qwen-max" }]"#,
)
.unwrap();
let guard = GuardEngine::from_config(
&GuardrailsConfig::from_toml_str(
r#"[guardrails.strict]
scanners = [{ type = "ban_substrings", substrings = ["forbidden"] }]"#,
)
.unwrap(),
)
.unwrap();
let catalog = Catalog::for_test(vec![("qwen", "http://127.0.0.1:1/v1".into())]);
let ledger = LedgerHandle::spawn(
Arc::new(InMemoryLedger::default()) as Arc<dyn LedgerStore>,
16,
);
let gw = Gateway::builder()
.routes(routes)
.catalog(catalog)
.pricing(PricingTable::default())
.ledger(ledger)
.guard(guard)
.build()
.unwrap();
let req = serde_json::from_value(serde_json::json!({
"model": "fast",
"messages": [{ "role": "user", "content": "this is forbidden" }]
}))
.unwrap();
let err = gw.chat(req, &RequestCtx::default()).await.unwrap_err();
assert!(matches!(err, GatewayError::ContentBlocked { .. }));
}
fn vertex_req(vertex: serde_json::Value) -> ChatRequest {
serde_json::from_value(serde_json::json!({
"model": "auto",
"messages": [{"role": "user", "content": "hi"}],
"vertex": vertex
}))
.unwrap()
}
fn vertex_leg(effort: Option<crate::routing::effort::Effort>) -> ChainLeg {
ChainLeg {
provider: "vertex".into(),
model: "gemini-2.5-pro".into(),
effort,
..Default::default()
}
}
#[test]
fn leg_effort_becomes_a_thinking_budget_and_keeps_the_vertex_block() {
let req = vertex_req(serde_json::json!({"response_schema": {"type": "object"}}));
let sent = with_leg_thinking(
&req,
&vertex_leg(Some(crate::routing::effort::Effort::High)),
);
let v = sent.vertex.as_ref().unwrap();
assert_eq!(
v.thinking_config,
Some(serde_json::json!({"thinkingBudget": 8192}))
);
assert_eq!(
v.response_schema,
Some(serde_json::json!({"type": "object"}))
);
}
#[test]
fn client_thinking_config_wins_over_leg_effort() {
let req = vertex_req(serde_json::json!({"thinking_config": {"thinkingLevel": "low"}}));
let sent = with_leg_thinking(&req, &vertex_leg(Some(crate::routing::effort::Effort::Max)));
assert!(matches!(sent, std::borrow::Cow::Borrowed(_)));
}
#[test]
fn effort_none_or_absent_leaves_the_request_untouched() {
let req = vertex_req(serde_json::json!({"response_schema": {"type": "object"}}));
assert!(matches!(
with_leg_thinking(
&req,
&vertex_leg(Some(crate::routing::effort::Effort::None))
),
std::borrow::Cow::Borrowed(_)
));
assert!(matches!(
with_leg_thinking(&req, &vertex_leg(None)),
std::borrow::Cow::Borrowed(_)
));
}
#[test]
fn null_thinking_config_or_missing_vertex_block_gets_the_leg_budget() {
let budget = Some(serde_json::json!({"thinkingBudget": 1024}));
let leg = vertex_leg(Some(crate::routing::effort::Effort::Low));
let null_config = vertex_req(serde_json::json!({"thinking_config": null}));
let no_vertex: ChatRequest = serde_json::from_value(serde_json::json!({
"model": "auto",
"messages": [{"role": "user", "content": "hi"}]
}))
.unwrap();
[null_config, no_vertex].iter().for_each(|req| {
assert_eq!(
with_leg_thinking(req, &leg)
.vertex
.as_ref()
.and_then(|v| v.thinking_config.clone()),
budget
)
});
}
}