use std::collections::BTreeMap;
use std::error::Error as StdError;
use std::sync::Arc;
use std::time::Instant;
use anyhow::{Error, Result};
use futures::{stream, stream::StreamExt};
use crate::{
http::service::metrics::Metrics,
model_card::ModelDeploymentCard,
protocols::{
TokenIdType,
common::{
extensions::{SESSION_AFFINITY_CONTEXT_KEY, SessionAffinityId},
llm_backend::{BackendOutput, LLMEngineOutput, PreprocessedRequest},
preprocessor::MultimodalData,
timing::RequestPhase,
},
},
session_affinity::explicit_target,
};
use dynamo_runtime::engine::Data;
use dynamo_runtime::error::{self, BackendError, DynamoError, ErrorType};
use dynamo_runtime::metrics::prometheus_names::frontend_service;
use dynamo_runtime::pipeline::{
AsyncEngineContext, AsyncEngineContextProvider, Context, ManyOut, Operator, PipelineOperator,
ResponseStream, ServerStreamingEngine, SingleIn, async_trait, attach_first_response_guard,
network::egress::route_span::{
RouteTraceContext, attach_route_trace_context, error_type_from_chain, error_type_name,
},
};
use dynamo_runtime::protocols::annotated::Annotated;
pub(crate) trait HasTokenIds {
fn token_ids(&self) -> &[TokenIdType];
fn worker_trace_link(&self) -> Option<&crate::protocols::common::preprocessor::TraceLink>;
}
impl HasTokenIds for BackendOutput {
fn token_ids(&self) -> &[TokenIdType] {
&self.token_ids
}
fn worker_trace_link(&self) -> Option<&crate::protocols::common::preprocessor::TraceLink> {
self.worker_trace_link.as_ref()
}
}
impl HasTokenIds for LLMEngineOutput {
fn token_ids(&self) -> &[TokenIdType] {
&self.token_ids
}
fn worker_trace_link(&self) -> Option<&crate::protocols::common::preprocessor::TraceLink> {
self.worker_trace_link.as_ref()
}
}
const NON_MIGRATABLE: &[ErrorType] = &[ErrorType::Cancelled, ErrorType::ResourceExhausted];
fn is_migratable(err: &(dyn StdError + 'static)) -> bool {
const MIGRATABLE: &[ErrorType] = &[
ErrorType::CannotConnect,
ErrorType::Disconnected,
ErrorType::ConnectionTimeout,
ErrorType::ResponseTimeout,
ErrorType::Backend(BackendError::EngineShutdown),
ErrorType::Backend(BackendError::StreamIncomplete),
ErrorType::WorkerOverloaded,
];
error::match_error_chain(err, MIGRATABLE, NON_MIGRATABLE)
}
fn is_migratable_for_request(
request: &PreprocessedRequest,
err: &(dyn StdError + 'static),
) -> bool {
if !is_migratable(err) {
return false;
}
let allows_phase = |phase| match explicit_target(request, phase) {
Ok(None) => true,
Ok(Some(target)) => {
tracing::debug!(
?phase,
worker_id = target.worker_id,
dp_rank = ?target.dp_rank,
"Migration disabled for explicitly pinned worker"
);
false
}
Err(error) => {
tracing::warn!(?phase, %error, "Migration disabled for invalid explicit worker target");
false
}
};
let Some(tracker) = request.tracker.as_ref() else {
return [
RequestPhase::Prefill,
RequestPhase::Decode,
RequestPhase::Aggregated,
]
.into_iter()
.all(allows_phase);
};
let phase = tracker.phase();
allows_phase(phase)
}
pub struct Migration {
migration_limit: u32,
max_seq_len: Option<u32>,
model_name: Arc<String>,
metrics: Arc<Metrics>,
}
impl Migration {
pub fn new(
migration_limit: u32,
max_seq_len: Option<u32>,
model_name: String,
metrics: Arc<Metrics>,
) -> Arc<Self> {
tracing::debug!(
"model {} migration limit {} max_seq_len {:?}",
model_name,
migration_limit,
max_seq_len
);
Arc::new(Self {
migration_limit,
max_seq_len,
model_name: Arc::new(model_name),
metrics,
})
}
pub fn from_mdc(
mdc: &ModelDeploymentCard,
migration_limit: u32,
max_seq_len: Option<u32>,
metrics: Arc<Metrics>,
) -> Arc<Self> {
Self::new(
migration_limit,
max_seq_len,
mdc.display_name.clone(),
metrics,
)
}
#[allow(clippy::type_complexity)]
pub(crate) fn into_operator_for<Resp>(
self: &Arc<Self>,
) -> Arc<
PipelineOperator<
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<Resp>>,
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<Resp>>,
>,
>
where
Resp: Data + HasTokenIds,
{
Operator::into_operator(self)
}
}
#[async_trait]
impl<Resp>
Operator<
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<Resp>>,
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<Resp>>,
> for Migration
where
Resp: Data + HasTokenIds,
{
async fn generate(
&self,
request: SingleIn<PreprocessedRequest>,
next: ServerStreamingEngine<PreprocessedRequest, Annotated<Resp>>,
) -> Result<ManyOut<Annotated<Resp>>> {
let (preprocessed_request, context) = request.transfer(());
let engine_ctx = context.context();
let engine_ctx_ = engine_ctx.clone();
let session_affinity = context
.get_optional::<SessionAffinityId>(SESSION_AFFINITY_CONTEXT_KEY)
.map_err(Error::msg)?
.map(|session_id| session_id.as_ref().clone());
let retry_manager = RetryManager::build(
engine_ctx,
context.metadata().clone(),
preprocessed_request,
next,
self.migration_limit,
self.max_seq_len,
self.model_name.clone(),
self.metrics.clone(),
session_affinity,
)
.await?;
let response_stream = stream::unfold(retry_manager, move |mut retry_manager| async move {
retry_manager
.next()
.await
.map(|response| (response, retry_manager))
})
.fuse();
Ok(ResponseStream::new(Box::pin(response_stream), engine_ctx_))
}
}
struct MigrationEvent {
migration_type: &'static str,
started_at: Instant,
}
impl MigrationEvent {
fn new(migration_type: &'static str) -> Self {
Self {
migration_type,
started_at: Instant::now(),
}
}
}
struct RetryManager<Resp>
where
Resp: Data + HasTokenIds,
{
context: Arc<dyn AsyncEngineContext>,
metadata: BTreeMap<String, String>,
request: PreprocessedRequest,
session_affinity: Option<SessionAffinityId>,
next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<Resp>>,
next_stream: Option<ManyOut<Annotated<Resp>>>,
retries_left: u64,
max_seq_len: Option<u32>,
model_name: Arc<String>,
metrics: Arc<Metrics>,
last_worker_link: Option<crate::protocols::common::preprocessor::TraceLink>,
active_route_trace: Option<Arc<RouteTraceContext>>,
next_attempt: u32,
completed_tokens: usize,
pending_migration: Option<MigrationCause>,
}
#[derive(Debug, Clone, Copy)]
struct MigrationCause {
reason: ErrorType,
from_worker_id: Option<u64>,
attempt: u32,
}
impl<Resp> RetryManager<Resp>
where
Resp: Data + HasTokenIds,
{
#[allow(clippy::too_many_arguments)]
pub async fn build(
context: Arc<dyn AsyncEngineContext>,
metadata: BTreeMap<String, String>,
mut preprocessed_request: PreprocessedRequest,
next: ServerStreamingEngine<PreprocessedRequest, Annotated<Resp>>,
mut retries_left: u32,
max_seq_len: Option<u32>,
model_name: Arc<String>,
metrics: Arc<Metrics>,
session_affinity: Option<SessionAffinityId>,
) -> Result<Self> {
if preprocessed_request
.sampling_options
.guided_decoding
.is_some()
{
if retries_left > 0 {
tracing::warn!(
"Guided-decoding request: migration disabled — FSM state is not transferable (applies to all backends)"
);
}
retries_left = 0;
}
if preprocessed_request.sampling_options.n.unwrap_or(1) > 1 {
if retries_left > 0 {
tracing::warn!(
"n>1 request: migration disabled - per-choice generation state is not transferable"
);
}
retries_left = 0;
}
if retries_left > 0 {
preprocessed_request.migration_state = Some(Default::default());
}
let mut slf = Self {
context,
metadata,
request: preprocessed_request,
session_affinity,
next_generate: next,
next_stream: None,
retries_left: u64::from(retries_left) + 1, max_seq_len,
model_name,
metrics,
last_worker_link: None,
active_route_trace: None,
next_attempt: 0,
completed_tokens: 0,
pending_migration: None,
};
slf.new_stream(None).await?;
slf.exceed_max_seq_len(0); Ok(slf)
}
pub async fn next(&mut self) -> Option<Annotated<Resp>> {
loop {
let response_stream = match self.next_stream.as_mut() {
Some(stream) => stream,
None => {
tracing::error!("next() called with next_stream is None - should not happen");
return Some(Annotated::from_error("next_stream is None"));
}
};
if let Some(response) = response_stream.next().await {
if let Some(err) = response.error.as_ref()
&& is_migratable_for_request(&self.request, err)
{
if self.retries_left == 0 {
let route_trace = self.active_route_trace.clone();
self.record_migration_exhausted(MigrationCause {
reason: err.error_type(),
from_worker_id: route_trace
.as_deref()
.and_then(RouteTraceContext::selected_worker_id),
attempt: self.failed_attempt(route_trace.as_deref()),
});
} else {
self.queue_migration(err.error_type(), self.active_route_trace.clone());
}
tracing::warn!(error = %err, "Stream disconnected, recreating stream");
self.metrics.inc_migration_ongoing_request(&self.model_name);
let migration_event =
MigrationEvent::new(frontend_service::migration_type::ONGOING_REQUEST);
if let Err(err) = self.new_stream(Some(migration_event)).await {
tracing::warn!(error = ?err, "Cannot recreate stream");
} else {
continue;
}
}
self.track_response(&response);
return Some(response);
}
return None;
}
}
async fn new_stream(&mut self, mut migration_event: Option<MigrationEvent>) -> Result<()> {
if self.retries_left == 0 {
if let Some(cause) = self.pending_migration.take() {
self.record_migration_exhausted(cause);
}
self.record_migration_outcome(
migration_event.as_ref(),
frontend_service::migration_outcome::FAILURE,
);
return Err(Error::msg("Migration limit exhausted"));
}
while self.retries_left > 0 {
self.retries_left -= 1;
if let Some(link) = self.last_worker_link.as_ref() {
self.request.migration_link = Some(link.clone());
}
let mut request = Context::with_id_and_metadata(
self.request.clone(),
self.context.id().to_string(),
self.metadata.clone(),
);
let migration = self.pending_migration.take();
let attempt = self.next_attempt;
self.next_attempt += 1;
let route_trace = attach_route_trace_context(
&mut request,
RouteTraceContext::new(
attempt,
migration.map(|migration| migration.reason),
migration.and_then(|migration| migration.from_worker_id),
self.completed_tokens,
),
);
if let Some(session_affinity) = self.session_affinity.as_ref() {
request.insert(SESSION_AFFINITY_CONTEXT_KEY, session_affinity.clone());
}
self.context.link_child(request.context());
if self.context.is_stopped() || self.context.is_killed() {
if let Some(cause) = migration {
tracing::info!(
target: "request_span",
{
{ "request.attempt" } = attempt,
{ "migration.is_retry" } = true,
{ "migration.reason" } = error_type_name(cause.reason),
{ "migration.from_worker_id" } = cause.from_worker_id,
{ "migration.tokens_completed" } = self.completed_tokens
},
"migration cancelled before worker dispatch"
);
} else {
tracing::info!(
target: "request_span",
{
{ "request.attempt" } = attempt,
{ "migration.is_retry" } = false
},
"request cancelled before worker dispatch"
);
}
self.record_migration_outcome(
migration_event.as_ref(),
frontend_service::migration_outcome::CANCELLED,
);
return Err(DynamoError::builder()
.error_type(ErrorType::Cancelled)
.message(format!(
"Context id {} is stopped or killed",
self.context.id()
))
.build()
.into());
}
let source_guards = self
.request
.multi_modal_data
.as_ref()
.into_iter()
.flat_map(|media| media.values())
.flatten()
.filter_map(|item| match item {
MultimodalData::Decoded(descriptor) => descriptor.source_storage.clone(),
_ => None,
})
.collect::<Vec<_>>();
if !source_guards.is_empty() {
attach_first_response_guard(&mut request, Arc::new(source_guards));
}
let response_stream = self.next_generate.generate(request).await;
match response_stream {
Ok(next_stream) => {
self.record_migration_outcome(
migration_event.as_ref(),
frontend_service::migration_outcome::SUCCESS,
);
self.active_route_trace = Some(route_trace);
self.next_stream = Some(next_stream);
return Ok(());
}
Err(err) if is_migratable_for_request(&self.request, err.as_ref()) => {
let reason = error_type_from_chain(err.as_ref());
if migration_event.is_none() {
migration_event = Some(MigrationEvent::new(
frontend_service::migration_type::NEW_REQUEST,
));
}
self.metrics.inc_migration_new_request(&self.model_name);
if self.retries_left == 0 {
let cause = MigrationCause {
reason,
from_worker_id: route_trace.selected_worker_id(),
attempt: route_trace.attempt(),
};
self.record_migration_exhausted(cause);
self.record_migration_outcome(
migration_event.as_ref(),
frontend_service::migration_outcome::FAILURE,
);
return Err(err);
}
self.queue_migration(reason, Some(route_trace));
tracing::warn!(error = %err, "Creating new stream, retrying");
}
Err(err) => {
let outcome =
if error::match_error_chain(err.as_ref(), &[ErrorType::Cancelled], &[]) {
frontend_service::migration_outcome::CANCELLED
} else {
frontend_service::migration_outcome::FAILURE
};
self.record_migration_outcome(migration_event.as_ref(), outcome);
return Err(err);
}
}
}
self.record_migration_outcome(
migration_event.as_ref(),
frontend_service::migration_outcome::FAILURE,
);
Err(Error::msg("Migration limit exhausted"))
}
fn failed_attempt(&self, route_trace: Option<&RouteTraceContext>) -> u32 {
route_trace.map_or_else(
|| self.next_attempt.saturating_sub(1),
RouteTraceContext::attempt,
)
}
fn queue_migration(&mut self, reason: ErrorType, route_trace: Option<Arc<RouteTraceContext>>) {
let from_worker_id = route_trace
.as_deref()
.and_then(RouteTraceContext::selected_worker_id);
self.pending_migration = Some(MigrationCause {
reason,
from_worker_id,
attempt: self.failed_attempt(route_trace.as_deref()),
});
tracing::info!(
target: "request_span",
{
{ "request.attempt" } = self.next_attempt,
{ "migration.is_retry" } = true,
{ "migration.reason" } = error_type_name(reason),
{ "migration.from_worker_id" } = from_worker_id,
{ "migration.tokens_completed" } = self.completed_tokens
},
"migration retry scheduled"
);
}
fn record_migration_exhausted(&self, cause: MigrationCause) {
tracing::warn!(
target: "request_span",
{
{ "request.attempt" } = cause.attempt,
{ "migration.is_retry" } = true,
{ "migration.reason" } = error_type_name(cause.reason),
{ "migration.from_worker_id" } = cause.from_worker_id,
{ "migration.tokens_completed" } = self.completed_tokens
},
"migration retries exhausted"
);
}
fn record_migration_outcome(&self, migration_event: Option<&MigrationEvent>, outcome: &str) {
if let Some(event) = migration_event {
self.metrics.observe_migration_duration(
&self.model_name,
event.migration_type,
outcome,
event.started_at.elapsed(),
);
}
}
fn track_response(&mut self, response: &Annotated<Resp>) {
let llm_engine_output = match response.data.as_ref() {
Some(output) => output,
None => return,
};
let token_ids = llm_engine_output.token_ids();
self.completed_tokens += token_ids.len();
if self.retries_left == 0 {
return;
}
if let Some(link) = llm_engine_output.worker_trace_link() {
self.last_worker_link = Some(link.clone());
}
let output_len = u32::try_from(token_ids.len()).unwrap_or(u32::MAX);
if self.exceed_max_seq_len(output_len) {
return;
}
if let Some(max_tokens) = self.request.stop_conditions.max_tokens {
self.request.stop_conditions.max_tokens = Some(max_tokens.saturating_sub(output_len));
}
if let Some(min_tokens) = self.request.stop_conditions.min_tokens {
self.request.stop_conditions.min_tokens = Some(min_tokens.saturating_sub(output_len));
}
for token_id in token_ids.iter() {
self.request.token_ids.push(*token_id);
}
}
fn exceed_max_seq_len(&mut self, new_output_len: u32) -> bool {
if let Some(max_seq_len) = self.max_seq_len {
let total_len = self.request.token_ids.len() as u32 + new_output_len;
if total_len > max_seq_len {
tracing::warn!(
"Sequence length {} exceeds migration max_seq_len {}, \
disabling migration",
total_len,
max_seq_len
);
self.metrics
.inc_migration_max_seq_len_exceeded(&self.model_name);
self.retries_left = 0; return true;
}
}
false
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::http::service::metrics::Metrics;
use crate::protocols::common::{
GuidedDecodingOptions, OutputOptions, SamplingOptions, StopConditions,
preprocessor::RoutingHints, timing::RequestTracker,
};
use dynamo_runtime::error::{DynamoError, ErrorType};
use dynamo_runtime::pipeline::AsyncEngine;
use dynamo_runtime::pipeline::context::Controller;
use dynamo_runtime::protocols::maybe_error::MaybeError;
use std::sync::atomic::{AtomicU32, Ordering};
use tokio::sync::mpsc;
const TEST_MODEL: &str = "test-model";
fn migration_duration_count(metrics: &Metrics, migration_type: &str, outcome: &str) -> u64 {
metrics.get_migration_duration_sample_count(TEST_MODEL, migration_type, outcome)
}
#[test]
fn stall_and_incomplete_stream_errors_are_migratable() {
let response_timeout = DynamoError::builder()
.error_type(ErrorType::ResponseTimeout)
.message("backend response inactivity timeout")
.build();
assert!(
is_migratable(&response_timeout),
"ResponseTimeout (stalled worker) must be migratable"
);
let stream_incomplete = DynamoError::builder()
.error_type(ErrorType::Backend(BackendError::StreamIncomplete))
.message("stream ended before completion")
.build();
assert!(
is_migratable(&stream_incomplete),
"StreamIncomplete (truncated stream from departed worker) must be migratable"
);
}
#[test]
fn pre_stream_failure_with_migration_sensitive_cause_is_still_migratable() {
use dynamo_runtime::pipeline::network::StreamPrologueError;
use dynamo_runtime::pipeline::network::egress::addressed_router::testing::{
migration_sensitive_error_types, pre_stream_failure_error,
};
for &error_type in migration_sensitive_error_types() {
let worker_error = DynamoError::builder()
.error_type(error_type)
.message("no capacity on the downstream worker")
.build();
assert!(!is_migratable(&worker_error), "{error_type:?} setup");
let err = pre_stream_failure_error(StreamPrologueError::new(
format!("Generate Error: {worker_error}"),
worker_error,
));
assert!(
is_migratable(&err),
"a {error_type:?} worker error must not make a pre-stream failure stop migrating"
);
}
let nested = DynamoError::builder()
.error_type(ErrorType::Backend(BackendError::InvalidArgument))
.message("downstream worker rejected the request")
.cause(
DynamoError::builder()
.error_type(ErrorType::ResourceExhausted)
.message("no capacity on the downstream worker")
.build(),
)
.build();
let err = pre_stream_failure_error(StreamPrologueError::new(
"Generate Error: downstream worker rejected the request",
nested,
));
assert!(
is_migratable(&err),
"a nested ResourceExhausted must not make a pre-stream failure stop migrating"
);
}
#[test]
fn migration_sensitive_types_match_the_exclusion_set() {
use dynamo_runtime::pipeline::network::egress::addressed_router::testing::migration_sensitive_error_types;
let router_types = migration_sensitive_error_types();
let missing_from_router: Vec<_> = NON_MIGRATABLE
.iter()
.filter(|t| !router_types.contains(t))
.collect();
let missing_from_here: Vec<_> = router_types
.iter()
.filter(|t| !NON_MIGRATABLE.contains(t))
.collect();
assert!(
missing_from_router.is_empty() && missing_from_here.is_empty(),
"NON_MIGRATABLE and MIGRATION_SENSITIVE_ERROR_TYPES must hold the same \
set: a type excluded from migration but still attachable as a \
pre-stream cause would stop that failure migrating. Missing from \
MIGRATION_SENSITIVE_ERROR_TYPES in addressed_router.rs: \
{missing_from_router:?}; missing from NON_MIGRATABLE here: \
{missing_from_here:?}"
);
}
#[test]
fn cancelled_and_exhausted_are_not_migratable() {
for et in [ErrorType::Cancelled, ErrorType::ResourceExhausted] {
let err = DynamoError::builder().error_type(et).message("x").build();
assert!(!is_migratable(&err), "{et:?} must not be migratable");
}
}
fn migratable_error(error_type: ErrorType) -> DynamoError {
DynamoError::builder()
.error_type(error_type)
.message("worker failed")
.build()
}
#[tokio::test]
async fn explicit_worker_pin_blocks_migration_for_worker_failures() {
let tracker = Arc::new(RequestTracker::new());
let mut request = create_mock_request(1);
request.tracker = Some(tracker.clone());
request.routing = Some(RoutingHints {
backend_instance_id: Some(7),
..Default::default()
});
for phase in [
RequestPhase::Aggregated,
RequestPhase::Prefill,
RequestPhase::Decode,
] {
let permit = tracker.set_phase(phase).await;
for error_type in [ErrorType::Disconnected, ErrorType::WorkerOverloaded] {
let error = migratable_error(error_type);
assert!(
!is_migratable_for_request(&request, &error),
"backend pin must block {error_type:?} migration during {phase:?}"
);
}
drop(permit);
}
}
#[tokio::test]
async fn phase_specific_pin_only_blocks_its_matching_phase() {
let error = migratable_error(ErrorType::WorkerOverloaded);
let prefill_tracker = Arc::new(RequestTracker::new());
let mut prefill_pinned = create_mock_request(1);
prefill_pinned.tracker = Some(prefill_tracker.clone());
prefill_pinned.routing = Some(RoutingHints {
prefill_worker_id: Some(11),
..Default::default()
});
let permit = prefill_tracker.set_phase(RequestPhase::Prefill).await;
assert!(!is_migratable_for_request(&prefill_pinned, &error));
drop(permit);
let permit = prefill_tracker.set_phase(RequestPhase::Decode).await;
assert!(is_migratable_for_request(&prefill_pinned, &error));
drop(permit);
let decode_tracker = Arc::new(RequestTracker::new());
let mut decode_pinned = create_mock_request(1);
decode_pinned.tracker = Some(decode_tracker.clone());
decode_pinned.routing = Some(RoutingHints {
decode_worker_id: Some(22),
..Default::default()
});
let permit = decode_tracker.set_phase(RequestPhase::Prefill).await;
assert!(is_migratable_for_request(&decode_pinned, &error));
drop(permit);
let permit = decode_tracker.set_phase(RequestPhase::Decode).await;
assert!(!is_migratable_for_request(&decode_pinned, &error));
drop(permit);
}
#[test]
fn trackerless_request_treats_any_explicit_worker_as_pinned() {
let error = migratable_error(ErrorType::WorkerOverloaded);
for routing in [
RoutingHints {
backend_instance_id: Some(7),
..Default::default()
},
RoutingHints {
prefill_worker_id: Some(11),
..Default::default()
},
RoutingHints {
decode_worker_id: Some(22),
..Default::default()
},
] {
let mut request = create_mock_request(1);
assert!(request.tracker.is_none());
request.routing = Some(routing);
assert!(!is_migratable_for_request(&request, &error));
}
let unpinned = create_mock_request(1);
assert!(is_migratable_for_request(&unpinned, &error));
}
fn create_mock_request(max_tokens: u32) -> PreprocessedRequest {
PreprocessedRequest::builder()
.model("mock".to_string())
.token_ids(vec![1, 2, 3])
.stop_conditions(StopConditions {
max_tokens: Some(max_tokens),
..Default::default()
})
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.eos_token_ids(vec![])
.annotations(vec![])
.build()
.unwrap()
}
fn create_mock_output(token_id: u32) -> Annotated<BackendOutput> {
Annotated::from_data(BackendOutput {
token_ids: vec![token_id],
tokens: vec![],
text: Some(format!("token_{token_id}")),
cum_log_probs: None,
log_probs: None,
top_logprobs: None,
finish_reason: None,
stop_reason: None,
index: None,
disaggregated_params: None,
encoder_result: None,
worker_trace_link: None,
completion_usage: None,
engine_data: None,
routing_data: None,
})
}
#[derive(Debug, Clone)]
enum MockBehavior {
Success,
FailThenSuccess,
FailThenSuccessWithAffinity,
FailThenCancel {
context: Arc<Controller>,
},
FailThenGenerateCancelled,
WorkerOverloadSequence {
worker_ids: Vec<u64>,
},
MidStreamFail {
fail_after: usize,
},
MidStreamFailAlways {
fail_after: usize,
},
MidStreamFailAlwaysStreamError {
fail_after: usize,
},
AlwaysFail,
}
struct MockEngine {
behavior: MockBehavior,
num_responses: usize,
token_offset: u32,
call_count: Arc<AtomicU32>,
context_id: String,
initial_min_tokens: Option<u32>,
}
impl MockEngine {
fn new(
behavior: MockBehavior,
num_responses: usize,
token_offset: u32,
context_id: String,
) -> Self {
Self {
behavior,
num_responses,
token_offset,
call_count: Arc::new(AtomicU32::new(0)),
context_id,
initial_min_tokens: None,
}
}
fn with_min_tokens(mut self, min_tokens: u32) -> Self {
self.initial_min_tokens = Some(min_tokens);
self
}
}
#[async_trait]
impl
AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<Annotated<BackendOutput>>, anyhow::Error>
for MockEngine
{
async fn generate(
&self,
request: SingleIn<PreprocessedRequest>,
) -> Result<ManyOut<Annotated<BackendOutput>>> {
let call_num = self.call_count.fetch_add(1, Ordering::SeqCst);
if matches!(self.behavior, MockBehavior::FailThenSuccessWithAffinity) {
let actual = request
.get::<SessionAffinityId>(SESSION_AFFINITY_CONTEXT_KEY)
.expect("session affinity context missing after migration wrapper");
assert_eq!(actual.as_str(), "session-123");
}
let (preprocessed_request, context) = request.transfer(());
assert_eq!(
context.id().to_string(),
self.context_id,
"Context ID mismatch"
);
let initial_tokens = 3; let responses_already_generated = preprocessed_request
.token_ids
.len()
.saturating_sub(initial_tokens);
let expected_max_tokens =
self.num_responses
.saturating_sub(responses_already_generated) as u32;
assert_eq!(
preprocessed_request.stop_conditions.max_tokens,
Some(expected_max_tokens),
"max_tokens should be {} but got {:?}",
expected_max_tokens,
preprocessed_request.stop_conditions.max_tokens
);
if let Some(initial_min_tokens) = self.initial_min_tokens {
let expected_min_tokens =
initial_min_tokens.saturating_sub(responses_already_generated as u32);
assert_eq!(
preprocessed_request.stop_conditions.min_tokens,
Some(expected_min_tokens),
"min_tokens should be rebased for each replacement request"
);
}
match &self.behavior {
MockBehavior::Success => {
self.send_responses(responses_already_generated, self.num_responses)
.await
}
MockBehavior::FailThenSuccess | MockBehavior::FailThenSuccessWithAffinity => {
if call_num == 0 {
return Err(anyhow::anyhow!(
DynamoError::builder()
.error_type(ErrorType::CannotConnect)
.message("no responders")
.build()
));
} else {
self.send_responses(responses_already_generated, self.num_responses)
.await
}
}
MockBehavior::FailThenCancel { context } => {
assert_eq!(call_num, 0, "cancelled retry must not reach the engine");
context.stop_generating();
Err(anyhow::anyhow!(
DynamoError::builder()
.error_type(ErrorType::CannotConnect)
.message("no responders")
.build()
))
}
MockBehavior::FailThenGenerateCancelled => {
let error_type = if call_num == 0 {
ErrorType::CannotConnect
} else {
ErrorType::Cancelled
};
Err(anyhow::anyhow!(
DynamoError::builder()
.error_type(error_type)
.message("request cancelled")
.build()
))
}
MockBehavior::WorkerOverloadSequence { worker_ids } => {
let excluded = preprocessed_request
.migration_state
.as_ref()
.expect("migration state missing")
.excluded_worker_ids();
assert_eq!(
excluded,
worker_ids[..call_num as usize],
"each retry must retain every previously rejected worker"
);
if let Some(&worker_id) = worker_ids.get(call_num as usize) {
let error = DynamoError::builder()
.error_type(ErrorType::WorkerOverloaded)
.message("selected worker is overloaded")
.build();
preprocessed_request
.migration_state
.as_ref()
.unwrap()
.record_failure(worker_id, Some(error.clone()));
return Err(anyhow::anyhow!(error));
}
self.send_responses(responses_already_generated, self.num_responses)
.await
}
MockBehavior::MidStreamFail { fail_after } => {
let (tx, rx) = mpsc::channel(1);
let token_offset = self.token_offset;
let fail_after = *fail_after;
let num_responses = self.num_responses;
if call_num == 0 {
tokio::spawn(async move {
for i in responses_already_generated..fail_after.min(num_responses) {
let response = create_mock_output(token_offset + 1 + i as u32);
if tx.send(response).await.is_err() {
break;
}
}
let error_response = Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("Stream ended before generation completed")
.build(),
);
let _ = tx.send(error_response).await;
});
} else {
tokio::spawn(async move {
for i in responses_already_generated..num_responses {
let response = create_mock_output(token_offset + 1 + i as u32);
if tx.send(response).await.is_err() {
break;
}
}
});
}
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let ctx = Arc::new(Controller::new(self.context_id.clone()));
Ok(dynamo_runtime::pipeline::ResponseStream::new(
Box::pin(stream),
ctx,
))
}
MockBehavior::MidStreamFailAlways { fail_after } => {
if call_num == 0 {
let (tx, rx) = mpsc::channel(1);
let token_offset = self.token_offset;
let fail_after = *fail_after;
let num_responses = self.num_responses;
tokio::spawn(async move {
for i in responses_already_generated..fail_after.min(num_responses) {
let response = create_mock_output(token_offset + 1 + i as u32);
if tx.send(response).await.is_err() {
break;
}
}
let error_response = Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("Stream ended before generation completed")
.build(),
);
let _ = tx.send(error_response).await;
});
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let ctx = Arc::new(Controller::new(self.context_id.clone()));
Ok(dynamo_runtime::pipeline::ResponseStream::new(
Box::pin(stream),
ctx,
))
} else {
Err(anyhow::anyhow!(
DynamoError::builder()
.error_type(ErrorType::CannotConnect)
.message("no responders")
.build()
))
}
}
MockBehavior::MidStreamFailAlwaysStreamError { fail_after } => {
let (tx, rx) = mpsc::channel(1);
let token_offset = self.token_offset;
let fail_after = *fail_after;
let num_responses = self.num_responses;
if call_num == 0 {
tokio::spawn(async move {
for i in responses_already_generated..fail_after.min(num_responses) {
let response = create_mock_output(token_offset + 1 + i as u32);
if tx.send(response).await.is_err() {
break;
}
}
let error_response = Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("Stream ended before generation completed")
.build(),
);
let _ = tx.send(error_response).await;
});
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let ctx = Arc::new(Controller::new(self.context_id.clone()));
Ok(dynamo_runtime::pipeline::ResponseStream::new(
Box::pin(stream),
ctx,
))
} else {
tokio::spawn(async move {
let error_response = Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("Stream ended before generation completed")
.build(),
);
let _ = tx.send(error_response).await;
});
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let ctx = Arc::new(Controller::new(self.context_id.clone()));
Ok(dynamo_runtime::pipeline::ResponseStream::new(
Box::pin(stream),
ctx,
))
}
}
MockBehavior::AlwaysFail => {
Err(anyhow::anyhow!(
DynamoError::builder()
.error_type(ErrorType::CannotConnect)
.message("no responders")
.build()
))
}
}
}
}
impl MockEngine {
async fn send_responses(
&self,
start: usize,
end: usize,
) -> Result<ManyOut<Annotated<BackendOutput>>> {
let (tx, rx) = mpsc::channel(1);
let token_offset = self.token_offset;
tokio::spawn(async move {
for i in start..end {
let response = create_mock_output(token_offset + 1 + i as u32);
if tx.send(response).await.is_err() {
break;
}
}
});
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let ctx = Arc::new(Controller::new(self.context_id.clone()));
Ok(dynamo_runtime::pipeline::ResponseStream::new(
Box::pin(stream),
ctx,
))
}
}
#[test]
fn worker_overload_migrates_but_pool_exhaustion_does_not() {
let worker_busy = DynamoError::builder()
.error_type(ErrorType::WorkerOverloaded)
.message("Selected worker is overloaded, please retry later")
.build();
assert!(
is_migratable(&worker_busy),
"one overloaded worker must fail over to another"
);
let pool_exhausted = DynamoError::builder()
.error_type(ErrorType::ResourceExhausted)
.message("All workers are busy, please retry later")
.build();
assert!(
!is_migratable(&pool_exhausted),
"pool-wide exhaustion must not migrate; no worker has room"
);
}
#[tokio::test]
async fn test_retry_manager_no_migration() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::Success,
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
0,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 10);
for (i, response) in responses.iter().enumerate() {
assert!(response.err().is_none());
if let Some(output) = &response.data {
assert_eq!(output.token_ids, vec![101 + i as u32]); }
}
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 0);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 0);
}
#[tokio::test]
async fn maximum_migration_limit_does_not_overflow() {
let context_id = uuid::Uuid::new_v4().to_string();
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::Success,
1,
100,
context_id.clone(),
));
let calls = mock_engine.call_count.clone();
let request =
Context::with_id_and_metadata(create_mock_request(1), context_id, BTreeMap::new());
let migration = Migration::new(
u32::MAX,
None,
TEST_MODEL.to_string(),
Arc::new(Metrics::new()),
);
let responses = migration
.generate(request, mock_engine)
.await
.unwrap()
.collect::<Vec<_>>()
.await;
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(responses.len(), 1);
assert!(responses[0].error.is_none());
}
#[tokio::test]
async fn test_migration_preserves_session_affinity_across_retry() {
let context_id = uuid::Uuid::new_v4().to_string();
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::FailThenSuccessWithAffinity,
1,
100,
context_id.clone(),
));
let calls = mock_engine.call_count.clone();
let mut request =
Context::with_id_and_metadata(create_mock_request(1), context_id, BTreeMap::new());
request.insert(
SESSION_AFFINITY_CONTEXT_KEY,
SessionAffinityId::new("session-123"),
);
let migration = Migration::new(1, None, TEST_MODEL.to_string(), Arc::new(Metrics::new()));
let mut stream = migration.generate(request, mock_engine).await.unwrap();
while stream.next().await.is_some() {}
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn explicit_backend_pin_does_not_retry_the_same_worker() {
let context_id = uuid::Uuid::new_v4().to_string();
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::FailThenSuccess,
1,
100,
context_id.clone(),
));
let calls = mock_engine.call_count.clone();
let mut content = create_mock_request(1);
content.routing = Some(RoutingHints {
backend_instance_id: Some(7),
..Default::default()
});
let request = Context::with_id_and_metadata(content, context_id, BTreeMap::new());
let migration = Migration::new(3, None, TEST_MODEL.to_string(), Arc::new(Metrics::new()));
let result = migration.generate(request, mock_engine).await;
assert!(result.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn worker_overload_excludes_the_rejected_worker_on_retry() {
let context_id = uuid::Uuid::new_v4().to_string();
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::WorkerOverloadSequence {
worker_ids: vec![7, 8],
},
1,
100,
context_id.clone(),
));
let calls = mock_engine.call_count.clone();
let request =
Context::with_id_and_metadata(create_mock_request(1), context_id, BTreeMap::new());
let migration = Migration::new(2, None, TEST_MODEL.to_string(), Arc::new(Metrics::new()));
let mut stream = migration.generate(request, mock_engine).await.unwrap();
let responses = stream.by_ref().collect::<Vec<_>>().await;
assert_eq!(calls.load(Ordering::SeqCst), 3);
assert_eq!(responses.len(), 1);
assert!(responses[0].error.is_none());
}
#[tokio::test]
async fn test_retry_manager_new_request_migration() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::FailThenSuccess,
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 10);
for (i, response) in responses.iter().enumerate() {
assert!(response.err().is_none());
if let Some(output) = &response.data {
assert_eq!(output.token_ids, vec![101 + i as u32]); }
}
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 1);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 0);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::NEW_REQUEST,
frontend_service::migration_outcome::SUCCESS,
),
1
);
}
#[tokio::test]
async fn test_retry_manager_ongoing_request_migration() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFail { fail_after: 5 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 10);
for (i, response) in responses.iter().enumerate() {
assert!(response.err().is_none());
if let Some(output) = &response.data {
assert_eq!(output.token_ids, vec![101 + i as u32]); }
}
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 0);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 1);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::ONGOING_REQUEST,
frontend_service::migration_outcome::SUCCESS,
),
1
);
}
#[tokio::test]
async fn retry_rebases_min_tokens_from_the_delivered_prefix() {
let context_id = uuid::Uuid::new_v4().to_string();
let mut request = create_mock_request(10);
request.stop_conditions.min_tokens = Some(7);
let mock_engine = Arc::new(
MockEngine::new(
MockBehavior::MidStreamFail { fail_after: 3 },
10,
100,
context_id.clone(),
)
.with_min_tokens(7),
);
let calls = mock_engine.call_count.clone();
let request = Context::with_id_and_metadata(request, context_id, BTreeMap::new());
let migration = Migration::new(1, None, TEST_MODEL.to_string(), Arc::new(Metrics::new()));
let responses = migration
.generate(request, mock_engine)
.await
.unwrap()
.collect::<Vec<_>>()
.await;
assert_eq!(calls.load(Ordering::SeqCst), 2);
assert_eq!(responses.len(), 10);
assert!(responses.iter().all(|response| response.error.is_none()));
}
#[tokio::test]
async fn test_retry_manager_new_request_migration_indefinite_failure() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(0);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::AlwaysFail,
0,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let retry_manager_result = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await;
assert!(retry_manager_result.is_err());
if let Err(error) = retry_manager_result {
assert!(error.to_string().contains("no responders"));
}
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 4);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 0);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::NEW_REQUEST,
frontend_service::migration_outcome::FAILURE,
),
1
);
}
#[tokio::test]
async fn test_retry_manager_ongoing_request_migration_indefinite_failure() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFailAlways { fail_after: 3 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
) .await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 4);
for (i, response) in responses[0..3].iter().enumerate() {
assert!(response.err().is_none());
if let Some(output) = &response.data {
assert_eq!(output.token_ids, vec![101 + i as u32]); }
}
let error_response = &responses[3];
let err = error_response.err().expect("expected error response");
assert_eq!(err.error_type(), ErrorType::Disconnected);
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 3);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 1);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::ONGOING_REQUEST,
frontend_service::migration_outcome::FAILURE,
),
1
);
}
#[tokio::test]
async fn test_retry_manager_ongoing_request_migration_indefinite_failure_stream_error() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFailAlwaysStreamError { fail_after: 3 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
) .await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 4);
for (i, response) in responses[0..3].iter().enumerate() {
assert!(response.err().is_none());
if let Some(output) = &response.data {
assert_eq!(output.token_ids, vec![101 + i as u32]); }
}
let error_response = &responses[3];
let err = error_response.err().expect("expected error response");
assert_eq!(err.error_type(), ErrorType::Disconnected);
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 0);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 4); assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::ONGOING_REQUEST,
frontend_service::migration_outcome::SUCCESS,
),
3
);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::ONGOING_REQUEST,
frontend_service::migration_outcome::FAILURE,
),
1
);
}
#[tokio::test]
async fn test_retry_manager_context_stopped_before_stream() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::Success,
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
ctx.stop_generating();
let metrics = Arc::new(Metrics::new());
let retry_manager_result = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await;
assert!(retry_manager_result.is_err());
if let Err(error) = retry_manager_result {
assert!(
error
.to_string()
.contains(&format!("Context id {} is stopped or killed", context_id))
);
let dynamo_err = error
.downcast_ref::<DynamoError>()
.expect("Error should be a DynamoError");
assert_eq!(
dynamo_err.error_type(),
ErrorType::Cancelled,
"Stopped/killed context should produce a Cancelled error"
);
}
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 0);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 0);
}
#[tokio::test]
async fn test_retry_manager_cancelled_during_migration() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let ctx = Arc::new(Controller::new(context_id.clone()));
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::FailThenCancel {
context: ctx.clone(),
},
10,
100,
context_id,
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let metrics = Arc::new(Metrics::new());
let result = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await;
let error = match result {
Ok(_) => panic!("cancelled migration must fail"),
Err(error) => error,
};
let dynamo_error = error
.downcast_ref::<DynamoError>()
.expect("error should be a DynamoError");
assert_eq!(dynamo_error.error_type(), ErrorType::Cancelled);
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 1);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::NEW_REQUEST,
frontend_service::migration_outcome::CANCELLED,
),
1
);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::NEW_REQUEST,
frontend_service::migration_outcome::FAILURE,
),
0
);
}
#[tokio::test]
async fn test_retry_manager_generate_cancelled_during_migration() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::FailThenGenerateCancelled,
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let metrics = Arc::new(Metrics::new());
let result = RetryManager::build(
Arc::new(Controller::new(context_id)),
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await;
let error = match result {
Ok(_) => panic!("cancelled migration must fail"),
Err(error) => error,
};
let dynamo_error = error
.downcast_ref::<DynamoError>()
.expect("error should be a DynamoError");
assert_eq!(dynamo_error.error_type(), ErrorType::Cancelled);
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 1);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::NEW_REQUEST,
frontend_service::migration_outcome::CANCELLED,
),
1
);
assert_eq!(
migration_duration_count(
&metrics,
frontend_service::migration_type::NEW_REQUEST,
frontend_service::migration_outcome::FAILURE,
),
0
);
}
#[tokio::test]
async fn test_retry_manager_no_migration_for_guided_decoding() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let mut request = create_mock_request(10);
request.sampling_options.guided_decoding = Some(GuidedDecodingOptions::new(
Some(serde_json::json!({"type": "object", "properties": {"name": {"type": "string"}}})),
None,
None,
None,
None,
None,
None,
));
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFail { fail_after: 3 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3, None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(
responses.len(),
4,
"Expected 3 successful + 1 error response (migration must be blocked for \
guided-decoding), but got {} responses",
responses.len()
);
for (i, response) in responses[0..3].iter().enumerate() {
assert!(
response.err().is_none(),
"Response {} should be successful",
i
);
}
let last = responses.last().unwrap();
let err = last
.err()
.expect("Last response should be a Disconnected error");
assert_eq!(
err.error_type(),
ErrorType::Disconnected,
"Error type should be Disconnected"
);
}
#[tokio::test]
async fn test_retry_manager_max_seq_len_exceeded() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFail { fail_after: 3 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
Some(5), Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(
responses.len(),
4,
"Expected 3 successful + 1 error (migration disabled by max_seq_len)"
);
for (i, response) in responses[0..3].iter().enumerate() {
assert!(response.err().is_none(), "Response {} should be OK", i);
}
let err = responses[3]
.err()
.expect("Last response should be Disconnected error");
assert_eq!(err.error_type(), ErrorType::Disconnected);
assert_eq!(metrics.get_migration_new_request_count(TEST_MODEL), 0);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 1);
assert_eq!(
metrics.get_migration_max_seq_len_exceeded_count(TEST_MODEL),
1
);
}
#[tokio::test]
async fn test_retry_manager_max_seq_len_at_limit() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFail { fail_after: 2 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
Some(5), Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 10);
for response in &responses {
assert!(response.err().is_none());
}
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 1);
assert_eq!(
retry_manager.request.token_ids.len(),
5,
"tracked token_ids should be exactly max_seq_len"
);
assert_eq!(
metrics.get_migration_max_seq_len_exceeded_count(TEST_MODEL),
1
);
}
#[tokio::test]
async fn test_retry_manager_max_seq_len_exceeded_by_prompt() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(10);
let mock_engine = Arc::new(MockEngine::new(
MockBehavior::MidStreamFail { fail_after: 3 },
10,
100,
context_id.clone(),
));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
mock_engine;
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
Some(2), Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(
responses.len(),
4,
"Expected 3 successful + 1 error (migration disabled by prompt exceeding max_seq_len)"
);
for (i, response) in responses[0..3].iter().enumerate() {
assert!(response.err().is_none(), "Response {} should be OK", i);
}
let err = responses[3]
.err()
.expect("Last response should be Disconnected error");
assert_eq!(err.error_type(), ErrorType::Disconnected);
assert_eq!(
metrics.get_migration_max_seq_len_exceeded_count(TEST_MODEL),
1
);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 1);
}
#[tokio::test]
async fn test_retry_manager_generic_over_llm_engine_output() {
dynamo_runtime::logging::init();
let context_id = uuid::Uuid::new_v4().to_string();
struct LlmEngineMock(String);
#[async_trait]
impl
AsyncEngine<
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<LLMEngineOutput>>,
anyhow::Error,
> for LlmEngineMock
{
async fn generate(
&self,
_request: SingleIn<PreprocessedRequest>,
) -> Result<ManyOut<Annotated<LLMEngineOutput>>> {
let responses = stream::iter((0..3u32).map(|i| {
Annotated::from_data(LLMEngineOutput {
token_ids: vec![200 + i],
..Default::default()
})
}));
let ctx = Arc::new(Controller::new(self.0.clone()));
Ok(ResponseStream::new(Box::pin(responses), ctx))
}
}
let request = create_mock_request(3);
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<LLMEngineOutput>> =
Arc::new(LlmEngineMock(context_id.clone()));
let ctx = Arc::new(Controller::new(context_id.clone()));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
1,
None,
Arc::new(TEST_MODEL.to_string()),
metrics,
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(r) = retry_manager.next().await {
responses.push(r);
}
assert_eq!(responses.len(), 3);
assert_eq!(
retry_manager.request.token_ids,
vec![1, 2, 3, 200, 201, 202]
);
}
#[tokio::test]
async fn test_retry_manager_cancellation_during_migration_skips_retry_dispatch() {
dynamo_runtime::logging::init();
struct CancelBeforeRetryEngine {
calls: Arc<AtomicU32>,
root: Arc<Controller>,
context_id: String,
}
#[async_trait]
impl
AsyncEngine<
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<BackendOutput>>,
anyhow::Error,
> for CancelBeforeRetryEngine
{
async fn generate(
&self,
_request: SingleIn<PreprocessedRequest>,
) -> Result<ManyOut<Annotated<BackendOutput>>> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
assert_eq!(call, 0, "cancelled migration must not dispatch a retry");
let root = self.root.clone();
let responses = async_stream::stream! {
yield create_mock_output(101);
root.stop();
yield Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("worker disconnected")
.build(),
);
};
Ok(ResponseStream::new(
Box::pin(responses),
Arc::new(Controller::new(self.context_id.clone())),
))
}
}
let context_id = uuid::Uuid::new_v4().to_string();
let root = Arc::new(Controller::new(context_id.clone()));
let calls = Arc::new(AtomicU32::new(0));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
Arc::new(CancelBeforeRetryEngine {
calls: calls.clone(),
root: root.clone(),
context_id,
});
let mut retry_manager = RetryManager::build(
root,
BTreeMap::new(),
create_mock_request(5),
next_generate,
2,
None,
Arc::new(TEST_MODEL.to_string()),
Arc::new(Metrics::new()),
None,
)
.await
.expect("initial stream should be created");
assert!(retry_manager.next().await.unwrap().err().is_none());
let failure = retry_manager
.next()
.await
.expect("disconnect should be returned when retry is cancelled")
.err()
.expect("second response should be the original disconnect");
assert_eq!(failure.error_type(), ErrorType::Disconnected);
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(retry_manager.completed_tokens, 1);
assert_eq!(retry_manager.next_attempt, 2);
assert!(retry_manager.pending_migration.is_none());
}
#[tokio::test]
async fn test_retry_manager_propagates_migration_link_over_two_hops() {
use crate::protocols::common::preprocessor::TraceLink;
use dynamo_runtime::pipeline::network::egress::route_span::get_route_trace_context;
use std::sync::Mutex;
dynamo_runtime::logging::init();
type CapturedRoute = (u32, Option<ErrorType>, Option<u64>, usize);
struct LinkingMockEngine {
captured_links: Arc<Mutex<Vec<Option<TraceLink>>>>,
captured_routes: Arc<Mutex<Vec<CapturedRoute>>>,
worker_links: Vec<TraceLink>,
context_id: String,
call_count: Arc<AtomicU32>,
}
#[async_trait]
impl
AsyncEngine<
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<BackendOutput>>,
anyhow::Error,
> for LinkingMockEngine
{
async fn generate(
&self,
request: SingleIn<PreprocessedRequest>,
) -> Result<ManyOut<Annotated<BackendOutput>>> {
let call_num = self.call_count.fetch_add(1, Ordering::SeqCst) as usize;
let route_trace = get_route_trace_context(&request)
.expect("migration wrapper must attach route trace context");
route_trace.set_selected_worker_id(100 + call_num as u64);
self.captured_routes.lock().unwrap().push((
route_trace.attempt(),
route_trace.migration_reason(),
route_trace.from_worker_id(),
route_trace.tokens_completed(),
));
let (preprocessed_request, _ctx) = request.transfer(());
self.captured_links
.lock()
.unwrap()
.push(preprocessed_request.migration_link.clone());
let (tx, rx) = mpsc::channel(1);
let context_id = self.context_id.clone();
let fail_this_call = call_num < 2;
let link = self.worker_links.get(call_num).cloned();
let responses_already_generated =
preprocessed_request.token_ids.len().saturating_sub(3);
let total_chunks: usize = 6;
tokio::spawn(async move {
let start = responses_already_generated;
let end = if fail_this_call {
(start + 2).min(total_chunks)
} else {
total_chunks
};
for i in start..end {
let mut out = create_mock_output(100 + 1 + i as u32);
if i == start
&& let (Some(link), Some(data)) = (&link, out.data.as_mut())
{
data.worker_trace_link = Some(link.clone());
}
if tx.send(out).await.is_err() {
return;
}
}
if fail_this_call {
let err = Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("Stream ended before generation completed")
.build(),
);
let _ = tx.send(err).await;
}
});
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let ctx = Arc::new(Controller::new(context_id));
Ok(dynamo_runtime::pipeline::ResponseStream::new(
Box::pin(stream),
ctx,
))
}
}
let context_id = uuid::Uuid::new_v4().to_string();
let request = create_mock_request(6);
let captured = Arc::new(Mutex::new(Vec::<Option<TraceLink>>::new()));
let captured_routes = Arc::new(Mutex::new(Vec::new()));
let link_a = TraceLink {
trace_id: "0123456789abcdef0123456789abcdef".to_string(),
span_id: "aaaaaaaaaaaaaaaa".to_string(),
};
let link_b = TraceLink {
trace_id: "0123456789abcdef0123456789abcdef".to_string(),
span_id: "bbbbbbbbbbbbbbbb".to_string(),
};
let engine = Arc::new(LinkingMockEngine {
captured_links: captured.clone(),
captured_routes: captured_routes.clone(),
worker_links: vec![link_a.clone(), link_b.clone()],
context_id: context_id.clone(),
call_count: Arc::new(AtomicU32::new(0)),
});
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
engine;
let ctx = Arc::new(Controller::new(context_id));
let metrics = Arc::new(Metrics::new());
let mut retry_manager = RetryManager::build(
ctx,
BTreeMap::new(),
request,
next_generate,
3,
None,
Arc::new(TEST_MODEL.to_string()),
metrics.clone(),
None,
)
.await
.expect("Failed to build RetryManager");
let mut responses = Vec::new();
while let Some(response) = retry_manager.next().await {
responses.push(response);
}
assert_eq!(responses.len(), 6, "expected all 6 chunks across 2 hops");
for response in &responses {
assert!(response.err().is_none(), "no chunk should be an error");
}
let links = captured.lock().unwrap();
assert_eq!(
links.len(),
3,
"engine.generate must be called 3 times for a 2-hop migration"
);
assert!(
links[0].is_none(),
"first attempt has no predecessor — migration_link must be None"
);
assert_eq!(
links[1].as_ref(),
Some(&link_a),
"second attempt must link back to worker A"
);
assert_eq!(
links[2].as_ref(),
Some(&link_b),
"third attempt must link back to worker B (latest worker, not original)"
);
drop(links);
assert_eq!(
*captured_routes.lock().unwrap(),
vec![
(0, None, None, 0),
(1, Some(ErrorType::Disconnected), Some(100), 2),
(2, Some(ErrorType::Disconnected), Some(101), 4),
],
"each retry must carry the failed worker, reason, and delivered-token count"
);
assert_eq!(metrics.get_migration_ongoing_request_count(TEST_MODEL), 2);
}
#[tokio::test]
async fn test_migration_exhausted_reports_the_attempt_that_failed() {
use std::sync::Mutex;
use tracing::field::{Field, Visit};
use tracing_subscriber::layer::{Context as LayerContext, Layer, SubscriberExt};
#[derive(Default)]
struct Captured {
exhausted_attempt: Option<u64>,
exhausted_tokens: Option<u64>,
scheduled_attempts: Vec<u64>,
}
struct AttemptVisitor {
message: Option<String>,
attempt: Option<u64>,
tokens: Option<u64>,
}
impl Visit for AttemptVisitor {
fn record_u64(&mut self, field: &Field, value: u64) {
match field.name() {
"request.attempt" => self.attempt = Some(value),
"migration.tokens_completed" => self.tokens = Some(value),
_ => {}
}
}
fn record_i64(&mut self, field: &Field, value: i64) {
match field.name() {
"request.attempt" => self.attempt = Some(value as u64),
"migration.tokens_completed" => self.tokens = Some(value as u64),
_ => {}
}
}
fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
if field.name() == "message" {
self.message = Some(format!("{value:?}"));
}
}
}
struct CaptureLayer(Arc<Mutex<Captured>>);
impl<S: tracing::Subscriber> Layer<S> for CaptureLayer {
fn on_event(&self, event: &tracing::Event<'_>, _ctx: LayerContext<'_, S>) {
let mut visitor = AttemptVisitor {
message: None,
attempt: None,
tokens: None,
};
event.record(&mut visitor);
let (Some(message), Some(attempt)) = (visitor.message, visitor.attempt) else {
return;
};
let mut captured = self.0.lock().unwrap();
if message.contains("migration retries exhausted") {
captured.exhausted_attempt = Some(attempt);
captured.exhausted_tokens = visitor.tokens;
} else if message.contains("migration retry scheduled") {
captured.scheduled_attempts.push(attempt);
}
}
}
struct AlwaysDisconnectEngine {
context_id: String,
}
#[async_trait]
impl
AsyncEngine<
SingleIn<PreprocessedRequest>,
ManyOut<Annotated<BackendOutput>>,
anyhow::Error,
> for AlwaysDisconnectEngine
{
async fn generate(
&self,
_request: SingleIn<PreprocessedRequest>,
) -> Result<ManyOut<Annotated<BackendOutput>>> {
let responses = async_stream::stream! {
yield create_mock_output(101);
yield Annotated::from_err(
DynamoError::builder()
.error_type(ErrorType::Disconnected)
.message("worker disconnected")
.build(),
);
};
Ok(ResponseStream::new(
Box::pin(responses),
Arc::new(Controller::new(self.context_id.clone())),
))
}
}
let captured = Arc::new(Mutex::new(Captured::default()));
let subscriber = tracing_subscriber::registry().with(CaptureLayer(Arc::clone(&captured)));
let context_id = uuid::Uuid::new_v4().to_string();
let root = Arc::new(Controller::new(context_id.clone()));
let next_generate: ServerStreamingEngine<PreprocessedRequest, Annotated<BackendOutput>> =
Arc::new(AlwaysDisconnectEngine {
context_id: context_id.clone(),
});
let retries = 1;
let manager = tracing::subscriber::with_default(subscriber, || {
futures::executor::block_on(async {
let mut manager = RetryManager::build(
root,
BTreeMap::new(),
create_mock_request(50),
next_generate,
retries,
None,
Arc::new(TEST_MODEL.to_string()),
Arc::new(Metrics::new()),
None,
)
.await
.expect("initial stream should be created");
while let Some(response) = manager.next().await {
if response.err().is_some() {
break;
}
}
manager
})
});
let captured = captured.lock().unwrap();
let last_dispatched = manager.next_attempt - 1;
assert_eq!(
captured.exhausted_attempt,
Some(u64::from(last_dispatched)),
"exhaustion must name the attempt that failed ({last_dispatched}), \
not the retry that never ran; scheduled={:?}",
captured.scheduled_attempts
);
assert_eq!(
captured.exhausted_tokens,
Some(2),
"exhaustion must count the token from the final allowed attempt, \
not just the attempts that had retries left"
);
assert_eq!(
captured.scheduled_attempts,
vec![u64::from(last_dispatched)],
"retry scheduled must name the upcoming attempt"
);
}
}