1use std::sync::Arc;
4
5use futures_util::{
6 StreamExt,
7 future::{Either, select},
8};
9use runifold_core::RetrySafety;
10use serde_json::{Value, json};
11
12use super::{
13 BreakerPermit, CircuitBreakerConfig, ModelCallContext, ModelError, ModelErrorKind,
14 ModelEventStream, ModelFallbackPolicy, ModelRef, ModelRequest, ModelRetryPolicy, ModelRoute,
15 ModelRouter, ModelStreamEvent, ProviderEvent, RoutePermit, RouterClock, RouterSleeper,
16};
17
18pub(super) struct RoutingRuntime {
19 pub(super) circuit_breaker: Option<CircuitBreakerConfig>,
20 pub(super) clock: Arc<dyn RouterClock>,
21 pub(super) retry_policy: Option<ModelRetryPolicy>,
22 pub(super) sleeper: Arc<dyn RouterSleeper>,
23}
24
25pub(super) fn routed_stream(
26 routes: Vec<ModelRoute>,
27 fallback: ModelFallbackPolicy,
28 runtime: RoutingRuntime,
29 request: ModelRequest,
30 context: ModelCallContext,
31) -> ModelEventStream {
32 Box::pin(async_stream::stream! {
33 let mut failures = Vec::new();
34 'routes: for (index, route) in routes.iter().enumerate() {
35 let max_attempts = runtime
36 .retry_policy
37 .as_ref()
38 .map_or(1, ModelRetryPolicy::max_attempts);
39 for route_attempt in 1..=max_attempts {
40 if context.cancellation().is_cancelled() {
41 let error = cancelled_retry();
42 yield Err(annotate_error(error, &request.model, failures));
43 return;
44 }
45 match start_route_attempt(route, &request, &context, &runtime).await {
46 RouteAttempt::CircuitOpen => {
47 failures.push(circuit_open_summary(route, route_attempt));
48 continue 'routes;
49 }
50 RouteAttempt::Committed {
51 first,
52 stream,
53 permit,
54 attempt_id,
55 } => {
56 let terminal = is_terminal(&first);
57 yield Ok(*first);
58 if terminal {
59 succeed_permit(permit);
60 return;
61 }
62 let circuit_probe =
63 permit.as_ref().is_some_and(BreakerPermit::is_probe);
64 let mut committed = committed_stream(
65 CommittedRoute {
66 route: route.clone(),
67 index,
68 route_attempt,
69 attempt_id,
70 circuit_probe,
71 },
72 stream,
73 permit,
74 failures,
75 request.model.clone(),
76 );
77 while let Some(item) = committed.next().await {
78 yield item;
79 }
80 return;
81 }
82 RouteAttempt::Failed(error) => {
83 failures.push(failure_summary(route, route_attempt, &error));
84 let retry = runtime.retry_policy.as_ref().is_some_and(|policy| {
85 route_attempt < max_attempts && policy.permits(&error)
86 });
87 if retry {
88 match wait_before_retry(
89 runtime.retry_policy.as_ref().expect("retry policy exists"),
90 &error,
91 route,
92 route_attempt,
93 &context,
94 &runtime.sleeper,
95 )
96 .await
97 {
98 RetryWait::Ready => continue,
99 RetryWait::Stop(stop) => {
100 yield Err(annotate_error(
101 stop,
102 &request.model,
103 failures,
104 ));
105 return;
106 }
107 }
108 }
109 if index + 1 < routes.len() && fallback.permits(&error) {
110 continue 'routes;
111 }
112 yield Err(annotate_error(error, &request.model, failures));
113 return;
114 }
115 }
116 }
117 }
118 let mut error = ModelError::local(
119 ModelErrorKind::Provider,
120 "all physical model routes are unavailable",
121 );
122 error.retry_safety = RetrySafety::Safe;
123 yield Err(annotate_error(error, &request.model, failures));
124 })
125}
126
127enum RouteAttempt {
128 CircuitOpen,
129 Failed(ModelError),
130 Committed {
131 first: Box<ModelStreamEvent>,
132 stream: ModelEventStream,
133 permit: Option<BreakerPermit>,
134 attempt_id: String,
135 },
136}
137
138enum RetryWait {
139 Ready,
140 Stop(ModelError),
141}
142
143async fn wait_before_retry(
144 policy: &ModelRetryPolicy,
145 error: &ModelError,
146 route: &ModelRoute,
147 failed_attempt: u32,
148 context: &ModelCallContext,
149 sleeper: &Arc<dyn RouterSleeper>,
150) -> RetryWait {
151 let entropy = retry_entropy(
152 &context.invocation_id().to_string(),
153 &route.name,
154 failed_attempt,
155 );
156 let backoff = policy.delay(failed_attempt, entropy);
157 let delay = retry_after(error).map_or(backoff, |server| server.max(backoff));
158 if let Some(remaining) = context.remaining()
159 && delay >= remaining
160 {
161 return RetryWait::Stop(retry_deadline(delay, remaining));
162 }
163 if context.cancellation().is_cancelled() {
164 return RetryWait::Stop(cancelled_retry());
165 }
166 if delay.is_zero() {
167 return RetryWait::Ready;
168 }
169 let cancellation = context.cancellation().clone();
170 match select(
171 Box::pin(cancellation.cancelled()),
172 Box::pin(sleeper.sleep(delay)),
173 )
174 .await
175 {
176 Either::Left(_) => RetryWait::Stop(cancelled_retry()),
177 Either::Right(_) => RetryWait::Ready,
178 }
179}
180
181fn retry_after(error: &ModelError) -> Option<std::time::Duration> {
182 error
183 .metadata
184 .get("retry.after_ms")
185 .and_then(Value::as_u64)
186 .map(std::time::Duration::from_millis)
187}
188
189fn retry_entropy(invocation: &str, route: &str, attempt: u32) -> u64 {
190 let mut hash = 0xcbf2_9ce4_8422_2325_u64;
191 for byte in invocation
192 .bytes()
193 .chain(route.bytes())
194 .chain(attempt.to_le_bytes())
195 {
196 hash ^= u64::from(byte);
197 hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
198 }
199 hash
200}
201
202fn cancelled_retry() -> ModelError {
203 ModelError::local(
204 ModelErrorKind::Cancelled,
205 "logical model invocation was cancelled during retry",
206 )
207}
208
209fn retry_deadline(delay: std::time::Duration, remaining: std::time::Duration) -> ModelError {
210 let mut error = ModelError::local(
211 ModelErrorKind::DeadlineExceeded,
212 "retry backoff would exceed the model invocation deadline",
213 );
214 error.metadata.insert(
215 "retry.delay_ms".into(),
216 Value::from(u64::try_from(delay.as_millis()).unwrap_or(u64::MAX)),
217 );
218 error.metadata.insert(
219 "retry.remaining_ms".into(),
220 Value::from(u64::try_from(remaining.as_millis()).unwrap_or(u64::MAX)),
221 );
222 error
223}
224
225async fn start_route_attempt(
226 route: &ModelRoute,
227 request: &ModelRequest,
228 context: &ModelCallContext,
229 runtime: &RoutingRuntime,
230) -> RouteAttempt {
231 let permit = match crate::circuit::acquire(
232 &route.health,
233 runtime.circuit_breaker.as_ref(),
234 &runtime.clock,
235 ) {
236 RoutePermit::Disabled => None,
237 RoutePermit::Acquired(permit) => Some(permit),
238 RoutePermit::Rejected => return RouteAttempt::CircuitOpen,
239 };
240 let mut routed_request = request.clone();
241 routed_request.model.clone_from(&route.target);
242 let attempt = context.child_attempt();
243 let attempt_id = attempt.invocation_id().to_string();
244 let opened = route.model.stream(routed_request, attempt).await;
245 let mut stream = match opened {
246 Ok(stream) => stream,
247 Err(error) => {
248 fail_permit(permit, &error);
249 return RouteAttempt::Failed(error);
250 }
251 };
252 match stream.next().await {
253 Some(Ok(first)) => RouteAttempt::Committed {
254 first: Box::new(first),
255 stream,
256 permit,
257 attempt_id,
258 },
259 Some(Err(error)) => {
260 fail_permit(permit, &error);
261 RouteAttempt::Failed(error)
262 }
263 None => {
264 let error = ModelError::local(
265 ModelErrorKind::Protocol,
266 "candidate model stream ended before its first event",
267 );
268 fail_permit(permit, &error);
269 RouteAttempt::Failed(error)
270 }
271 }
272}
273
274struct CommittedRoute {
275 route: ModelRoute,
276 index: usize,
277 route_attempt: u32,
278 attempt_id: String,
279 circuit_probe: bool,
280}
281
282fn committed_stream(
283 committed: CommittedRoute,
284 mut stream: ModelEventStream,
285 permit: Option<BreakerPermit>,
286 mut failures: Vec<Value>,
287 logical: ModelRef,
288) -> ModelEventStream {
289 Box::pin(async_stream::stream! {
290 yield Ok(selected_event(
291 &committed.route,
292 committed.index,
293 committed.route_attempt,
294 &committed.attempt_id,
295 committed.circuit_probe,
296 &failures,
297 ));
298 while let Some(item) = stream.next().await {
299 match item {
300 Ok(event) => {
301 let terminal = is_terminal(&event);
302 yield Ok(event);
303 if terminal {
304 succeed_permit(permit);
305 return;
306 }
307 }
308 Err(mut error) => {
309 error.retry_safety = RetrySafety::UnsafeAfterVisibleOutput;
310 fail_permit(permit, &error);
311 failures.push(failure_summary(
312 &committed.route,
313 committed.route_attempt,
314 &error,
315 ));
316 yield Err(annotate_error(error, &logical, failures));
317 return;
318 }
319 }
320 }
321 let mut error = ModelError::local(
322 ModelErrorKind::Protocol,
323 "selected model stream ended without a terminal event",
324 );
325 error.retry_safety = RetrySafety::UnsafeAfterVisibleOutput;
326 fail_permit(permit, &error);
327 failures.push(failure_summary(
328 &committed.route,
329 committed.route_attempt,
330 &error,
331 ));
332 yield Err(annotate_error(error, &logical, failures));
333 })
334}
335
336fn fail_permit(permit: Option<BreakerPermit>, error: &ModelError) {
337 if let Some(permit) = permit {
338 permit.failure(error);
339 }
340}
341
342fn succeed_permit(permit: Option<BreakerPermit>) {
343 if let Some(permit) = permit {
344 permit.success();
345 }
346}
347
348impl std::fmt::Debug for ModelRouter {
349 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
350 formatter
351 .debug_struct("ModelRouter")
352 .field("logical", &self.logical)
353 .field("routes", &self.routes)
354 .field("policy", &self.policy)
355 .field("circuit_breaker", &self.circuit_breaker)
356 .field("retry_policy", &self.retry_policy)
357 .finish_non_exhaustive()
358 }
359}
360
361fn selected_event(
362 route: &ModelRoute,
363 index: usize,
364 route_attempt: u32,
365 attempt_id: &str,
366 circuit_probe: bool,
367 failures: &[Value],
368) -> ModelStreamEvent {
369 ModelStreamEvent::Provider {
370 event: ProviderEvent {
371 provider: "runifold.router".into(),
372 name: "route.selected".into(),
373 payload: json!({
374 "route": route.name,
375 "target": route.target,
376 "index": index,
377 "route_attempt": route_attempt,
378 "attempt_id": attempt_id,
379 "circuit_probe": circuit_probe,
380 "prior_failures": failures,
381 }),
382 },
383 }
384}
385
386fn failure_summary(route: &ModelRoute, route_attempt: u32, error: &ModelError) -> Value {
387 json!({
388 "route": route.name,
389 "target": route.target,
390 "route_attempt": route_attempt,
391 "kind": &error.kind,
392 "retry_safety": error.retry_safety,
393 })
394}
395
396fn circuit_open_summary(route: &ModelRoute, route_attempt: u32) -> Value {
397 json!({
398 "route": route.name,
399 "target": route.target,
400 "route_attempt": route_attempt,
401 "kind": "circuit_open",
402 })
403}
404
405const fn is_terminal(event: &ModelStreamEvent) -> bool {
406 matches!(event, ModelStreamEvent::ResponseCompleted { .. })
407}
408
409fn annotate_error(mut error: ModelError, logical: &ModelRef, failures: Vec<Value>) -> ModelError {
410 error.metadata.insert(
411 "runifold.router.logical_model".into(),
412 serde_json::to_value(logical).expect("ModelRef is serializable"),
413 );
414 error
415 .metadata
416 .insert("runifold.router.failures".into(), Value::Array(failures));
417 error
418}