use std::sync::{Arc, LazyLock};
use systemprompt_identifiers::{ModelId, ProviderId, RouteId, UserId};
use systemprompt_traits::DatabaseHandle;
#[derive(Debug, Clone)]
pub struct GatewayGuardRequest<'a> {
pub user_id: &'a UserId,
pub model: &'a ModelId,
pub route_id: Option<&'a RouteId>,
pub provider: &'a ProviderId,
pub streaming: bool,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub enum GatewayDenyKind {
#[default]
Quota,
Forbidden,
Unavailable,
}
#[derive(Debug, Clone)]
pub struct GatewayDenyReason {
pub message: String,
pub retry_after_seconds: i32,
pub kind: GatewayDenyKind,
}
impl GatewayDenyReason {
pub fn unavailable(message: impl Into<String>) -> Self {
Self {
message: message.into(),
retry_after_seconds: 5,
kind: GatewayDenyKind::Unavailable,
}
}
pub fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
retry_after_seconds: 0,
kind: GatewayDenyKind::Quota,
}
}
pub fn forbidden(message: impl Into<String>) -> Self {
Self {
message: message.into(),
retry_after_seconds: 0,
kind: GatewayDenyKind::Forbidden,
}
}
}
#[async_trait::async_trait]
pub trait GatewayRequestGuard: Send + Sync {
async fn check(
&self,
db: &dyn DatabaseHandle,
request: &GatewayGuardRequest<'_>,
) -> Result<(), GatewayDenyReason>;
}
#[derive(Debug, Clone, Copy)]
pub struct GatewayRequestGuardRegistration {
pub factory: fn() -> Arc<dyn GatewayRequestGuard>,
}
inventory::collect!(GatewayRequestGuardRegistration);
#[macro_export]
macro_rules! register_gateway_guard {
($guard_type:ty) => {
::inventory::submit! {
$crate::GatewayRequestGuardRegistration {
factory: || ::std::sync::Arc::new(<$guard_type>::default())
as ::std::sync::Arc<dyn $crate::GatewayRequestGuard>,
}
}
};
($guard_expr:expr) => {
::inventory::submit! {
$crate::GatewayRequestGuardRegistration {
factory: || ::std::sync::Arc::new($guard_expr)
as ::std::sync::Arc<dyn $crate::GatewayRequestGuard>,
}
}
};
}
static GATEWAY_GUARDS: LazyLock<Vec<Arc<dyn GatewayRequestGuard>>> = LazyLock::new(|| {
inventory::iter::<GatewayRequestGuardRegistration>
.into_iter()
.map(|registration| (registration.factory)())
.collect()
});
#[must_use]
pub fn gateway_guards() -> &'static [Arc<dyn GatewayRequestGuard>] {
&GATEWAY_GUARDS
}
pub async fn run_gateway_guards(
db: &dyn DatabaseHandle,
request: &GatewayGuardRequest<'_>,
) -> Result<(), GatewayDenyReason> {
for guard in gateway_guards() {
guard.check(db, request).await?;
}
Ok(())
}