use std::sync::{Arc, Mutex};
use runifold_core::RetrySafety;
use thiserror::Error;
use crate::circuit::{BreakerPermit, BreakerState, RoutePermit, SharedBreakerState};
use crate::{
CircuitBreakerConfig, Model, ModelCallContext, ModelCapabilities, ModelError, ModelErrorKind,
ModelEventStream, ModelFuture, ModelRef, ModelRequest, ModelRetryPolicy, ModelRouteHealth,
ModelStreamEvent, ProviderEvent, RouterClock, RouterSleeper, SystemRouterClock,
SystemRouterSleeper,
};
mod capabilities;
mod execution;
use capabilities::intersect_capabilities;
use execution::{RoutingRuntime, routed_stream};
#[derive(Clone)]
pub struct ModelRoute {
name: String,
model: Arc<dyn Model>,
target: ModelRef,
health: SharedBreakerState,
}
impl ModelRoute {
pub fn new(name: impl Into<String>, model: Arc<dyn Model>, target: ModelRef) -> Self {
Self {
name: name.into(),
model,
target,
health: Arc::new(Mutex::new(BreakerState::default())),
}
}
pub fn name(&self) -> &str {
&self.name
}
pub const fn target(&self) -> &ModelRef {
&self.target
}
}
impl std::fmt::Debug for ModelRoute {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ModelRoute")
.field("name", &self.name)
.field("target", &self.target)
.finish_non_exhaustive()
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct ModelFallbackPolicy {
unknown_safety_kinds: Vec<ModelErrorKind>,
}
impl ModelFallbackPolicy {
pub const fn safe_only() -> Self {
Self {
unknown_safety_kinds: Vec::new(),
}
}
#[must_use]
pub fn allow_unknown(mut self, kind: ModelErrorKind) -> Self {
if !self.unknown_safety_kinds.contains(&kind) {
self.unknown_safety_kinds.push(kind);
}
self
}
fn permits(&self, error: &ModelError) -> bool {
if error.kind == ModelErrorKind::Cancelled {
return false;
}
match error.retry_safety {
RetrySafety::Safe => true,
RetrySafety::Unknown => self.unknown_safety_kinds.contains(&error.kind),
_ => false,
}
}
}
#[derive(Clone, Debug, Error, Eq, PartialEq)]
#[non_exhaustive]
pub enum ModelRouterBuildError {
#[error("logical model provider and name cannot be empty")]
EmptyLogicalModel,
#[error("model route name cannot be empty")]
EmptyRouteName,
#[error("physical model provider and name cannot be empty")]
EmptyTarget,
#[error("model route `{0}` is already registered")]
DuplicateRoute(String),
#[error("model router requires at least one route")]
NoRoutes,
}
pub struct ModelRouterBuilder {
logical: ModelRef,
routes: Vec<ModelRoute>,
policy: ModelFallbackPolicy,
circuit_breaker: Option<CircuitBreakerConfig>,
clock: Arc<dyn RouterClock>,
retry_policy: Option<ModelRetryPolicy>,
sleeper: Arc<dyn RouterSleeper>,
error: Option<ModelRouterBuildError>,
}
impl ModelRouterBuilder {
#[must_use]
pub fn route(
mut self,
name: impl Into<String>,
model: Arc<dyn Model>,
target: ModelRef,
) -> Self {
if self.error.is_some() {
return self;
}
let route = ModelRoute::new(name, model, target);
if route.name.trim().is_empty() {
self.error = Some(ModelRouterBuildError::EmptyRouteName);
} else if route.target.provider.trim().is_empty() || route.target.name.trim().is_empty() {
self.error = Some(ModelRouterBuildError::EmptyTarget);
} else if self
.routes
.iter()
.any(|existing| existing.name == route.name)
{
self.error = Some(ModelRouterBuildError::DuplicateRoute(route.name));
} else {
self.routes.push(route);
}
self
}
#[must_use]
pub fn fallback_policy(mut self, policy: ModelFallbackPolicy) -> Self {
self.policy = policy;
self
}
#[must_use]
pub fn circuit_breaker(mut self, config: CircuitBreakerConfig) -> Self {
self.circuit_breaker = Some(config);
self
}
#[must_use]
pub fn clock(mut self, clock: Arc<dyn RouterClock>) -> Self {
self.clock = clock;
self
}
#[must_use]
pub fn retry_policy(mut self, policy: ModelRetryPolicy) -> Self {
self.retry_policy = Some(policy);
self
}
#[must_use]
pub fn sleeper(mut self, sleeper: Arc<dyn RouterSleeper>) -> Self {
self.sleeper = sleeper;
self
}
pub fn build(self) -> Result<ModelRouter, ModelRouterBuildError> {
if let Some(error) = self.error {
return Err(error);
}
if self.logical.provider.trim().is_empty() || self.logical.name.trim().is_empty() {
return Err(ModelRouterBuildError::EmptyLogicalModel);
}
if self.routes.is_empty() {
return Err(ModelRouterBuildError::NoRoutes);
}
Ok(ModelRouter {
logical: self.logical,
routes: self.routes,
policy: self.policy,
circuit_breaker: self.circuit_breaker,
clock: self.clock,
retry_policy: self.retry_policy,
sleeper: self.sleeper,
})
}
}
impl std::fmt::Debug for ModelRouterBuilder {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("ModelRouterBuilder")
.field("logical", &self.logical)
.field("routes", &self.routes)
.field("policy", &self.policy)
.field("circuit_breaker", &self.circuit_breaker)
.field("retry_policy", &self.retry_policy)
.field("error", &self.error)
.finish_non_exhaustive()
}
}
#[derive(Clone)]
pub struct ModelRouter {
logical: ModelRef,
routes: Vec<ModelRoute>,
policy: ModelFallbackPolicy,
circuit_breaker: Option<CircuitBreakerConfig>,
clock: Arc<dyn RouterClock>,
retry_policy: Option<ModelRetryPolicy>,
sleeper: Arc<dyn RouterSleeper>,
}
impl ModelRouter {
pub fn builder(logical: ModelRef) -> ModelRouterBuilder {
ModelRouterBuilder {
logical,
routes: Vec::new(),
policy: ModelFallbackPolicy::default(),
circuit_breaker: None,
clock: Arc::new(SystemRouterClock),
retry_policy: None,
sleeper: Arc::new(SystemRouterSleeper),
error: None,
}
}
pub const fn logical_model(&self) -> &ModelRef {
&self.logical
}
pub fn routes(&self) -> &[ModelRoute] {
&self.routes
}
pub fn route_health(&self) -> Vec<ModelRouteHealth> {
let now = self.clock.now();
self.routes
.iter()
.map(|route| {
crate::circuit::snapshot(
&route.health,
route.name.clone(),
route.target.clone(),
self.circuit_breaker.as_ref(),
now,
)
})
.collect()
}
fn validate_request(&self, request: &ModelRequest) -> Result<(), ModelError> {
if request.model != self.logical {
return Err(ModelError::local(
ModelErrorKind::InvalidRequest,
format!(
"router for `{}/{}` cannot invoke logical model `{}/{}`",
self.logical.provider,
self.logical.name,
request.model.provider,
request.model.name
),
));
}
Ok(())
}
}
impl Model for ModelRouter {
fn capabilities<'a>(
&'a self,
model: &'a ModelRef,
) -> ModelFuture<'a, Result<ModelCapabilities, ModelError>> {
Box::pin(async move {
if model != &self.logical {
return Err(ModelError::local(
ModelErrorKind::InvalidRequest,
"capabilities requested for the wrong logical model",
));
}
let mut capabilities = Vec::with_capacity(self.routes.len());
for route in &self.routes {
capabilities.push(route.model.capabilities(&route.target).await?);
}
Ok(intersect_capabilities(capabilities))
})
}
fn stream(
&self,
request: ModelRequest,
context: ModelCallContext,
) -> ModelFuture<'_, Result<ModelEventStream, ModelError>> {
let validation = self.validate_request(&request);
let stream = routed_stream(
self.routes.clone(),
self.policy.clone(),
RoutingRuntime {
circuit_breaker: self.circuit_breaker.clone(),
clock: self.clock.clone(),
retry_policy: self.retry_policy.clone(),
sleeper: self.sleeper.clone(),
},
request,
context,
);
Box::pin(async move {
validation?;
Ok(stream)
})
}
}
#[cfg(test)]
mod tests;