use std::convert::Infallible;
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use axum::extract::{Request, State};
use axum::http::{HeaderMap, HeaderValue, StatusCode};
use axum::middleware::Next;
use axum::response::sse::{Event, Sse};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post};
use axum::Json;
use futures::StreamExt;
use serde_json::{json, Value};
use std::sync::atomic::{AtomicU64, Ordering};
use tower_http::cors::CorsLayer;
use tokio::sync::mpsc;
use crate::error::ShimError;
use crate::gateway::{
ChunkStream, Dispatch, DispatchError, GatewayConfig, GatewayError, GatewayRequest, Scheduler,
StreamChunk,
};
use crate::log::{Logger, RequestTimer};
use crate::models;
use crate::proxy::convert::{chunk_to_events, request_to_value, value_to_response};
use crate::proxy::error::ApiError;
use crate::proxy::ratelimit::{build_limiter, estimate_request_tokens, penalty_duration};
use crate::proxy::types::{ChatRequest, HealthResponse, ModelEntry, ModelsResponse, StreamEvent};
use crate::router::Router;
pub struct RealDispatch {
router: Arc<Router>,
logger: Option<Logger>,
}
impl RealDispatch {
fn map_err(err: ShimError) -> DispatchError {
match err {
ShimError::ProviderError { status: 429, body } => DispatchError {
message: body,
retry_after: Some(penalty_duration()),
},
other => DispatchError::new(other.to_string()),
}
}
}
#[async_trait]
impl Dispatch for RealDispatch {
async fn dispatch(&self, _provider: &str, payload: Value) -> Result<Value, DispatchError> {
crate::completion_with_logger(self.router.as_ref(), &payload, self.logger.as_ref())
.await
.map_err(Self::map_err)
}
async fn dispatch_stream(
&self,
_provider: &str,
payload: Value,
) -> Result<ChunkStream, DispatchError> {
let upstream = crate::stream(self.router.as_ref(), &payload)
.await
.map_err(Self::map_err)?;
let mapped = upstream.map(|item| item.map_err(|e| GatewayError::Upstream(e.to_string())));
Ok(Box::pin(mapped))
}
}
enum Backend {
Local(Arc<Scheduler>),
#[cfg(feature = "redis-coordination")]
Distributed(Arc<crate::gateway::distributed::DistributedGateway>),
}
pub struct GatewayState {
router: Arc<Router>,
backend: Backend,
keystore: crate::gateway::auth::KeyStore,
quota: crate::gateway::quota::TenantQuota,
idempotency: crate::gateway::idempotency::IdempotencyCache,
#[cfg_attr(not(feature = "redis-coordination"), allow(dead_code))]
idempotency_ttl_secs: u64,
overloaded_retry_after: Duration,
}
impl GatewayState {
pub fn from_env(router: Router, logger: Option<Logger>) -> Arc<Self> {
let config = GatewayConfig::from_env();
let router = Arc::new(router);
let dispatch = Arc::new(RealDispatch {
router: router.clone(),
logger,
});
let scheduler = Scheduler::new(config.clone(), build_limiter(), dispatch);
Arc::new(Self {
router,
backend: Backend::Local(scheduler),
keystore: crate::gateway::auth::KeyStore::from_env(),
quota: crate::gateway::quota::TenantQuota::new(),
idempotency: crate::gateway::idempotency::IdempotencyCache::new(
std::time::Duration::from_secs(idem_ttl_secs()),
),
idempotency_ttl_secs: idem_ttl_secs(),
overloaded_retry_after: config.overloaded_retry_after,
})
}
#[cfg(feature = "redis-coordination")]
pub async fn distributed_from_env(
router: Router,
logger: Option<Logger>,
redis_url: &str,
) -> Result<Arc<Self>, String> {
let config = GatewayConfig::from_env();
let router = Arc::new(router);
let dispatch = Arc::new(RealDispatch {
router: router.clone(),
logger,
});
let gateway = crate::gateway::distributed::DistributedGateway::connect(
redis_url,
dispatch,
build_limiter(),
config.clone(),
)
.await
.map_err(|e| e.to_string())?;
let providers: Vec<String> = router
.provider_keys()
.into_iter()
.map(String::from)
.collect();
gateway.spawn_workers(providers);
Ok(Arc::new(Self {
router,
backend: Backend::Distributed(gateway),
keystore: crate::gateway::auth::KeyStore::from_env(),
quota: crate::gateway::quota::TenantQuota::new(),
idempotency: crate::gateway::idempotency::IdempotencyCache::new(
std::time::Duration::from_secs(idem_ttl_secs()),
),
idempotency_ttl_secs: idem_ttl_secs(),
overloaded_retry_after: config.overloaded_retry_after,
}))
}
async fn submit(&self, req: GatewayRequest) -> Result<Value, GatewayError> {
match &self.backend {
Backend::Local(scheduler) => scheduler.submit(req).await,
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => gateway.submit(req).await,
}
}
async fn submit_stream(
&self,
req: GatewayRequest,
) -> Result<mpsc::Receiver<StreamChunk>, GatewayError> {
match &self.backend {
Backend::Local(scheduler) => scheduler.submit_stream(req).await,
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => gateway.submit_stream(req).await,
}
}
async fn update_queue_depth_gauges(&self) {
use crate::gateway::metrics;
let depths = match &self.backend {
Backend::Local(scheduler) => scheduler.lane_depths(),
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => {
let providers: Vec<String> = self
.router
.provider_keys()
.into_iter()
.map(String::from)
.collect();
gateway.queue_depths(&providers).await
}
};
for (provider, depth) in depths {
metrics::gauge_set(
metrics::QUEUE_DEPTH,
&[("provider", &provider)],
depth as i64,
);
}
}
async fn idem_lookup(&self, key: &str) -> Option<Value> {
match &self.backend {
Backend::Local(_) => self.idempotency.get(key),
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => gateway.idem_get(key).await,
}
}
async fn idem_store(&self, key: &str, value: &Value) {
match &self.backend {
Backend::Local(_) => self.idempotency.put(key, value.clone()),
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => {
gateway
.idem_put(key, value, self.idempotency_ttl_secs)
.await
}
}
}
async fn ready(&self) -> bool {
match &self.backend {
Backend::Local(_) => true,
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => gateway.ping().await,
}
}
async fn stats(&self) -> Value {
let providers: Vec<String> = self
.router
.provider_keys()
.into_iter()
.map(String::from)
.collect();
let (mode, depths) = match &self.backend {
Backend::Local(scheduler) => ("in-memory", scheduler.lane_depths()),
#[cfg(feature = "redis-coordination")]
Backend::Distributed(gateway) => {
("distributed", gateway.queue_depths(&providers).await)
}
};
let mut lanes = Vec::new();
for (provider, depth) in depths {
#[cfg_attr(not(feature = "redis-coordination"), allow(unused_mut))]
let mut lane = json!({ "provider": provider, "queue_depth": depth });
#[cfg(feature = "redis-coordination")]
if let Backend::Distributed(gateway) = &self.backend {
lane["dead_letter"] = json!(gateway.dead_letter_len(&provider).await);
}
lanes.push(lane);
}
json!({
"mode": mode,
"auth_enforced": self.keystore.is_enforced(),
"providers": providers,
"lanes": lanes,
})
}
fn enforce_quota(
&self,
identity: &crate::gateway::auth::Identity,
provider: &str,
permits: u32,
) -> Result<(), ApiError> {
if let Err(retry) = self.quota.check(
&identity.tenant,
provider,
identity.rpm,
identity.tpm,
permits,
) {
crate::gateway::metrics::incr(
crate::gateway::metrics::REJECTED,
&[("provider", provider), ("reason", "tenant_quota")],
);
return Err(ApiError::RateLimited(retry));
}
Ok(())
}
}
fn idem_ttl_secs() -> u64 {
std::env::var("LLMSHIM_GATEWAY_IDEMPOTENCY_TTL_SECS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(300)
}
fn build_request(
state: &GatewayState,
headers: &HeaderMap,
req: &ChatRequest,
) -> Result<(String, GatewayRequest, crate::gateway::auth::Identity), ApiError> {
let identity = state.keystore.identify(headers).map_err(|_| {
crate::gateway::metrics::incr(
crate::gateway::metrics::REJECTED,
&[("provider", "unknown"), ("reason", "unauthorized")],
);
ApiError::Unauthorized
})?;
let provider_name = {
let (provider, _model) = state.router.resolve(&req.model)?;
provider.name().to_string()
};
let gw = GatewayRequest {
provider: provider_name.clone(),
tier: identity.tier,
permits: estimate_request_tokens(req),
payload: request_to_value(req),
};
Ok((provider_name, gw, identity))
}
fn gateway_err_to_api(state: &GatewayState, err: GatewayError) -> ApiError {
match err {
GatewayError::Overloaded(retry_after) => ApiError::Overloaded(retry_after),
GatewayError::Timeout => ApiError::Overloaded(state.overloaded_retry_after),
GatewayError::Shutdown => ApiError::from(ShimError::ProviderError {
status: 503,
body: "gateway shutting down".to_string(),
}),
GatewayError::Upstream(message) => ApiError::from(ShimError::ProviderError {
status: 502,
body: message,
}),
}
}
async fn chat(
State(state): State<Arc<GatewayState>>,
headers: HeaderMap,
Json(req): Json<ChatRequest>,
) -> Result<Response, ApiError> {
if req.stream {
return Ok(chat_stream_inner(state, headers, req).await);
}
let idem_key = headers
.get("idempotency-key")
.and_then(|v| v.to_str().ok())
.map(str::to_string);
let (provider_name, gw, identity) = build_request(&state, &headers, &req)?;
state.enforce_quota(&identity, &provider_name, gw.permits)?;
if let Some(key) = &idem_key {
if let Some(cached) = state.idem_lookup(key).await {
return Ok((
[("idempotency-replayed", "true")],
Json(value_to_response(&cached, &provider_name, 0)),
)
.into_response());
}
}
let timer = RequestTimer::start();
match state.submit(gw).await {
Ok(resp) => {
if let Some(key) = &idem_key {
state.idem_store(key, &resp).await;
}
let elapsed = timer.elapsed().as_millis() as u64;
Ok(Json(value_to_response(&resp, &provider_name, elapsed)).into_response())
}
Err(err) => Err(gateway_err_to_api(&state, err)),
}
}
async fn chat_stream(
State(state): State<Arc<GatewayState>>,
headers: HeaderMap,
Json(req): Json<ChatRequest>,
) -> Response {
chat_stream_inner(state, headers, req).await
}
async fn chat_stream_inner(
state: Arc<GatewayState>,
headers: HeaderMap,
req: ChatRequest,
) -> Response {
let (provider_name, gw, identity) = match build_request(&state, &headers, &req) {
Ok(t) => t,
Err(e) => return e.into_response(),
};
if let Err(e) = state.enforce_quota(&identity, &provider_name, gw.permits) {
return e.into_response();
}
let mut rx = match state.submit_stream(gw).await {
Ok(rx) => rx,
Err(err) => return gateway_err_to_api(&state, err).into_response(),
};
let event_stream = async_stream::stream! {
while let Some(item) = rx.recv().await {
match item {
Ok(chunk) => {
for event in chunk_to_events(&chunk) {
let event_type = stream_event_type(&event);
if let Ok(data) = serde_json::to_string(&event) {
yield Ok(Event::default().event(event_type).data(data));
}
}
}
Err(e) => {
let error_event = StreamEvent::Error { message: e.to_string() };
if let Ok(data) = serde_json::to_string(&error_event) {
yield Ok(Event::default().event("error").data(data));
}
break;
}
}
}
};
fn pin_item<S: futures::Stream<Item = Result<Event, Infallible>>>(s: S) -> S {
s
}
Sse::new(pin_item(event_stream)).into_response()
}
fn stream_event_type(event: &StreamEvent) -> &'static str {
match event {
StreamEvent::Content { .. } => "content",
StreamEvent::Reasoning { .. } => "reasoning",
StreamEvent::ToolCall { .. } => "tool_call",
StreamEvent::Usage(_) => "usage",
StreamEvent::Done { .. } => "done",
StreamEvent::Error { .. } => "error",
}
}
async fn list_models(State(state): State<Arc<GatewayState>>) -> Json<ModelsResponse> {
let provider_keys = state.router.provider_keys();
let entries = models::available_models(&provider_keys)
.into_iter()
.map(|m| ModelEntry {
id: m.id.to_string(),
provider: m.provider.to_string(),
name: m.name.to_string(),
})
.collect();
Json(ModelsResponse { models: entries })
}
async fn metrics(State(state): State<Arc<GatewayState>>) -> Response {
state.update_queue_depth_gauges().await;
let body = crate::gateway::metrics::render();
(
[(
axum::http::header::CONTENT_TYPE,
"text/plain; version=0.0.4; charset=utf-8",
)],
body,
)
.into_response()
}
async fn ready(State(state): State<Arc<GatewayState>>) -> Response {
if state.ready().await {
(StatusCode::OK, "ready\n").into_response()
} else {
(StatusCode::SERVICE_UNAVAILABLE, "not ready\n").into_response()
}
}
async fn stats(State(state): State<Arc<GatewayState>>) -> Json<Value> {
Json(state.stats().await)
}
async fn request_id(mut req: Request, next: Next) -> Response {
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = req
.headers()
.get("x-request-id")
.and_then(|v| v.to_str().ok())
.map(str::to_string)
.unwrap_or_else(|| {
format!(
"gw-{}-{}",
std::process::id(),
COUNTER.fetch_add(1, Ordering::Relaxed)
)
});
if let Ok(hv) = HeaderValue::from_str(&id) {
req.headers_mut().insert("x-request-id", hv.clone());
let mut resp = next.run(req).await;
resp.headers_mut().insert("x-request-id", hv);
return resp;
}
next.run(req).await
}
async fn health(State(state): State<Arc<GatewayState>>) -> Json<HealthResponse> {
Json(HealthResponse {
status: "ok".to_string(),
providers: state
.router
.provider_keys()
.into_iter()
.map(String::from)
.collect(),
})
}
pub fn app(state: Arc<GatewayState>) -> axum::Router {
axum::Router::new()
.route("/v1/chat", post(chat))
.route("/v1/chat/stream", post(chat_stream))
.route("/v1/models", get(list_models))
.route("/v1/gateway/stats", get(stats))
.route("/metrics", get(metrics))
.route("/health", get(health))
.route("/ready", get(ready))
.layer(axum::middleware::from_fn(request_id))
.layer(CorsLayer::permissive())
.with_state(state)
}