use std::collections::HashSet;
use std::sync::{Arc, Mutex, OnceLock};
use derive_builder::Builder;
use dynamo_kv_router::{
config::RouterConfigOverride,
protocols::{BlockExtraInfo, RoutingConstraints, WorkerId},
router_hint::{ROUTER_HINT_EXTRA_ARGS_KEY, RouterHint},
};
use dynamo_runtime::error::{DynamoError, ErrorType, match_error_chain};
use serde::{Deserialize, Serialize};
const KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY: &str = "kv_transfer_params";
use uuid::Uuid;
use super::extensions::{AgentContext, RouterParams};
use super::timing::RequestTracker;
use super::{OutputOptions, SamplingOptions, StopConditions};
use crate::preprocessor::media::RdmaMediaDataDescriptor;
use crate::protocols::TokenIdType;
#[derive(Serialize, Deserialize, Debug, Clone, Default, Builder)]
#[builder(default)]
pub struct RoutingHints {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend_instance_id: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prefill_worker_id: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub decode_worker_id: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dp_rank: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prefill_dp_rank: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub expected_output_tokens: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub lora_name: Option<String>,
#[serde(
default,
rename = "cache_salt",
skip_serializing_if = "Option::is_none"
)]
pub cache_namespace: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub priority_jump: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub strict_priority: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub priority: Option<i32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allowed_worker_ids: Option<HashSet<WorkerId>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub routing_constraints: Option<RoutingConstraints>,
}
#[derive(Serialize, Deserialize, Debug, Clone, Default)]
pub struct BootstrapInfo {
pub bootstrap_host: String,
pub bootstrap_port: u16,
pub bootstrap_room: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub handoff_id: Option<Uuid>,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)]
pub struct TraceLink {
pub trace_id: String,
pub span_id: String,
}
#[derive(Debug, Clone, Default)]
pub(crate) struct MigrationState {
inner: Arc<OnceLock<Mutex<MigrationStateInner>>>,
}
#[derive(Debug, Default)]
struct MigrationStateInner {
excluded_worker_ids: Vec<WorkerId>,
last_error: Option<DynamoError>,
}
impl MigrationState {
pub(crate) fn record_failure(&self, worker_id: WorkerId, error: Option<DynamoError>) {
let mut inner = self
.inner
.get_or_init(Default::default)
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !inner.excluded_worker_ids.contains(&worker_id) {
inner.excluded_worker_ids.push(worker_id);
}
if error.is_some() {
inner.last_error = error;
}
}
pub(crate) fn excluded_worker_ids(&self) -> Vec<WorkerId> {
let Some(inner) = self.inner.get() else {
return Vec::new();
};
inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.excluded_worker_ids
.clone()
}
pub(crate) fn exhausted_error(&self) -> Option<DynamoError> {
let inner = self.inner.get()?;
let last_error = inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.last_error
.clone()?;
let (error_type, message) = if match_error_chain(
&last_error,
&[ErrorType::WorkerOverloaded],
&[ErrorType::ResourceExhausted],
) {
(
ErrorType::ResourceExhausted,
"all eligible workers rejected the request as overloaded",
)
} else {
(
ErrorType::Unavailable,
"no untried eligible worker remains after migration",
)
};
Some(
DynamoError::builder()
.error_type(error_type)
.message(message)
.build(),
)
}
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct PrefillResult {
pub disaggregated_params: serde_json::Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt_tokens_details: Option<dynamo_protocols::types::PromptTokensDetails>,
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, Builder)]
#[builder(default)]
pub struct MmRoutingInfo {
pub routing_token_ids: Vec<TokenIdType>,
pub block_mm_infos: Vec<Option<BlockExtraInfo>>,
#[serde(default)]
pub expanded_prompt_len: usize,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub enum MultimodalData {
Url(url::Url),
#[serde(rename(serialize = "Url"))]
RawUrl(String),
Decoded(RdmaMediaDataDescriptor),
UuidOnly(String),
}
pub type MultimodalDataMap = std::collections::HashMap<String, Vec<MultimodalData>>;
pub type MultimodalUuidMap = std::collections::HashMap<String, Vec<Option<String>>>;
#[derive(Serialize, Deserialize, Debug, Clone, Builder)]
pub struct PreprocessedRequest {
#[serde(default)]
pub model: String,
#[builder(default)]
#[serde(skip)]
pub(crate) migration_state: Option<MigrationState>,
#[builder(default)]
#[serde(skip)]
pub(crate) staged_kv_cleanup: bool,
pub token_ids: Vec<TokenIdType>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt_embeds: Option<String>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multi_modal_data: Option<MultimodalDataMap>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multi_modal_uuids: Option<MultimodalUuidMap>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mm_routing_info: Option<MmRoutingInfo>,
#[serde(default)]
pub stop_conditions: StopConditions,
#[serde(default)]
pub sampling_options: SamplingOptions,
#[serde(default)]
pub output_options: OutputOptions,
#[builder(default)]
#[serde(default)]
pub eos_token_ids: Vec<TokenIdType>,
#[builder(default)]
pub mdc_sum: Option<String>,
#[builder(default)]
#[serde(default)]
pub annotations: Vec<String>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub routing: Option<RoutingHints>,
#[builder(default)]
pub router_config_override: Option<RouterConfigOverride>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prefill_result: Option<PrefillResult>,
#[builder(default)]
#[serde(
default,
skip_serializing_if = "Option::is_none",
deserialize_with = "deserialize_optional_object"
)]
pub encoder_result: Option<serde_json::Value>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub migration_link: Option<TraceLink>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub bootstrap_info: Option<BootstrapInfo>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub extra_args: Option<serde_json::Value>,
#[builder(default)]
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub require_reasoning: bool,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub router: Option<RouterParams>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub agent_context: Option<AgentContext>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mm_processor_kwargs: Option<serde_json::Value>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub media_io_kwargs: Option<serde_json::Value>,
#[builder(default)]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_timestamp_ms: Option<f64>,
#[builder(default)]
#[serde(skip)]
pub tracker: Option<Arc<RequestTracker>>,
#[builder(default)]
#[serde(
default,
rename = "_HEALTH_CHECK",
skip_serializing_if = "std::ops::Not::not"
)]
pub is_probe: bool,
}
fn deserialize_optional_object<'de, D>(
deserializer: D,
) -> Result<Option<serde_json::Value>, D::Error>
where
D: serde::Deserializer<'de>,
{
let value = Option::<serde_json::Value>::deserialize(deserializer)?;
if let Some(v) = &value
&& !v.is_object()
{
return Err(serde::de::Error::custom(
"encoder_result must be a JSON object",
));
}
Ok(value)
}
impl PreprocessedRequest {
pub fn has_annotation(&self, annotation: &str) -> bool {
self.annotations.contains(&annotation.to_string())
}
pub fn get_annotation_value(&self, key: &str) -> Option<String> {
let prefix = format!("{}:", key);
self.annotations
.iter()
.find(|a| a.starts_with(&prefix))
.map(|a| a[prefix.len()..].to_string())
}
pub fn builder() -> PreprocessedRequestBuilder {
PreprocessedRequestBuilder::default()
}
pub fn routing_mut(&mut self) -> &mut RoutingHints {
self.routing.get_or_insert_with(RoutingHints::default)
}
pub fn attach_router_hint(&mut self, hint: &RouterHint) -> serde_json::Result<()> {
let hint_value = serde_json::to_value(hint)?;
let mut map = extra_args_object(self.extra_args.take());
let mut kv_transfer_params = match map.remove(KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY) {
Some(serde_json::Value::Object(params)) => params,
Some(_) | None => serde_json::Map::new(),
};
kv_transfer_params.insert(ROUTER_HINT_EXTRA_ARGS_KEY.to_string(), hint_value);
map.insert(
KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY.to_string(),
serde_json::Value::Object(kv_transfer_params),
);
self.extra_args = Some(serde_json::Value::Object(map));
Ok(())
}
pub fn block_mm_routing_info(&self) -> (&[TokenIdType], Option<&[Option<BlockExtraInfo>]>) {
let Some(mm) = self.mm_routing_info.as_ref() else {
return (&self.token_ids, None);
};
let tokens = mm.routing_token_ids.as_slice();
if tokens.is_empty() {
return (&self.token_ids, None);
}
(tokens, Some(mm.block_mm_infos.as_slice()))
}
}
fn extra_args_object(
extra_args: Option<serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
match extra_args {
Some(serde_json::Value::Object(map)) => map,
_ => serde_json::Map::new(),
}
}
#[derive(Serialize, Deserialize, Debug, Clone, Builder)]
pub struct PreprocessedEmbeddingRequest {
pub token_ids: Vec<Vec<TokenIdType>>,
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub encoding_format: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[builder(default)]
pub truncate_prompt_tokens: Option<i64>,
pub dimensions: Option<u32>,
#[builder(default)]
pub mdc_sum: Option<String>,
#[builder(default)]
pub annotations: Vec<String>,
}
impl PreprocessedEmbeddingRequest {
pub fn has_annotation(&self, annotation: &str) -> bool {
self.annotations.contains(&annotation.to_string())
}
}
impl PreprocessedEmbeddingRequest {
pub fn builder() -> PreprocessedEmbeddingRequestBuilder {
PreprocessedEmbeddingRequestBuilder::default()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn embedding_encoding_format_serde_omits_none() {
let mut request = PreprocessedEmbeddingRequest {
token_ids: vec![vec![1, 2, 3]],
model: "test-model".to_string(),
encoding_format: None,
truncate_prompt_tokens: None,
dimensions: None,
mdc_sum: None,
annotations: Vec::new(),
};
let omitted = serde_json::to_value(&request).unwrap();
assert!(omitted.get("encoding_format").is_none());
assert!(omitted.get("truncate_prompt_tokens").is_none());
let round_trip: PreprocessedEmbeddingRequest = serde_json::from_value(omitted).unwrap();
assert!(round_trip.encoding_format.is_none());
assert!(round_trip.truncate_prompt_tokens.is_none());
request.encoding_format = Some("float".to_string());
request.truncate_prompt_tokens = Some(-1);
let explicit = serde_json::to_value(&request).unwrap();
assert_eq!(explicit["encoding_format"], "float");
assert_eq!(explicit["truncate_prompt_tokens"], -1);
}
#[test]
fn attach_router_hint_preserves_extra_args_object() {
use dynamo_kv_router::{
protocols::ExternalSequenceBlockHash,
router_hint::{ROUTER_HINT_EXTRA_ARGS_KEY, RouterHint},
};
let mut req = PreprocessedRequest::builder()
.model("t".to_string())
.token_ids(vec![1])
.stop_conditions(StopConditions::default())
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.extra_args(Some(serde_json::json!({
"caller": "kept",
"kv_transfer_params": {"existing": "kept"}
})))
.build()
.unwrap();
let hint = RouterHint {
source_control_endpoint: "tcp://127.0.0.1:23280".to_string(),
block_hashes: vec![ExternalSequenceBlockHash(11), ExternalSequenceBlockHash(22)],
};
req.attach_router_hint(&hint).unwrap();
let extra_args = req.extra_args.unwrap();
assert_eq!(extra_args["caller"], "kept");
assert_eq!(
extra_args[KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY]["existing"],
"kept"
);
assert_eq!(
extra_args[KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY][ROUTER_HINT_EXTRA_ARGS_KEY]["source_control_endpoint"],
"tcp://127.0.0.1:23280"
);
assert_eq!(
extra_args[KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY][ROUTER_HINT_EXTRA_ARGS_KEY]["block_hashes"],
serde_json::json!([11, 22])
);
}
#[test]
fn attach_router_hint_replaces_non_object_kv_transfer_params() {
use dynamo_kv_router::{
protocols::ExternalSequenceBlockHash,
router_hint::{ROUTER_HINT_EXTRA_ARGS_KEY, RouterHint},
};
for invalid_params in [
serde_json::Value::Null,
serde_json::json!("invalid"),
serde_json::json!(["invalid"]),
] {
let mut req = PreprocessedRequest::builder()
.model("t".to_string())
.token_ids(vec![1])
.stop_conditions(StopConditions::default())
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.extra_args(Some(serde_json::json!({
"caller": "kept",
"kv_transfer_params": invalid_params
})))
.build()
.unwrap();
let hint = RouterHint {
source_control_endpoint: "tcp://127.0.0.1:23280".to_string(),
block_hashes: vec![ExternalSequenceBlockHash(33)],
};
req.attach_router_hint(&hint).unwrap();
let extra_args = req.extra_args.unwrap();
assert_eq!(extra_args["caller"], "kept");
assert_eq!(
extra_args[KV_TRANSFER_PARAMS_EXTRA_ARGS_KEY][ROUTER_HINT_EXTRA_ARGS_KEY]["block_hashes"],
serde_json::json!([33])
);
}
}
#[test]
fn bootstrap_info_carries_only_stable_handoff_identity() {
let handoff_id = Uuid::from_u128(42);
let info = BootstrapInfo {
bootstrap_host: "127.0.0.1".to_string(),
bootstrap_port: 1234,
bootstrap_room: 7,
handoff_id: Some(handoff_id),
};
let value = serde_json::to_value(&info).unwrap();
assert_eq!(value["handoff_id"], handoff_id.to_string());
assert!(value.get("mocker_handoff_protocol_version").is_none());
assert!(value.get("mocker_handoff_role").is_none());
assert!(value.get("mocker_handoff_engine_type").is_none());
assert_eq!(
serde_json::from_value::<BootstrapInfo>(value)
.unwrap()
.handoff_id,
Some(handoff_id)
);
}
#[test]
fn is_probe_serde_round_trip() {
let mut req = PreprocessedRequest::builder()
.model("t".to_string())
.token_ids(vec![1])
.stop_conditions(StopConditions::default())
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.build()
.unwrap();
assert!(!req.is_probe);
let normal = serde_json::to_string(&req).unwrap();
assert!(!normal.contains("_HEALTH_CHECK"), "got: {normal}");
let back: PreprocessedRequest = serde_json::from_str(&normal).unwrap();
assert!(!back.is_probe);
req.is_probe = true;
let probe = serde_json::to_string(&req).unwrap();
assert!(probe.contains("\"_HEALTH_CHECK\":true"), "got: {probe}");
let back: PreprocessedRequest = serde_json::from_str(&probe).unwrap();
assert!(back.is_probe);
}
#[test]
fn require_reasoning_serde_round_trip() {
let mut req = PreprocessedRequest::builder()
.model("t".to_string())
.token_ids(vec![1])
.stop_conditions(StopConditions::default())
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.build()
.unwrap();
let normal = serde_json::to_value(&req).unwrap();
assert!(
!normal
.as_object()
.unwrap()
.contains_key("require_reasoning")
);
let back: PreprocessedRequest = serde_json::from_value(normal).unwrap();
assert!(!back.require_reasoning);
req.require_reasoning = true;
let guided = serde_json::to_value(&req).unwrap();
assert_eq!(guided["require_reasoning"], true);
let back: PreprocessedRequest = serde_json::from_value(guided).unwrap();
assert!(back.require_reasoning);
}
#[test]
fn minimal_canary_payload_deserializes() {
let req: PreprocessedRequest = serde_json::from_value(serde_json::json!({
"token_ids": [1],
"_HEALTH_CHECK": true,
}))
.unwrap();
assert_eq!(req.token_ids, vec![1]);
assert!(req.is_probe);
assert_eq!(req.model, "");
}
#[test]
fn encoder_result_round_trips_through_serde() {
let payload = serde_json::json!({
"embedding_handle": {
"shape": [1, 1024],
"dtype": "fp16",
"uri": "nixl://encoder-0/embedding-42",
},
"processed_token_ids": [128_000_u32, 200_001_u32, 200_002_u32],
});
let req = PreprocessedRequest::builder()
.model("test/model".to_string())
.token_ids(vec![1, 2, 3])
.stop_conditions(StopConditions::default())
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.encoder_result(Some(payload.clone()))
.build()
.unwrap();
let json = serde_json::to_value(&req).unwrap();
assert_eq!(json["encoder_result"], payload);
let back: PreprocessedRequest = serde_json::from_value(json).unwrap();
assert_eq!(back.encoder_result, Some(payload));
}
#[test]
fn encoder_result_is_absent_when_none() {
let req = PreprocessedRequest::builder()
.model("test/model".to_string())
.token_ids(vec![1, 2, 3])
.stop_conditions(StopConditions::default())
.sampling_options(SamplingOptions::default())
.output_options(OutputOptions::default())
.build()
.unwrap();
assert!(req.encoder_result.is_none());
let json = serde_json::to_value(&req).unwrap();
assert!(
!json.as_object().unwrap().contains_key("encoder_result"),
"encoder_result must be absent from wire when None; got {json}"
);
}
#[test]
fn routing_hints_cache_namespace_serializes_as_cache_salt() {
let hints = RoutingHints {
cache_namespace: Some("tenant-a".to_string()),
..Default::default()
};
let value = serde_json::to_value(&hints).unwrap();
assert_eq!(value["cache_salt"], "tenant-a");
assert!(value.get("cache_namespace").is_none());
let decoded: RoutingHints = serde_json::from_value(value).unwrap();
assert_eq!(decoded.cache_namespace.as_deref(), Some("tenant-a"));
}
}