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;
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,
}
#[derive(Debug, Clone, Default)]
pub struct RequestCtx {
pub tenant: Option<String>,
pub workspace: Option<String>,
pub request_id: Option<String>,
}
#[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>,
}
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 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(),
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(),
});
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 leg = legs
.iter()
.find(|l| l.provider == "vertex")
.ok_or_else(|| GatewayError::NativeFeatureUnsupported {
feature: "native-vertex".into(),
route: req.model.clone(),
})?;
let provider = self
.vertex_native
.as_ref()
.ok_or_else(|| GatewayError::BadRequest("native vertex lane not configured".into()))?;
provider
.stream_generate(&leg.model, req)
.await
.map(|stream| CommittedStream::single("vertex".into(), leg.model.clone(), stream))
}
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)?;
let request_id = ctx
.request_id
.clone()
.unwrap_or_else(|| Uuid::new_v4().to_string());
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", 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(),
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)?;
let request_id = ctx
.request_id
.clone()
.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 committed = self.native_committed(&req, &legs).await?;
(collect_committed(committed).await?, "native", 1)
}
};
self.record(
ctx,
&req.model,
lane_str,
&request_id,
&completion,
legs_n,
started,
);
Ok(completion)
}
}
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 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()),
})
}
}
pub(crate) struct StreamSideEffects {
ledger: LedgerHandle,
pricing: Arc<PricingTable>,
route: String,
tenant: String,
workspace: 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>,
provider: String,
model: String,
lane: &'static str,
request_id: String,
legs_attempted: u32,
started: Instant,
) -> Self {
Self {
ledger,
pricing,
route,
tenant,
workspace,
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(),
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(),
});
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,
"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,
"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()),
workspace: None,
request_id: Some("corr-123".into()),
};
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].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);
}
}