use std::sync::{Arc, OnceLock};
use async_trait::async_trait;
use systemprompt_models::profile::GatewayRoute;
use systemprompt_models::wire::canonical::CanonicalRequest;
#[derive(Debug, thiserror::Error)]
pub enum RouteSelectorError {
#[error("route selector '{name}' failed: {message}")]
Failed { name: &'static str, message: String },
}
#[async_trait]
pub trait RouteSelector: Send + Sync {
fn name(&self) -> &'static str;
async fn refine(
&self,
matched: &GatewayRoute,
request: &CanonicalRequest,
) -> Result<Option<GatewayRoute>, RouteSelectorError>;
}
#[derive(Debug, Clone, Copy)]
pub struct RouteSelectorRegistration {
pub name: &'static str,
pub factory: fn() -> Arc<dyn RouteSelector>,
}
inventory::collect!(RouteSelectorRegistration);
#[macro_export]
macro_rules! register_route_selector {
($factory:expr, name = $name:expr $(,)?) => {
::inventory::submit! {
$crate::RouteSelectorRegistration {
name: $name,
factory: || ::std::sync::Arc::new($factory()),
}
}
};
}
pub struct RouteSelectorEngine {
selectors: Vec<Arc<dyn RouteSelector>>,
}
impl std::fmt::Debug for RouteSelectorEngine {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RouteSelectorEngine")
.field("selectors", &self.selectors.len())
.finish()
}
}
impl RouteSelectorEngine {
#[must_use]
pub fn global() -> &'static Self {
static ENGINE: OnceLock<RouteSelectorEngine> = OnceLock::new();
ENGINE.get_or_init(|| Self {
selectors: inventory::iter::<RouteSelectorRegistration>()
.map(|reg| (reg.factory)())
.collect(),
})
}
#[must_use]
pub fn has_selectors(&self) -> bool {
!self.selectors.is_empty()
}
pub async fn refine(
&self,
matched: &GatewayRoute,
request: &CanonicalRequest,
) -> Option<(GatewayRoute, &'static str)> {
for selector in &self.selectors {
match selector.refine(matched, request).await {
Ok(Some(route)) => {
tracing::info!(
selector = selector.name(),
from_route = %matched.id,
to_route = %route.id,
"gateway route refined by selector"
);
return Some((route, selector.name()));
},
Ok(None) => {},
Err(e) => {
tracing::warn!(
selector = selector.name(),
error = %e,
"route selector errored; keeping matched route"
);
},
}
}
None
}
}