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#[derive(Clone)]
22pub struct ModelRoute {
23 name: String,
24 model: Arc<dyn Model>,
25 target: ModelRef,
26 health: SharedBreakerState,
27}
28
29impl ModelRoute {
30 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 pub fn name(&self) -> &str {
42 &self.name
43 }
44
45 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#[derive(Clone, Debug, Default, Eq, PartialEq)]
67pub struct ModelFallbackPolicy {
68 unknown_safety_kinds: Vec<ModelErrorKind>,
69}
70
71impl ModelFallbackPolicy {
72 pub const fn safe_only() -> Self {
74 Self {
75 unknown_safety_kinds: Vec::new(),
76 }
77 }
78
79 #[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#[derive(Clone, Debug, Error, Eq, PartialEq)]
105#[non_exhaustive]
106pub enum ModelRouterBuildError {
107 #[error("logical model provider and name cannot be empty")]
109 EmptyLogicalModel,
110 #[error("model route name cannot be empty")]
112 EmptyRouteName,
113 #[error("physical model provider and name cannot be empty")]
115 EmptyTarget,
116 #[error("model route `{0}` is already registered")]
118 DuplicateRoute(String),
119 #[error("model router requires at least one route")]
121 NoRoutes,
122}
123
124pub 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 #[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 #[must_use]
167 pub fn fallback_policy(mut self, policy: ModelFallbackPolicy) -> Self {
168 self.policy = policy;
169 self
170 }
171
172 #[must_use]
174 pub fn circuit_breaker(mut self, config: CircuitBreakerConfig) -> Self {
175 self.circuit_breaker = Some(config);
176 self
177 }
178
179 #[must_use]
181 pub fn clock(mut self, clock: Arc<dyn RouterClock>) -> Self {
182 self.clock = clock;
183 self
184 }
185
186 #[must_use]
188 pub fn retry_policy(mut self, policy: ModelRetryPolicy) -> Self {
189 self.retry_policy = Some(policy);
190 self
191 }
192
193 #[must_use]
195 pub fn sleeper(mut self, sleeper: Arc<dyn RouterSleeper>) -> Self {
196 self.sleeper = sleeper;
197 self
198 }
199
200 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#[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 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 pub const fn logical_model(&self) -> &ModelRef {
277 &self.logical
278 }
279
280 pub fn routes(&self) -> &[ModelRoute] {
282 &self.routes
283 }
284
285 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;