Skip to main content

runifold_model/
router.rs

1use std::sync::{Arc, Mutex};
2
3use runifold_core::RetrySafety;
4use thiserror::Error;
5
6use crate::circuit::{BreakerPermit, BreakerState, RoutePermit, SharedBreakerState};
7use crate::{
8    CircuitBreakerConfig, Model, ModelCallContext, ModelCapabilities, ModelError, ModelErrorKind,
9    ModelEventStream, ModelFuture, ModelRef, ModelRequest, ModelRetryPolicy, ModelRouteHealth,
10    ModelStreamEvent, ProviderEvent, RouterClock, RouterSleeper, SystemRouterClock,
11    SystemRouterSleeper,
12};
13
14mod capabilities;
15mod execution;
16
17use capabilities::intersect_capabilities;
18use execution::{RoutingRuntime, routed_stream};
19
20/// One named physical endpoint eligible for a logical model invocation.
21#[derive(Clone)]
22pub struct ModelRoute {
23    name: String,
24    model: Arc<dyn Model>,
25    target: ModelRef,
26    health: SharedBreakerState,
27}
28
29impl ModelRoute {
30    /// Creates a physical route.
31    pub fn new(name: impl Into<String>, model: Arc<dyn Model>, target: ModelRef) -> Self {
32        Self {
33            name: name.into(),
34            model,
35            target,
36            health: Arc::new(Mutex::new(BreakerState::default())),
37        }
38    }
39
40    /// Returns the stable route name.
41    pub fn name(&self) -> &str {
42        &self.name
43    }
44
45    /// Returns the provider-qualified physical target.
46    pub const fn target(&self) -> &ModelRef {
47        &self.target
48    }
49}
50
51impl std::fmt::Debug for ModelRoute {
52    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
53        formatter
54            .debug_struct("ModelRoute")
55            .field("name", &self.name)
56            .field("target", &self.target)
57            .finish_non_exhaustive()
58    }
59}
60
61/// Explicit authority for selecting another physical model after a failure.
62///
63/// Safe errors are always eligible. Errors with unknown retry safety are
64/// eligible only when their kind is explicitly added. Cancellation and errors
65/// marked unsafe are never eligible.
66#[derive(Clone, Debug, Default, Eq, PartialEq)]
67pub struct ModelFallbackPolicy {
68    unknown_safety_kinds: Vec<ModelErrorKind>,
69}
70
71impl ModelFallbackPolicy {
72    /// Creates the conservative policy: only errors explicitly marked safe.
73    pub const fn safe_only() -> Self {
74        Self {
75            unknown_safety_kinds: Vec::new(),
76        }
77    }
78
79    /// Allows fallback for one error kind whose retry safety is unknown.
80    ///
81    /// This is explicit authority to risk duplicate provider cost. It never
82    /// overrides cancellation or an error marked unsafe.
83    #[must_use]
84    pub fn allow_unknown(mut self, kind: ModelErrorKind) -> Self {
85        if !self.unknown_safety_kinds.contains(&kind) {
86            self.unknown_safety_kinds.push(kind);
87        }
88        self
89    }
90
91    fn permits(&self, error: &ModelError) -> bool {
92        if error.kind == ModelErrorKind::Cancelled {
93            return false;
94        }
95        match error.retry_safety {
96            RetrySafety::Safe => true,
97            RetrySafety::Unknown => self.unknown_safety_kinds.contains(&error.kind),
98            _ => false,
99        }
100    }
101}
102
103/// Invalid logical router configuration.
104#[derive(Clone, Debug, Error, Eq, PartialEq)]
105#[non_exhaustive]
106pub enum ModelRouterBuildError {
107    /// Logical provider or model name is blank.
108    #[error("logical model provider and name cannot be empty")]
109    EmptyLogicalModel,
110    /// A route name is blank.
111    #[error("model route name cannot be empty")]
112    EmptyRouteName,
113    /// A physical provider or model name is blank.
114    #[error("physical model provider and name cannot be empty")]
115    EmptyTarget,
116    /// A route name was registered more than once.
117    #[error("model route `{0}` is already registered")]
118    DuplicateRoute(String),
119    /// No physical route was registered.
120    #[error("model router requires at least one route")]
121    NoRoutes,
122}
123
124/// Fluent, validation-preserving assembly of a [`ModelRouter`].
125pub struct ModelRouterBuilder {
126    logical: ModelRef,
127    routes: Vec<ModelRoute>,
128    policy: ModelFallbackPolicy,
129    circuit_breaker: Option<CircuitBreakerConfig>,
130    clock: Arc<dyn RouterClock>,
131    retry_policy: Option<ModelRetryPolicy>,
132    sleeper: Arc<dyn RouterSleeper>,
133    error: Option<ModelRouterBuildError>,
134}
135
136impl ModelRouterBuilder {
137    /// Adds a physical route in selection order.
138    #[must_use]
139    pub fn route(
140        mut self,
141        name: impl Into<String>,
142        model: Arc<dyn Model>,
143        target: ModelRef,
144    ) -> Self {
145        if self.error.is_some() {
146            return self;
147        }
148        let route = ModelRoute::new(name, model, target);
149        if route.name.trim().is_empty() {
150            self.error = Some(ModelRouterBuildError::EmptyRouteName);
151        } else if route.target.provider.trim().is_empty() || route.target.name.trim().is_empty() {
152            self.error = Some(ModelRouterBuildError::EmptyTarget);
153        } else if self
154            .routes
155            .iter()
156            .any(|existing| existing.name == route.name)
157        {
158            self.error = Some(ModelRouterBuildError::DuplicateRoute(route.name));
159        } else {
160            self.routes.push(route);
161        }
162        self
163    }
164
165    /// Sets fallback authority.
166    #[must_use]
167    pub fn fallback_policy(mut self, policy: ModelFallbackPolicy) -> Self {
168        self.policy = policy;
169        self
170    }
171
172    /// Enables an independent circuit breaker for every physical route.
173    #[must_use]
174    pub fn circuit_breaker(mut self, config: CircuitBreakerConfig) -> Self {
175        self.circuit_breaker = Some(config);
176        self
177    }
178
179    /// Replaces the monotonic clock used by circuit-breaker policy.
180    #[must_use]
181    pub fn clock(mut self, clock: Arc<dyn RouterClock>) -> Self {
182        self.clock = clock;
183        self
184    }
185
186    /// Enables bounded same-route retries before fallback selection.
187    #[must_use]
188    pub fn retry_policy(mut self, policy: ModelRetryPolicy) -> Self {
189        self.retry_policy = Some(policy);
190        self
191    }
192
193    /// Replaces the asynchronous timer used by retry backoff.
194    #[must_use]
195    pub fn sleeper(mut self, sleeper: Arc<dyn RouterSleeper>) -> Self {
196        self.sleeper = sleeper;
197        self
198    }
199
200    /// Validates and builds the router.
201    ///
202    /// # Errors
203    ///
204    /// Returns [`ModelRouterBuildError`] for blank identities, duplicate route
205    /// names, or an empty route list.
206    pub fn build(self) -> Result<ModelRouter, ModelRouterBuildError> {
207        if let Some(error) = self.error {
208            return Err(error);
209        }
210        if self.logical.provider.trim().is_empty() || self.logical.name.trim().is_empty() {
211            return Err(ModelRouterBuildError::EmptyLogicalModel);
212        }
213        if self.routes.is_empty() {
214            return Err(ModelRouterBuildError::NoRoutes);
215        }
216        Ok(ModelRouter {
217            logical: self.logical,
218            routes: self.routes,
219            policy: self.policy,
220            circuit_breaker: self.circuit_breaker,
221            clock: self.clock,
222            retry_policy: self.retry_policy,
223            sleeper: self.sleeper,
224        })
225    }
226}
227
228impl std::fmt::Debug for ModelRouterBuilder {
229    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
230        formatter
231            .debug_struct("ModelRouterBuilder")
232            .field("logical", &self.logical)
233            .field("routes", &self.routes)
234            .field("policy", &self.policy)
235            .field("circuit_breaker", &self.circuit_breaker)
236            .field("retry_policy", &self.retry_policy)
237            .field("error", &self.error)
238            .finish_non_exhaustive()
239    }
240}
241
242/// Ordered, provider-neutral fallback routing behind the canonical [`Model`]
243/// boundary.
244#[derive(Clone)]
245pub struct ModelRouter {
246    logical: ModelRef,
247    routes: Vec<ModelRoute>,
248    policy: ModelFallbackPolicy,
249    circuit_breaker: Option<CircuitBreakerConfig>,
250    clock: Arc<dyn RouterClock>,
251    retry_policy: Option<ModelRetryPolicy>,
252    sleeper: Arc<dyn RouterSleeper>,
253}
254
255impl ModelRouter {
256    /// Starts a router builder for one logical model identity.
257    pub fn builder(logical: ModelRef) -> ModelRouterBuilder {
258        ModelRouterBuilder {
259            logical,
260            routes: Vec::new(),
261            policy: ModelFallbackPolicy::default(),
262            circuit_breaker: None,
263            clock: Arc::new(SystemRouterClock),
264            retry_policy: None,
265            sleeper: Arc::new(SystemRouterSleeper),
266            error: None,
267        }
268    }
269
270    /// Returns the identity applications use in [`ModelRequest`].
271    pub const fn logical_model(&self) -> &ModelRef {
272        &self.logical
273    }
274
275    /// Returns physical routes in deterministic selection order.
276    pub fn routes(&self) -> &[ModelRoute] {
277        &self.routes
278    }
279
280    /// Returns a point-in-time health snapshot for every physical route.
281    pub fn route_health(&self) -> Vec<ModelRouteHealth> {
282        let now = self.clock.now();
283        self.routes
284            .iter()
285            .map(|route| {
286                crate::circuit::snapshot(
287                    &route.health,
288                    route.name.clone(),
289                    route.target.clone(),
290                    self.circuit_breaker.as_ref(),
291                    now,
292                )
293            })
294            .collect()
295    }
296
297    fn validate_request(&self, request: &ModelRequest) -> Result<(), ModelError> {
298        if request.model != self.logical {
299            return Err(ModelError::local(
300                ModelErrorKind::InvalidRequest,
301                format!(
302                    "router for `{}/{}` cannot invoke logical model `{}/{}`",
303                    self.logical.provider,
304                    self.logical.name,
305                    request.model.provider,
306                    request.model.name
307                ),
308            ));
309        }
310        Ok(())
311    }
312}
313
314impl Model for ModelRouter {
315    fn capabilities<'a>(
316        &'a self,
317        model: &'a ModelRef,
318    ) -> ModelFuture<'a, Result<ModelCapabilities, ModelError>> {
319        Box::pin(async move {
320            if model != &self.logical {
321                return Err(ModelError::local(
322                    ModelErrorKind::InvalidRequest,
323                    "capabilities requested for the wrong logical model",
324                ));
325            }
326            let mut capabilities = Vec::with_capacity(self.routes.len());
327            for route in &self.routes {
328                capabilities.push(route.model.capabilities(&route.target).await?);
329            }
330            Ok(intersect_capabilities(capabilities))
331        })
332    }
333
334    fn stream(
335        &self,
336        request: ModelRequest,
337        context: ModelCallContext,
338    ) -> ModelFuture<'_, Result<ModelEventStream, ModelError>> {
339        let validation = self.validate_request(&request);
340        let stream = routed_stream(
341            self.routes.clone(),
342            self.policy.clone(),
343            RoutingRuntime {
344                circuit_breaker: self.circuit_breaker.clone(),
345                clock: self.clock.clone(),
346                retry_policy: self.retry_policy.clone(),
347                sleeper: self.sleeper.clone(),
348            },
349            request,
350            context,
351        );
352        Box::pin(async move {
353            validation?;
354            Ok(stream)
355        })
356    }
357}
358
359#[cfg(test)]
360mod tests;