use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::{Duration, SystemTime};
use tower::{Layer, Service};
use super::budget::{BudgetLedger, CostCheckContext, provider_of, should_hedge, user_id_of};
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;
fn allow_hedge(&self, _req: &LlmRequest) -> bool {
true
}
}
#[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 BudgetAwareHedge<P, L: BudgetLedger> {
inner: P,
ledger: Arc<L>,
estimated_cost_usd: f64,
safety_margin_pct: f64,
}
impl<P: HedgePolicy, L: BudgetLedger> BudgetAwareHedge<P, L> {
#[cfg_attr(alef, alef(skip))]
#[must_use]
pub fn new(inner: P, ledger: Arc<L>, estimated_cost_usd: f64, safety_margin_pct: f64) -> Self {
Self {
inner,
ledger,
estimated_cost_usd,
safety_margin_pct,
}
}
}
impl<P: HedgePolicy, L: BudgetLedger> HedgePolicy for BudgetAwareHedge<P, L> {
fn delay_for_attempt(&self, attempt: u32, latency_so_far: Duration) -> Option<Duration> {
self.inner.delay_for_attempt(attempt, latency_so_far)
}
fn max_attempts(&self) -> u32 {
self.inner.max_attempts()
}
fn allow_hedge(&self, req: &LlmRequest) -> bool {
if !self.inner.allow_hedge(req) {
return false;
}
let model = req.model().unwrap_or("");
let provider = provider_of(model);
let ctx = CostCheckContext {
model,
provider,
tenant_id: req.tenant_id().map(|t| t.as_ref()),
user_id: user_id_of(req),
api_key_id: None,
timestamp: SystemTime::now(),
};
should_hedge(
self.ledger.as_ref(),
&ctx,
self.estimated_cost_usd,
self.safety_margin_pct,
)
}
}
#[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 {
if !policy.allow_hedge(&req) {
tracing::debug!(attempt, "hedge suppressed: policy vetoed additional attempt");
break;
}
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::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)
);
}
struct ToggleableHedge {
inner: FixedDelayHedge,
allowed: Arc<std::sync::atomic::AtomicBool>,
}
impl HedgePolicy for ToggleableHedge {
fn delay_for_attempt(&self, attempt: u32, latency_so_far: Duration) -> Option<Duration> {
self.inner.delay_for_attempt(attempt, latency_so_far)
}
fn max_attempts(&self) -> u32 {
self.inner.max_attempts()
}
fn allow_hedge(&self, _req: &LlmRequest) -> bool {
self.allowed.load(Ordering::SeqCst)
}
}
#[tokio::test]
async fn allow_hedge_false_suppresses_additional_attempts() {
let policy = Arc::new(ToggleableHedge {
inner: FixedDelayHedge::new(Duration::ZERO, 3),
allowed: Arc::new(std::sync::atomic::AtomicBool::new(false)),
});
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,
"allow_hedge=false must suppress every attempt beyond the primary, even with max_attempts=3"
);
}
#[tokio::test]
async fn budget_aware_hedge_suppresses_when_ledger_is_near_limit() {
use crate::tower::budget::{CostRecordContext, DimensionLimits, InMemoryBudgetLedger};
let mut limits = DimensionLimits::default();
limits.per_user.insert("alice".to_owned(), 10.0);
let ledger = Arc::new(InMemoryBudgetLedger::new(limits, Duration::from_secs(3600)));
ledger
.record(&CostRecordContext {
model: "gpt-4",
provider: "openai",
tenant_id: None,
user_id: Some("alice"),
api_key_id: None,
cost_usd: 9.50,
tokens_in: 100,
tokens_out: 50,
timestamp: std::time::SystemTime::now(),
})
.await;
let near_limit_policy =
BudgetAwareHedge::new(FixedDelayHedge::new(Duration::ZERO, 2), Arc::clone(&ledger), 0.50, 0.10);
let mut req = chat_req("gpt-4");
req.user = Some("alice".into());
let llm_req = LlmRequest::Chat(req);
assert!(
!near_limit_policy.allow_hedge(&llm_req),
"hedging must be suppressed once spend + 2x estimated cost would exceed 90% of the user budget"
);
let far_from_limit_policy = BudgetAwareHedge::new(FixedDelayHedge::new(Duration::ZERO, 2), ledger, 0.01, 0.10);
let mut req2 = chat_req("gpt-4");
req2.user = Some("bob".into());
let llm_req2 = LlmRequest::Chat(req2);
assert!(
far_from_limit_policy.allow_hedge(&llm_req2),
"a user with no recorded spend and a tiny estimated cost must be allowed to hedge"
);
}
}