Skip to main content

runifold_model/router/
execution.rs

1//! Runtime route selection, retries, circuit permits, and stream commitment.
2
3use 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: 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        if delay >= remaining {
160            return RetryWait::Stop(retry_deadline(delay, remaining));
161        }
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,
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}