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///
245/// A router owns process-local retry and circuit-breaker state. Applications
246/// should build it once and reuse this value or its clones for the lifetime of
247/// the service. Clones share route health; rebuilding a router intentionally
248/// starts with fresh health state.
249#[derive(Clone)]
250pub struct ModelRouter {
251    logical: ModelRef,
252    routes: Vec<ModelRoute>,
253    policy: ModelFallbackPolicy,
254    circuit_breaker: Option<CircuitBreakerConfig>,
255    clock: Arc<dyn RouterClock>,
256    retry_policy: Option<ModelRetryPolicy>,
257    sleeper: Arc<dyn RouterSleeper>,
258}
259
260impl ModelRouter {
261    /// Starts a router builder for one logical model identity.
262    pub fn builder(logical: ModelRef) -> ModelRouterBuilder {
263        ModelRouterBuilder {
264            logical,
265            routes: Vec::new(),
266            policy: ModelFallbackPolicy::default(),
267            circuit_breaker: None,
268            clock: Arc::new(SystemRouterClock),
269            retry_policy: None,
270            sleeper: Arc::new(SystemRouterSleeper),
271            error: None,
272        }
273    }
274
275    /// Returns the identity applications use in [`ModelRequest`].
276    pub const fn logical_model(&self) -> &ModelRef {
277        &self.logical
278    }
279
280    /// Returns physical routes in deterministic selection order.
281    pub fn routes(&self) -> &[ModelRoute] {
282        &self.routes
283    }
284
285    /// Returns a point-in-time health snapshot for every physical route.
286    pub fn route_health(&self) -> Vec<ModelRouteHealth> {
287        let now = self.clock.now();
288        self.routes
289            .iter()
290            .map(|route| {
291                crate::circuit::snapshot(
292                    &route.health,
293                    route.name.clone(),
294                    route.target.clone(),
295                    self.circuit_breaker.as_ref(),
296                    now,
297                )
298            })
299            .collect()
300    }
301
302    fn validate_request(&self, request: &ModelRequest) -> Result<(), ModelError> {
303        if request.model != self.logical {
304            return Err(ModelError::local(
305                ModelErrorKind::InvalidRequest,
306                format!(
307                    "router for `{}/{}` cannot invoke logical model `{}/{}`",
308                    self.logical.provider,
309                    self.logical.name,
310                    request.model.provider,
311                    request.model.name
312                ),
313            ));
314        }
315        Ok(())
316    }
317}
318
319impl Model for ModelRouter {
320    fn capabilities<'a>(
321        &'a self,
322        model: &'a ModelRef,
323    ) -> ModelFuture<'a, Result<ModelCapabilities, ModelError>> {
324        Box::pin(async move {
325            if model != &self.logical {
326                return Err(ModelError::local(
327                    ModelErrorKind::InvalidRequest,
328                    "capabilities requested for the wrong logical model",
329                ));
330            }
331            let mut capabilities = Vec::with_capacity(self.routes.len());
332            for route in &self.routes {
333                capabilities.push(route.model.capabilities(&route.target).await?);
334            }
335            Ok(intersect_capabilities(capabilities))
336        })
337    }
338
339    fn stream(
340        &self,
341        request: ModelRequest,
342        context: ModelCallContext,
343    ) -> ModelFuture<'_, Result<ModelEventStream, ModelError>> {
344        let validation = self.validate_request(&request);
345        let stream = routed_stream(
346            self.routes.clone(),
347            self.policy.clone(),
348            RoutingRuntime {
349                circuit_breaker: self.circuit_breaker.clone(),
350                clock: self.clock.clone(),
351                retry_policy: self.retry_policy.clone(),
352                sleeper: self.sleeper.clone(),
353            },
354            request,
355            context,
356        );
357        Box::pin(async move {
358            validation?;
359            Ok(stream)
360        })
361    }
362}
363
364#[cfg(test)]
365mod tests;