use std::{future::Future, sync::Mutex, time::Duration};
use tokio::time::Instant;
use dynamo_runtime::{
error::{DynamoError, ErrorType},
pipeline::{AsyncEngineContext, Error},
};
use crate::{preprocessor::PreprocessedRequest, protocols::common::timing::RequestPhase};
const CLEANUP_DISPATCH_TIMEOUT: Duration = Duration::from_secs(120);
#[derive(Debug, Default)]
pub(super) struct CleanupBudget {
started: Mutex<Option<Instant>>,
}
impl CleanupBudget {
pub(super) fn remaining(&self) -> Duration {
let mut started = self.started.lock().unwrap();
let start = *started.get_or_insert_with(Instant::now);
CLEANUP_DISPATCH_TIMEOUT.saturating_sub(start.elapsed())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum StagedKv {
Present,
Absent,
}
impl StagedKv {
pub(super) fn for_request(request: &PreprocessedRequest) -> Self {
if request.staged_kv_cleanup {
Self::Present
} else {
Self::Absent
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum DispatchCancellation {
CancelWhenStopped,
DispatchWhenStopped,
}
impl DispatchCancellation {
pub(super) fn for_request(phase: RequestPhase, staged_kv: StagedKv) -> Self {
match (phase, staged_kv) {
(RequestPhase::Decode, StagedKv::Present) => Self::DispatchWhenStopped,
(RequestPhase::Decode, StagedKv::Absent)
| (RequestPhase::Prefill, _)
| (RequestPhase::Aggregated, _) => Self::CancelWhenStopped,
}
}
}
pub(super) async fn await_with_cleanup_policy<T>(
context: &dyn AsyncEngineContext,
phase: RequestPhase,
staged_kv: StagedKv,
stage: &'static str,
budget: &CleanupBudget,
operation: impl Future<Output = T>,
) -> Result<T, Error> {
match DispatchCancellation::for_request(phase, staged_kv) {
DispatchCancellation::CancelWhenStopped => cancel_on_stop(context, operation).await,
DispatchCancellation::DispatchWhenStopped => {
tokio::pin!(operation);
tokio::select! {
biased;
result = &mut operation => Ok(result),
_ = context.stopped() => {
match tokio::time::timeout(budget.remaining(), &mut operation).await {
Ok(result) => Ok(result),
Err(_) => Err(cleanup_budget_exhausted(context.id(), stage)),
}
}
}
}
}
}
fn cleanup_budget_exhausted(context_id: &str, stage: &'static str) -> Error {
tracing::warn!(
request_id = %context_id,
stage,
budget_secs = CLEANUP_DISPATCH_TIMEOUT.as_secs(),
"decode cleanup budget exhausted before the worker was reached; staged KV \
blocks will not be released until they expire"
);
DynamoError::builder()
.error_type(ErrorType::Cancelled)
.message(format!(
"Request {context_id} exhausted its decode cleanup budget at {stage}"
))
.build()
.into()
}
pub(super) fn cancelled_error(context_id: &str) -> Error {
DynamoError::builder()
.error_type(ErrorType::Cancelled)
.message(format!("Request {context_id} was cancelled"))
.build()
.into()
}
pub(super) async fn cancel_on_stop<T>(
context: &dyn AsyncEngineContext,
operation: impl Future<Output = T>,
) -> Result<T, Error> {
tokio::pin!(operation);
tokio::select! {
biased;
result = &mut operation => Ok(result),
_ = context.stopped() => Err(cancelled_error(context.id())),
}
}
#[cfg(test)]
mod tests {
use std::{
future::Future,
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::Duration,
};
use dynamo_runtime::{
error::{DynamoError, ErrorType},
pipeline::{AsyncEngineContext, context::Controller},
};
use super::{CLEANUP_DISPATCH_TIMEOUT, CleanupBudget, cancel_on_stop};
struct PendingUntilDropped(Arc<AtomicBool>);
impl Future for PendingUntilDropped {
type Output = ();
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
Poll::Pending
}
}
impl Drop for PendingUntilDropped {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
#[tokio::test(start_paused = true)]
async fn cleanup_budget_decays_and_saturates_at_zero() {
let budget = CleanupBudget::default();
assert_eq!(budget.remaining(), CLEANUP_DISPATCH_TIMEOUT);
tokio::time::advance(Duration::from_secs(90)).await;
assert_eq!(
budget.remaining(),
CLEANUP_DISPATCH_TIMEOUT - Duration::from_secs(90)
);
tokio::time::advance(Duration::from_secs(60)).await;
assert_eq!(budget.remaining(), Duration::ZERO);
}
#[tokio::test]
async fn drops_pending_operation_when_context_stops() {
let context = Controller::new("cancelled-request".to_string());
context.stop();
let dropped = Arc::new(AtomicBool::new(false));
let error = cancel_on_stop(&context, PendingUntilDropped(dropped.clone()))
.await
.unwrap_err();
let error = error
.downcast_ref::<DynamoError>()
.expect("cancellation should return DynamoError");
assert_eq!(error.error_type(), ErrorType::Cancelled);
assert!(dropped.load(Ordering::SeqCst));
}
#[tokio::test]
async fn ready_operation_wins_if_context_is_already_stopped() {
let context = Controller::new("completed-request".to_string());
context.stop();
let result = cancel_on_stop(&context, std::future::ready(42))
.await
.unwrap();
assert_eq!(result, 42);
}
}