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 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}