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