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 uuid::Uuid;
use crate::error::{GatewayError, LegFailure};
use crate::guard::GuardEngine;
use crate::ledger::{LedgerHandle, UsageEntry};
use crate::observability::GenAiSpan;
use crate::pricing::PricingTable;
use crate::providers::Catalog;
use crate::routing::classify::{classify, Lane};
use crate::routing::executor::{
execute_buffered_with_timeouts, CommittedStream, Completion, LegError, StreamTimeouts,
};
use crate::routing::request::ChatRequest;
use crate::routing::stream::{Accumulator, FinishReason, StreamItem};
use crate::routing::table::{ChainLeg, RouteTable};
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) 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>,
}
#[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 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())
}
}
#[derive(Default)]
pub struct GatewayBuilder {
routes: Option<RouteTable>,
catalog: Option<Catalog>,
pricing: Option<PricingTable>,
ledger: Option<LedgerHandle>,
vertex_native: Option<VertexNativeProvider>,
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>,
}
impl Gateway {
pub fn builder() -> GatewayBuilder {
GatewayBuilder::default()
}
pub fn model_aliases(&self) -> Vec<String> {
self.routes.aliases()
}
fn resolve_legs(&self, req: &ChatRequest) -> Result<Vec<ChainLeg>, GatewayError> {
self.routes
.legs(&req.model)
.ok_or_else(|| GatewayError::UnknownModel(req.model.clone()))
.map(<[ChainLeg]>::to_vec)
}
fn guard_input(&self, req: &ChatRequest) -> Result<(), GatewayError> {
let policy = self.routes.policy_of(&req.model).unwrap_or("default");
self.guard.guard(policy, req)
}
fn tenant_of<'a>(&'a self, ctx: &'a RequestCtx) -> &'a str {
ctx.tenant.as_deref().unwrap_or(&self.default_tenant)
}
#[allow(clippy::too_many_arguments)]
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);
self.ledger.enqueue(UsageEntry {
ts: Utc::now(),
tenant: self.tenant_of(ctx).to_string(),
workspace: ctx.workspace.clone(),
user: ctx.user.clone(),
thread: ctx.thread.clone(),
message: ctx.message.clone(),
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(),
});
let span_lane = if lane == "native" {
Lane::NativeVertex
} else {
Lane::Standard
};
GenAiSpan::from_completion(
c,
span_lane,
route,
self.tenant_of(ctx),
ctx.workspace.as_deref(),
legs,
false,
)
.emit_metrics(started.elapsed().as_secs_f64());
}
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, req, 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 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 legs = self.resolve_legs(&req)?;
self.guard_input(&req)?;
let request_id = ctx
.resolved_request_id()
.unwrap_or_else(|| Uuid::new_v4().to_string());
let vertex_leg_count = legs.iter().filter(|l| l.provider == "vertex").count() as u32;
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),
),
};
let model = committed.model.clone();
let guard = StreamSideEffects::new(
self.ledger.clone(),
self.pricing.clone(),
req.model.clone(),
self.tenant_of(ctx).to_string(),
ctx.workspace.clone(),
ctx.user.clone(),
ctx.thread.clone(),
ctx.message.clone(),
committed.provider.clone(),
committed.model.clone(),
lane_str,
request_id,
legs_attempted,
started,
);
Ok(GuardedStream::new(committed.stream, model, guard))
}
pub async fn chat(
&self,
req: ChatRequest,
ctx: &RequestCtx,
) -> Result<Completion, GatewayError> {
let started = Instant::now();
let legs = self.resolve_legs(&req)?;
self.guard_input(&req)?;
let request_id = ctx
.resolved_request_id()
.unwrap_or_else(|| Uuid::new_v4().to_string());
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),
)
}
};
self.record(
ctx,
&req.model,
lane_str,
&request_id,
&completion,
legs_n,
started,
);
Ok(completion)
}
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) => {
metrics::counter!(
"synapse_embeddings_total",
"route" => alias.clone(),
"model" => leg.model.clone(),
"provider" => leg.provider.clone(),
)
.increment(1);
metrics::histogram!(
"synapse_embedding_duration_seconds",
"route" => alias.clone(),
"model" => leg.model.clone(),
"provider" => leg.provider.clone(),
)
.record(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,
);
self.ledger.enqueue(UsageEntry {
ts: Utc::now(),
tenant: self.tenant_of(ctx).to_string(),
workspace: ctx.workspace.clone(),
user: ctx.user.clone(),
thread: ctx.thread.clone(),
message: ctx.message.clone(),
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(),
});
}
}
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,
}
}
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 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 build(self) -> anyhow::Result<Gateway> {
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"))?,
),
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),
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)),
})
}
}
pub(crate) struct StreamSideEffects {
ledger: LedgerHandle,
pricing: Arc<PricingTable>,
route: String,
tenant: String,
workspace: Option<String>,
user: Option<String>,
thread: Option<String>,
message: Option<String>,
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>,
route: String,
tenant: String,
workspace: Option<String>,
user: Option<String>,
thread: Option<String>,
message: Option<String>,
provider: String,
model: String,
lane: &'static str,
request_id: String,
legs_attempted: u32,
started: Instant,
) -> Self {
Self {
ledger,
pricing,
route,
tenant,
workspace,
user,
thread,
message,
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.workspace.clone(),
user: self.user.clone(),
thread: self.thread.clone(),
message: self.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(),
});
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.workspace.as_deref(),
self.legs_attempted,
true,
)
.emit_metrics(self.started.elapsed().as_secs_f64());
}
}
pub struct GuardedStream {
inner: BoxStream<'static, Result<StreamItem, LegError>>,
model: String,
guard: StreamSideEffects,
}
impl GuardedStream {
pub(crate) fn new(
inner: BoxStream<'static, Result<StreamItem, LegError>>,
model: String,
guard: StreamSideEffects,
) -> Self {
Self {
inner,
model,
guard,
}
}
pub fn model(&self) -> &str {
&self.model
}
}
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,
}
}
}
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,
"route".into(),
"tenant".into(),
None,
None,
None,
None,
"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()),
"route".into(),
"acme".into(),
None,
None,
None,
None,
"p".into(),
"m".into(),
"standard",
"rid".into(),
1,
std::time::Instant::now(),
);
{
let mut gs = GuardedStream::new(inner, "m".into(), guard);
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");
}
#[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")
.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()),
user: Some("user-42".into()),
thread: Some("thread-9".into()),
message: Some("msg-7".into()),
request_id: Some("corr-123".into()),
..Default::default()
};
let c = gw.chat(req, &ctx).await.unwrap();
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].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")
.build()
.unwrap();
let req = serde_json::from_value(serde_json::json!(
{"model":"fast","stream":true,"messages":[{"role":"user","content":"hi"}]}))
.unwrap();
let mut stream = gw.chat_stream(req, &RequestCtx::default()).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;
assert_eq!(store.entries().len(), 1);
}
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 resp = gw.embed(req, RequestCtx::default()).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);
}
#[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 { .. }));
}
}