use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use tower::{Layer, Service};
use super::types::{LlmRequest, LlmResponse};
use crate::client::BoxFuture;
use crate::error::{LiterLlmError, Result};
#[cfg_attr(alef, alef(skip))]
pub trait HedgePolicy: Send + Sync + 'static {
fn delay_for_attempt(&self, attempt: u32, latency_so_far: Duration) -> Option<Duration>;
fn max_attempts(&self) -> u32;
}
#[cfg_attr(alef, alef(skip))]
pub struct FixedDelayHedge {
delay: Duration,
max_attempts: u32,
}
impl FixedDelayHedge {
#[must_use]
pub fn new(delay: Duration, max_attempts: u32) -> Self {
Self {
delay,
max_attempts: max_attempts.max(1),
}
}
}
impl HedgePolicy for FixedDelayHedge {
fn delay_for_attempt(&self, attempt: u32, _latency_so_far: Duration) -> Option<Duration> {
if attempt > self.max_attempts {
return None;
}
Some(self.delay * (attempt - 1))
}
fn max_attempts(&self) -> u32 {
self.max_attempts
}
}
#[cfg_attr(alef, alef(skip))]
pub struct HedgeLayer<P> {
policy: Arc<P>,
}
impl<P: HedgePolicy> HedgeLayer<P> {
#[must_use]
pub fn new(policy: Arc<P>) -> Self {
Self { policy }
}
}
impl<P: HedgePolicy, S> Layer<S> for HedgeLayer<P> {
type Service = HedgeService<P, S>;
fn layer(&self, inner: S) -> Self::Service {
HedgeService {
inner,
policy: Arc::clone(&self.policy),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct HedgeService<P, S> {
inner: S,
policy: Arc<P>,
}
impl<P: HedgePolicy, S: Clone> Clone for HedgeService<P, S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
policy: Arc::clone(&self.policy),
}
}
}
impl<P, S> Service<LlmRequest> for HedgeService<P, S>
where
P: HedgePolicy + 'static,
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + Clone + 'static,
S::Future: Send + 'static,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<()>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
let policy = Arc::clone(&self.policy);
let max_attempts = policy.max_attempts();
let standby = self.inner.clone(); let primary = std::mem::replace(&mut self.inner, standby);
let inner_for_hedges = self.inner.clone();
Box::pin(async move {
tracing::debug!(hedge.max_attempts = max_attempts, "starting hedged request");
hedge_race(req, primary, inner_for_hedges, policy, max_attempts).await
})
}
}
async fn hedge_race<S>(
req: LlmRequest,
mut primary: S,
inner_for_hedges: S,
policy: Arc<impl HedgePolicy>,
max_attempts: u32,
) -> Result<LlmResponse>
where
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + Clone + 'static,
S::Future: Send + 'static,
{
use std::time::Instant;
use tower::ServiceExt as _;
let dispatch_time = Instant::now();
if max_attempts == 1 {
tracing::debug!("hedge fast path: max_attempts=1, calling primary directly");
return primary.call(req).await;
}
let mut join_set: tokio::task::JoinSet<(u32, Result<LlmResponse>)> = tokio::task::JoinSet::new();
{
let req_clone = req.clone();
join_set.spawn(async move {
let result = primary.call(req_clone).await;
(1u32, result)
});
}
for attempt in 2..=max_attempts {
let latency_so_far = dispatch_time.elapsed();
let Some(hedge_delay) = policy.delay_for_attempt(attempt, latency_so_far) else {
break;
};
let req_clone = req.clone();
let mut svc_clone = inner_for_hedges.clone();
join_set.spawn(async move {
if hedge_delay > Duration::ZERO {
tokio::time::sleep(hedge_delay).await;
}
tracing::debug!(attempt, "launching hedged request");
let model = req_clone.model().unwrap_or("").to_owned();
let system = model.split_once('/').map(|(p, _)| p.to_owned()).unwrap_or_default();
super::metrics::record_retry_attempt(&system, &model, req_clone.operation_name());
let ready_result = svc_clone.ready().await;
let result = match ready_result {
Ok(ready_svc) => ready_svc.call(req_clone).await,
Err(e) => Err(e),
};
(attempt, result)
});
}
let mut last_err: Option<LiterLlmError> = None;
while let Some(join_result) = join_set.join_next().await {
match join_result {
Ok((attempt, Ok(resp))) => {
tracing::debug!(attempt, "hedged request succeeded first");
join_set.abort_all();
return Ok(resp);
}
Ok((attempt, Err(e))) => {
tracing::debug!(attempt, error = %e, "hedged attempt failed");
last_err = Some(e);
}
Err(join_err) if join_err.is_cancelled() => {
}
Err(join_err) => {
tracing::error!(error = %join_err, "hedged task panicked");
last_err = Some(LiterLlmError::InternalError {
message: format!("hedge task panicked: {join_err}"),
});
}
}
}
Err(last_err.unwrap_or(LiterLlmError::InternalError {
message: "all hedged attempts failed with no error recorded".into(),
}))
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tower::{Layer as _, Service as _, ServiceExt as _};
use super::*;
use crate::tower::service::LlmService;
use crate::tower::tests_common::{MockClient, chat_req};
use crate::tower::types::LlmRequest;
#[tokio::test]
async fn hedge_returns_first_success() {
let policy = Arc::new(FixedDelayHedge::new(Duration::from_millis(200), 2));
let inner = LlmService::new(MockClient::ok());
let mut svc = HedgeLayer::new(policy).layer(inner);
let resp = svc
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("should succeed");
assert!(matches!(resp, LlmResponse::Chat(_)));
}
#[tokio::test]
async fn hedge_single_attempt_policy_does_not_spawn_extra() {
let policy = Arc::new(FixedDelayHedge::new(Duration::from_millis(100), 1));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = HedgeLayer::new(policy).layer(inner);
svc.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("should succeed");
assert_eq!(
call_count.load(Ordering::SeqCst),
1,
"max_attempts=1 should only call inner service once"
);
}
#[tokio::test]
async fn hedge_propagates_error_when_all_attempts_fail() {
let policy = Arc::new(FixedDelayHedge::new(Duration::from_millis(10), 2));
let inner = LlmService::new(MockClient::failing_timeout());
let mut svc = HedgeLayer::new(policy).layer(inner);
let err = svc
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect_err("all attempts should fail");
assert!(
matches!(err, LiterLlmError::Timeout),
"expected Timeout from failed hedge, got {err:?}"
);
}
#[tokio::test]
async fn fixed_delay_hedge_policy_respects_max_attempts() {
let policy = FixedDelayHedge::new(Duration::from_millis(100), 3);
assert_eq!(policy.delay_for_attempt(1, Duration::ZERO), Some(Duration::ZERO));
assert_eq!(
policy.delay_for_attempt(2, Duration::ZERO),
Some(Duration::from_millis(100))
);
assert_eq!(
policy.delay_for_attempt(3, Duration::ZERO),
Some(Duration::from_millis(200))
);
assert_eq!(policy.delay_for_attempt(4, Duration::ZERO), None);
}
#[tokio::test]
async fn hedge_with_two_attempts_calls_inner_at_most_twice() {
let policy = Arc::new(FixedDelayHedge::new(Duration::from_millis(5), 2));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = HedgeLayer::new(policy).layer(inner);
svc.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("should succeed");
let count = call_count.load(Ordering::SeqCst);
assert!((1..=2).contains(&count), "expected 1 or 2 calls, got {count}");
}
#[tokio::test]
async fn hedge_max_attempts_one_does_not_spawn_extra() {
let policy = Arc::new(FixedDelayHedge::new(Duration::from_millis(0), 1));
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let mut svc = HedgeLayer::new(policy).layer(inner);
for _ in 0..2 {
svc.ready()
.await
.expect("service should become ready")
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("should succeed");
}
assert_eq!(
call_count.load(Ordering::SeqCst),
2,
"max_attempts=1 must not spawn additional tasks; expected exactly 2 calls total"
);
}
#[tokio::test]
async fn hedge_respects_inner_readiness_via_ready_and_call() {
use tower::limit::ConcurrencyLimitLayer;
let mock = MockClient::ok();
let call_count = Arc::clone(&mock.call_count);
let inner = LlmService::new(mock);
let limited = ConcurrencyLimitLayer::new(1).layer(inner);
let policy = Arc::new(FixedDelayHedge::new(Duration::ZERO, 2));
let mut svc = HedgeLayer::new(policy).layer(limited);
let resp = svc
.ready()
.await
.expect("service should become ready")
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("hedged request should succeed");
assert!(matches!(resp, LlmResponse::Chat(_)));
let count = call_count.load(Ordering::SeqCst);
assert!(
(1..=2).contains(&count),
"expected 1 or 2 inner calls with ConcurrencyLimit(1), got {count}"
);
}
#[tokio::test]
async fn hedge_no_double_permit_consumption() {
use std::sync::atomic::AtomicUsize;
use tower::limit::ConcurrencyLimit;
let peak = Arc::new(AtomicUsize::new(0));
let current = Arc::new(AtomicUsize::new(0));
#[derive(Clone)]
struct PeakTracker {
peak: Arc<AtomicUsize>,
current: Arc<AtomicUsize>,
}
impl tower::Service<LlmRequest> for PeakTracker {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = crate::client::BoxFuture<'static, crate::error::Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut std::task::Context<'_>) -> std::task::Poll<crate::error::Result<()>> {
std::task::Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
let peak = Arc::clone(&self.peak);
let current = Arc::clone(&self.current);
Box::pin(async move {
let now = current.fetch_add(1, Ordering::SeqCst) + 1;
let mut prev = peak.load(Ordering::SeqCst);
while now > prev {
match peak.compare_exchange(prev, now, Ordering::SeqCst, Ordering::SeqCst) {
Ok(_) => break,
Err(p) => prev = p,
}
}
tokio::task::yield_now().await;
current.fetch_sub(1, Ordering::SeqCst);
Ok(LlmResponse::Chat(crate::tower::tests_common::make_chat_response(
"gpt-4",
)))
})
}
}
let tracker = PeakTracker {
peak: Arc::clone(&peak),
current: Arc::clone(¤t),
};
let limited: ConcurrencyLimit<PeakTracker> = ConcurrencyLimit::new(tracker, 1);
let policy = Arc::new(FixedDelayHedge::new(Duration::ZERO, 2));
let mut svc = HedgeLayer::new(policy).layer(limited);
let resp = svc
.ready()
.await
.expect("service should become ready")
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("hedged request must succeed");
assert!(matches!(resp, LlmResponse::Chat(_)));
assert_eq!(
peak.load(Ordering::SeqCst),
1,
"ConcurrencyLimit(1) must cap concurrent calls at 1 even with hedging"
);
}
#[tokio::test]
async fn hedge_loser_is_dropped_before_winner_returns() {
use std::sync::atomic::AtomicUsize;
use std::task::Poll;
use tokio::sync::Notify;
struct DropGuard {
live_count: Arc<AtomicUsize>,
}
impl Drop for DropGuard {
fn drop(&mut self) {
self.live_count.fetch_sub(1, Ordering::SeqCst);
}
}
let live = Arc::new(AtomicUsize::new(0));
let total_calls = Arc::new(AtomicUsize::new(0));
let winner_signal = Arc::new(Notify::new());
#[derive(Clone)]
struct SlowOrFast {
live: Arc<AtomicUsize>,
total_calls: Arc<AtomicUsize>,
attempt: Arc<AtomicUsize>,
winner_signal: Arc<Notify>,
}
impl tower::Service<LlmRequest> for SlowOrFast {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = crate::client::BoxFuture<'static, crate::error::Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut std::task::Context<'_>) -> Poll<crate::error::Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
let attempt = self.attempt.fetch_add(1, Ordering::SeqCst) + 1;
self.total_calls.fetch_add(1, Ordering::SeqCst);
self.live.fetch_add(1, Ordering::SeqCst);
let guard = DropGuard {
live_count: Arc::clone(&self.live),
};
let winner_signal = Arc::clone(&self.winner_signal);
Box::pin(async move {
let _g = guard;
if attempt == 1 {
tokio::time::sleep(Duration::from_millis(20)).await;
winner_signal.notify_one();
Ok(LlmResponse::Chat(crate::tower::tests_common::make_chat_response(
"gpt-4",
)))
} else {
std::future::pending::<()>().await;
unreachable!("loser must be cancelled before completing");
}
})
}
}
let inner = SlowOrFast {
live: Arc::clone(&live),
total_calls: Arc::clone(&total_calls),
attempt: Arc::new(AtomicUsize::new(0)),
winner_signal: Arc::clone(&winner_signal),
};
let policy = Arc::new(FixedDelayHedge::new(Duration::ZERO, 2));
let mut svc = HedgeLayer::new(policy).layer(inner);
let resp = svc
.ready()
.await
.expect("ready")
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("winner must succeed");
assert!(matches!(resp, LlmResponse::Chat(_)));
assert_eq!(
total_calls.load(Ordering::SeqCst),
2,
"expected primary + 1 hedged attempt"
);
for _ in 0..50 {
if live.load(Ordering::SeqCst) == 0 {
break;
}
tokio::task::yield_now().await;
tokio::time::sleep(Duration::from_millis(2)).await;
}
assert_eq!(
live.load(Ordering::SeqCst),
0,
"loser future must be dropped after winner returns; {} still alive",
live.load(Ordering::SeqCst)
);
}
}