use std::collections::{HashMap, HashSet};
use axum::http::HeaderMap;
use derive_builder::Builder;
use dynamo_protocols::types::StopReason;
use serde::{Deserialize, Serialize};
use utoipa::ToSchema;
use crate::protocols::TokenIdType;
use crate::protocols::agents::{
AgentContextHeaderValues, agent_context_header_values, session_affinity_header_value,
};
use crate::protocols::common::FinishReason;
use crate::protocols::common::llm_backend::PromptLogprobs;
use crate::protocols::common::timing::{RequestTracker, TimingInfo};
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Default)]
#[serde(deny_unknown_fields)]
pub struct RoutingConstraints {
#[serde(default, skip_serializing_if = "HashSet::is_empty")]
pub required_taints: HashSet<String>,
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
pub preferred_taints: HashMap<String, f32>,
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq)]
#[serde(deny_unknown_fields)]
pub struct RouterParams {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ttft_target: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub itl_target: Option<f64>,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
#[serde(deny_unknown_fields)]
pub struct MetadataUpload {
#[serde(deserialize_with = "deserialize_metadata_upload_url")]
pub url: String,
}
fn deserialize_metadata_upload_url<'de, D>(deserializer: D) -> Result<String, D::Error>
where
D: serde::Deserializer<'de>,
{
let url = String::deserialize(deserializer)?;
let url = url.trim();
if url.is_empty() {
return Err(serde::de::Error::custom(
"metadata_upload.url must not be empty",
));
}
Ok(url.to_string())
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)]
#[serde(deny_unknown_fields)]
pub struct KvHints {
pub evict_session: bool,
}
#[derive(Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum InputTrigger {
UserMessage,
ToolResult,
Other,
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq, Eq)]
pub struct AgentCompaction {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub trigger: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reason: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub implementation: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub phase: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub strategy: Option<String>,
}
#[derive(Serialize, Deserialize, Builder, Debug, Clone, PartialEq, Eq)]
pub struct AgentContext {
pub session_id: String,
#[builder(default, setter(strip_option))]
#[serde(skip_serializing_if = "Option::is_none")]
pub parent_session_id: Option<String>,
#[builder(default, setter(strip_option))]
#[serde(skip_serializing_if = "Option::is_none")]
pub session_final: Option<bool>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub compaction: Option<AgentCompaction>,
#[builder(default, setter(strip_option))]
#[serde(skip_serializing_if = "Option::is_none")]
pub kv_hints: Option<KvHints>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_trigger: Option<InputTrigger>,
}
impl AgentContext {
pub fn builder() -> AgentContextBuilder {
AgentContextBuilder::default()
}
}
#[derive(Serialize, Deserialize, Builder, Debug, Clone, Default, PartialEq)]
#[serde(deny_unknown_fields)]
pub struct AgentHints {
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub priority: Option<i32>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub strict_priority: Option<u32>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub osl: Option<u32>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub speculative_prefill: Option<bool>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub latency_sensitivity: Option<f64>,
}
#[derive(Serialize, Deserialize, Builder, Debug, Clone, PartialEq)]
#[serde(deny_unknown_fields)]
pub struct NvExt {
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub greed_sampling: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub use_raw_prompt: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub annotations: Option<Vec<String>>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub backend_instance_id: Option<u64>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub token_data: Option<Vec<u32>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub max_thinking_tokens: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub cache_salt: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub extra_fields: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[builder(default, setter(strip_option))]
pub metadata_upload: Option<MetadataUpload>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prefill_worker_id: Option<u64>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub decode_worker_id: Option<u64>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dp_rank: Option<u32>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prefill_dp_rank: Option<u32>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub agent_hints: Option<AgentHints>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_timestamp_ms: Option<f64>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub routing_constraints: Option<RoutingConstraints>,
#[builder(default, setter(strip_option))]
#[serde(default, skip_serializing_if = "Option::is_none")]
pub router: Option<RouterParams>,
}
impl Default for NvExt {
fn default() -> Self {
NvExt::builder().build().unwrap()
}
}
impl NvExt {
pub fn builder() -> NvExtBuilder {
NvExtBuilder::default()
}
pub fn has_query_instance_id_annotation(&self) -> bool {
self.annotations.as_ref().is_some_and(|annotations| {
annotations
.iter()
.any(|annotation| annotation.starts_with("query_instance_id:"))
})
}
pub fn has_non_cache_salt_fields(&self) -> bool {
let Self {
greed_sampling,
use_raw_prompt,
annotations,
backend_instance_id,
token_data,
max_thinking_tokens,
cache_salt: _,
extra_fields,
metadata_upload,
prefill_worker_id,
decode_worker_id,
dp_rank,
prefill_dp_rank,
agent_hints,
request_timestamp_ms,
routing_constraints,
router,
} = self;
greed_sampling.is_some()
|| use_raw_prompt.is_some()
|| annotations.is_some()
|| backend_instance_id.is_some()
|| token_data.is_some()
|| max_thinking_tokens.is_some()
|| extra_fields.is_some()
|| metadata_upload.is_some()
|| prefill_worker_id.is_some()
|| decode_worker_id.is_some()
|| dp_rank.is_some()
|| prefill_dp_rank.is_some()
|| agent_hints.is_some()
|| request_timestamp_ms.is_some()
|| routing_constraints.is_some()
|| router.is_some()
}
}
impl NvExtBuilder {
pub fn add_annotation(&mut self, annotation: impl Into<String>) -> &mut Self {
self.annotations
.get_or_insert_with(|| Some(vec![]))
.as_mut()
.expect("annotations should always be Some(Vec)")
.push(annotation.into());
self
}
}
pub fn parse_nvext(raw: Option<serde_json::Value>) -> anyhow::Result<Option<NvExt>> {
raw.map(serde_json::from_value)
.transpose()
.map_err(|err| anyhow::anyhow!("invalid nvext: {err}"))
}
pub const HEADER_WORKER_INSTANCE_ID: &str = "x-dynamo-worker-instance-id";
pub const HEADER_PREFILL_INSTANCE_ID: &str = "x-dynamo-prefill-instance-id";
pub const HEADER_DP_RANK: &str = "x-dynamo-dp-rank";
pub const HEADER_PREFILL_DP_RANK: &str = "x-dynamo-prefill-dp-rank";
pub const HEADER_REQUEST_PRIORITY: &str = "x-dynamo-request-priority";
pub const HEADER_REQUEST_STRICT_PRIORITY: &str = "x-dynamo-request-strict-priority";
pub const HEADER_TENANT_ID: &str = "x-tenant-id";
pub const HEADER_WORKER_INSTANCE_ID_ALIAS: &str = "x-worker-instance-id";
pub const HEADER_PREFILL_INSTANCE_ID_ALIAS: &str = "x-prefill-instance-id";
pub const HEADER_DP_RANK_ALIAS: &str = "x-dp-rank";
pub const HEADER_DATA_PARALLEL_RANK_ALIAS: &str = "x-data-parallel-rank";
pub const HEADER_PREFILL_DP_RANK_ALIAS: &str = "x-prefill-dp-rank";
const UNSET_DP_RANK_SENTINEL: u32 = u32::MAX;
pub fn last_non_empty_trimmed_value<'a>(values: impl Iterator<Item = &'a str>) -> Option<&'a str> {
values
.filter_map(|value| {
let value = value.trim();
(!value.is_empty()).then_some(value)
})
.last()
}
pub fn has_non_cache_salt_routing_headers(headers: &HeaderMap) -> bool {
[
HEADER_WORKER_INSTANCE_ID,
HEADER_WORKER_INSTANCE_ID_ALIAS,
HEADER_PREFILL_INSTANCE_ID,
HEADER_PREFILL_INSTANCE_ID_ALIAS,
HEADER_DP_RANK,
HEADER_DP_RANK_ALIAS,
HEADER_DATA_PARALLEL_RANK_ALIAS,
HEADER_PREFILL_DP_RANK,
HEADER_PREFILL_DP_RANK_ALIAS,
HEADER_REQUEST_PRIORITY,
HEADER_REQUEST_STRICT_PRIORITY,
]
.iter()
.any(|header| headers.contains_key(*header))
}
impl From<AgentContextHeaderValues> for AgentContext {
fn from(values: AgentContextHeaderValues) -> Self {
let kv_hints = (values.session_final == Some(true)).then_some(KvHints {
evict_session: true,
});
Self {
session_id: values.session_id,
parent_session_id: values.parent_session_id,
session_final: values.session_final,
compaction: values.compaction,
kv_hints,
input_trigger: None,
}
}
}
pub const AGENT_CONTEXT_CONTEXT_KEY: &str = "dynamo.llm.agent_context";
pub const SESSION_AFFINITY_CONTEXT_KEY: &str = "dynamo.llm.session_affinity";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SessionAffinityId(String);
impl SessionAffinityId {
pub(crate) fn new(value: impl Into<String>) -> Self {
Self(value.into())
}
pub fn as_str(&self) -> &str {
&self.0
}
}
pub fn agent_context_from_headers(headers: &HeaderMap) -> Option<AgentContext> {
agent_context_header_values(headers).map(AgentContext::from)
}
pub fn session_affinity_from_headers(headers: &HeaderMap) -> Option<SessionAffinityId> {
session_affinity_header_value(headers).map(SessionAffinityId::new)
}
pub fn apply_header_routing_overrides(nvext: Option<NvExt>, headers: &HeaderMap) -> Option<NvExt> {
let worker_id = headers
.get(HEADER_WORKER_INSTANCE_ID)
.or_else(|| headers.get(HEADER_WORKER_INSTANCE_ID_ALIAS))
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok());
let prefill_id = headers
.get(HEADER_PREFILL_INSTANCE_ID)
.or_else(|| headers.get(HEADER_PREFILL_INSTANCE_ID_ALIAS))
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok());
let dp_rank = headers
.get(HEADER_DP_RANK)
.or_else(|| headers.get(HEADER_DP_RANK_ALIAS))
.or_else(|| headers.get(HEADER_DATA_PARALLEL_RANK_ALIAS))
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u32>().ok());
let prefill_dp_rank = headers
.get(HEADER_PREFILL_DP_RANK)
.or_else(|| headers.get(HEADER_PREFILL_DP_RANK_ALIAS))
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u32>().ok());
let prefill_dp_rank = prefill_dp_rank.filter(|rank| *rank != UNSET_DP_RANK_SENTINEL);
let priority_header = headers
.get(HEADER_REQUEST_PRIORITY)
.and_then(|v| v.to_str().ok());
let strict_priority_header = headers
.get(HEADER_REQUEST_STRICT_PRIORITY)
.and_then(|v| v.to_str().ok());
let priority = priority_header.and_then(|s| s.trim().parse::<i32>().ok());
let strict_priority = strict_priority_header.and_then(|s| s.trim().parse::<u32>().ok());
if worker_id.is_none()
&& prefill_id.is_none()
&& dp_rank.is_none()
&& prefill_dp_rank.is_none()
&& priority.is_none()
&& strict_priority.is_none()
{
return nvext;
}
let mut ext = nvext.unwrap_or_default();
if let Some(id) = worker_id {
ext.backend_instance_id = Some(id);
ext.decode_worker_id = Some(id);
}
if let Some(id) = prefill_id {
ext.prefill_worker_id = Some(id);
}
if let Some(rank) = dp_rank {
ext.dp_rank = Some(rank);
}
if let Some(rank) = prefill_dp_rank {
ext.prefill_dp_rank = Some(rank);
}
if priority.is_some() || strict_priority.is_some() {
let resolved = resolve_request_priority(
ext.agent_hints.as_ref(),
priority_header,
strict_priority_header,
);
let hints = ext.agent_hints.get_or_insert_with(AgentHints::default);
hints.priority = resolved.priority;
hints.strict_priority = resolved.strict_priority;
}
Some(ext)
}
pub fn apply_cache_salt_header_override(
nvext: Option<NvExt>,
headers: &HeaderMap,
) -> Option<NvExt> {
let Some(cache_salt) = last_non_empty_trimmed_value(
headers
.get_all(HEADER_TENANT_ID)
.iter()
.filter_map(|value| value.to_str().ok()),
)
.map(str::to_owned) else {
return nvext;
};
let mut nvext = nvext.unwrap_or_default();
nvext.cache_salt = Some(cache_salt);
Some(nvext)
}
pub fn retain_cache_salt(nvext: Option<NvExt>) -> Option<NvExt> {
nvext.and_then(|nvext| {
nvext
.cache_salt
.filter(|cache_salt| !cache_salt.is_empty())
.map(|cache_salt| NvExt {
cache_salt: Some(cache_salt),
..Default::default()
})
})
}
pub fn apply_frontend_nvext_policy(
nvext: Option<NvExt>,
headers: &HeaderMap,
nvext_enabled: bool,
) -> Option<NvExt> {
let nvext = apply_cache_salt_header_override(nvext, headers);
if nvext_enabled {
apply_header_routing_overrides(nvext, headers)
} else {
retain_cache_salt(nvext)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct ResolvedPriority {
pub priority: Option<i32>,
pub strict_priority: Option<u32>,
pub priority_jump: Option<f64>,
}
pub fn resolve_request_priority(
hints: Option<&AgentHints>,
priority_header: Option<&str>,
strict_priority_header: Option<&str>,
) -> ResolvedPriority {
let priority = priority_header
.and_then(|h| h.trim().parse::<i32>().ok())
.or_else(|| hints.and_then(|h| h.priority));
let strict_priority = strict_priority_header
.and_then(|h| h.trim().parse::<u32>().ok())
.or_else(|| hints.and_then(|h| h.strict_priority));
let priority_jump = priority
.map(|p| p as f64)
.or_else(|| hints.and_then(|h| h.latency_sensitivity));
ResolvedPriority {
priority,
strict_priority,
priority_jump,
}
}
pub trait NvExtProvider {
fn nvext(&self) -> Option<&NvExt>;
fn raw_prompt(&self) -> Option<String>;
fn unsupported_fields(&self) -> Option<&std::collections::HashMap<String, serde_json::Value>> {
None
}
}
pub fn request_cache_salt<R: NvExtProvider>(request: &R) -> Option<&str> {
request
.nvext()
.and_then(|nvext| nvext.cache_salt.as_deref())
.filter(|salt| !salt.is_empty())
.or_else(|| {
request
.unsupported_fields()
.and_then(|fields| fields.get("cache_salt"))
.and_then(|value| value.as_str())
.filter(|salt| !salt.is_empty())
})
}
pub fn routing_constraints_to_kv(
constraints: RoutingConstraints,
) -> dynamo_kv_router::protocols::RoutingConstraints {
dynamo_kv_router::protocols::RoutingConstraints {
required_taints: constraints.required_taints,
preferred_taints: constraints.preferred_taints,
}
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct WorkerIdInfo {
#[serde(skip_serializing_if = "Option::is_none")]
pub prefill_worker_id: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prefill_dp_rank: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub decode_worker_id: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub decode_dp_rank: Option<u32>,
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone)]
pub struct NvExtResponse {
#[serde(skip_serializing_if = "Option::is_none")]
pub worker_id: Option<WorkerIdInfo>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timing: Option<TimingInfo>,
#[serde(skip_serializing_if = "Option::is_none")]
pub token_ids: Option<Vec<u32>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub routed_experts: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub engine_data: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop_reason: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub detailed_finish_reason: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub completion_token_ids: Option<Vec<TokenIdType>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_token_ids: Option<Vec<TokenIdType>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_logprobs: Option<PromptLogprobs>,
}
pub(crate) fn merge_response_nvext(
target: &mut Option<serde_json::Value>,
incoming: Option<serde_json::Value>,
) {
let Some(incoming) = incoming else {
return;
};
match (target.as_mut(), incoming) {
(Some(serde_json::Value::Object(target_obj)), serde_json::Value::Object(incoming_obj)) => {
for (key, value) in incoming_obj {
match key.as_str() {
"completion_token_ids" => {
let entry = target_obj
.entry(&key)
.or_insert_with(|| serde_json::Value::Array(Vec::new()));
if let (serde_json::Value::Array(acc), serde_json::Value::Array(new)) =
(entry, value)
{
acc.extend(new);
}
}
"prompt_logprobs" => {
target_obj.insert(key, value);
}
_ => {
target_obj.insert(key, value);
}
}
}
}
(_, incoming) => {
*target = Some(incoming);
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct NvExtResponseFieldSelection {
pub worker_id: bool,
pub timing: bool,
pub token_ids: bool,
pub routed_experts: bool,
pub engine_data: bool,
pub stop_reason: bool,
pub detailed_finish_reason: bool,
pub completion_token_ids: bool,
pub prompt_token_ids: bool,
pub prompt_logprobs: bool,
}
#[derive(Debug, Default)]
pub struct NvExtResponseInput<'a> {
pub tracker: Option<&'a RequestTracker>,
pub finish_reason: Option<&'a FinishReason>,
pub engine_data: Option<serde_json::Value>,
pub stop_reason: Option<StopReason>,
pub completion_token_ids: Option<&'a [TokenIdType]>,
pub prompt_logprobs: Option<PromptLogprobs>,
}
impl NvExtResponseFieldSelection {
pub fn from_nvext(nvext: Option<&NvExt>) -> Self {
let Some(ext) = nvext else {
return Self::default();
};
let mut selection = Self::default();
if let Some(fields) = ext.extra_fields.as_ref() {
for field in fields {
match field.as_str() {
"worker_id" => selection.worker_id = true,
"timing" => selection.timing = true,
"routed_experts" => selection.routed_experts = true,
"engine_data" => selection.engine_data = true,
"stop_reason" => selection.stop_reason = true,
"detailed_finish_reason" => selection.detailed_finish_reason = true,
"completion_token_ids" => selection.completion_token_ids = true,
"prompt_token_ids" => selection.prompt_token_ids = true,
"prompt_logprobs" => selection.prompt_logprobs = true,
_ => {}
}
}
}
if ext.has_query_instance_id_annotation() {
selection.worker_id = true;
selection.token_ids = true;
}
selection
}
pub fn build_response_nvext(&self, input: NvExtResponseInput<'_>) -> Option<NvExtResponse> {
let finish_reason_present = input.finish_reason.is_some();
let worker_id = if self.worker_id {
input.tracker.and_then(RequestTracker::get_worker_info)
} else {
None
};
let token_ids = if self.token_ids {
input
.tracker
.and_then(|tracker| tracker.query_token_ids().map(<[u32]>::to_vec))
} else {
None
};
let routed_experts = if self.routed_experts {
input
.engine_data
.as_ref()
.and_then(|data| data.get("routed_experts"))
.cloned()
} else {
None
};
let timing = if finish_reason_present && self.timing {
input.tracker.map(RequestTracker::get_timing_info)
} else {
None
};
let engine_data = if self.engine_data {
input.engine_data
} else {
None
};
let stop_reason = if self.stop_reason {
input
.stop_reason
.and_then(|reason| serde_json::to_value(reason).ok())
} else {
None
};
let detailed_finish_reason = if self.detailed_finish_reason {
input.finish_reason.map(ToString::to_string)
} else {
None
};
let completion_token_ids = if self.completion_token_ids {
input.completion_token_ids.map(<[u32]>::to_vec)
} else {
None
};
let prompt_token_ids = if self.prompt_token_ids && finish_reason_present {
input
.tracker
.and_then(|tracker| tracker.prompt_token_ids().map(<[u32]>::to_vec))
} else {
None
};
let prompt_logprobs = if self.prompt_logprobs && finish_reason_present {
input.prompt_logprobs
} else {
None
};
if worker_id.is_none()
&& token_ids.is_none()
&& routed_experts.is_none()
&& timing.is_none()
&& engine_data.is_none()
&& stop_reason.is_none()
&& detailed_finish_reason.is_none()
&& completion_token_ids.is_none()
&& prompt_token_ids.is_none()
&& prompt_logprobs.is_none()
{
return None;
}
Some(NvExtResponse {
worker_id,
timing,
token_ids,
routed_experts,
engine_data,
stop_reason,
detailed_finish_reason,
completion_token_ids,
prompt_token_ids,
prompt_logprobs,
})
}
}
pub(crate) fn validate_completion_token_ids_single_choice(
total_choices: usize,
nvext: Option<&NvExt>,
) -> anyhow::Result<()> {
let requested = nvext
.and_then(|ext| ext.extra_fields.as_ref())
.is_some_and(|fields| fields.iter().any(|field| field == "completion_token_ids"));
if requested && total_choices > 1 {
anyhow::bail!(
"`nvext.extra_fields=[\"completion_token_ids\"]` requires exactly one generated choice"
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocols::agents::{
HEADER_CLAUDE_CODE_AGENT_ID, HEADER_CLAUDE_CODE_PARENT_AGENT_ID,
HEADER_CLAUDE_CODE_SESSION_ID, HEADER_CODEX_PARENT_THREAD_ID, HEADER_CODEX_THREAD_ID,
HEADER_CODEX_TURN_METADATA, HEADER_DYNAMO_PARENT_SESSION_ID, HEADER_DYNAMO_SESSION_FINAL,
HEADER_DYNAMO_SESSION_ID, HEADER_OPENCODE_PARENT_SESSION_ID, HEADER_OPENCODE_SESSION_ID,
};
#[derive(Default)]
struct CacheSaltRequest {
nvext: Option<NvExt>,
unsupported_fields: HashMap<String, serde_json::Value>,
}
impl NvExtProvider for CacheSaltRequest {
fn nvext(&self) -> Option<&NvExt> {
self.nvext.as_ref()
}
fn raw_prompt(&self) -> Option<String> {
None
}
fn unsupported_fields(&self) -> Option<&HashMap<String, serde_json::Value>> {
Some(&self.unsupported_fields)
}
}
#[test]
fn agent_context_accepts_missing_input_trigger() {
let context: AgentContext = serde_json::from_str(r#"{"session_id":"root"}"#).unwrap();
assert_eq!(context.input_trigger, None);
}
#[test]
fn agent_context_accepts_nested_compaction() {
let context = serde_json::from_str::<AgentContext>(
r#"{"session_id":"root","compaction":{"trigger":"manual"}}"#,
)
.unwrap();
assert_eq!(
context.compaction.and_then(|compaction| compaction.trigger),
Some("manual".to_string())
);
}
#[test]
fn request_cache_salt_uses_canonical_precedence_and_empty_fallbacks() {
let mut request = CacheSaltRequest::default();
assert_eq!(request_cache_salt(&request), None);
request
.unsupported_fields
.insert("cache_salt".to_string(), serde_json::json!("tenant-legacy"));
assert_eq!(request_cache_salt(&request), Some("tenant-legacy"));
request.nvext = Some(NvExt {
cache_salt: Some("tenant-nvext".to_string()),
..Default::default()
});
assert_eq!(request_cache_salt(&request), Some("tenant-nvext"));
request.nvext.as_mut().unwrap().cache_salt = Some(String::new());
assert_eq!(request_cache_salt(&request), Some("tenant-legacy"));
request
.unsupported_fields
.insert("cache_salt".to_string(), serde_json::json!(""));
assert_eq!(request_cache_salt(&request), None);
}
#[test]
fn parse_nvext_rejects_unknown_fields_in_llm_layer() {
let err = parse_nvext(Some(serde_json::json!({
"unsupported_future_field": true
})))
.unwrap_err();
assert!(err.to_string().contains("invalid nvext"));
assert!(err.to_string().contains("unknown field"));
}
#[test]
fn agent_hints_strict_priority_serde() {
let hints: AgentHints = serde_json::from_str(r#"{"strict_priority":3}"#).unwrap();
assert_eq!(hints.strict_priority, Some(3));
assert_eq!(
serde_json::to_string(&hints).unwrap(),
r#"{"strict_priority":3}"#
);
assert!(serde_json::from_str::<AgentHints>(r#"{"strict_priority":-1}"#).is_err());
}
#[test]
fn nvext_agent_context_is_rejected() {
for json in [
r#"{"agent_context":{"session_id":"run-123"}}"#,
r#"{"agent_context":{"session_id":"run-123","parent_session_id":"root-1"}}"#,
r#"{"agent_context":{"session_id":"run-123","session_final":true}}"#,
] {
let err = serde_json::from_str::<NvExt>(json).unwrap_err();
assert!(err.to_string().contains("unknown field `agent_context`"));
}
}
#[test]
fn metadata_upload_parses_url() {
let nvext: NvExt = serde_json::from_value(serde_json::json!({
"metadata_upload": {
"url": " s3://bucket/root/rollouts "
}
}))
.unwrap();
let upload = nvext.metadata_upload.as_ref().unwrap();
assert_eq!(upload.url, "s3://bucket/root/rollouts");
assert!(!NvExtResponseFieldSelection::from_nvext(Some(&nvext)).engine_data);
assert!(
serde_json::from_value::<NvExt>(serde_json::json!({
"metadata_upload": {}
}))
.is_err()
);
assert!(
serde_json::from_value::<NvExt>(serde_json::json!({
"metadata_upload": {
"url": ""
}
}))
.is_err()
);
assert!(
serde_json::from_value::<NvExt>(serde_json::json!({
"metadata_upload": {
"url": "s3://bucket/root/rollouts",
"format": "json"
}
}))
.is_err()
);
}
#[test]
fn apply_header_routing_overrides_sets_worker_fields() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_WORKER_INSTANCE_ID, "123".parse().unwrap());
headers.insert(HEADER_PREFILL_INSTANCE_ID, "456".parse().unwrap());
headers.insert(HEADER_DP_RANK, "3".parse().unwrap());
headers.insert(HEADER_PREFILL_DP_RANK, "5".parse().unwrap());
headers.insert(HEADER_WORKER_INSTANCE_ID_ALIAS, "1".parse().unwrap());
headers.insert(HEADER_PREFILL_INSTANCE_ID_ALIAS, "2".parse().unwrap());
headers.insert(HEADER_DP_RANK_ALIAS, "4".parse().unwrap());
headers.insert(HEADER_PREFILL_DP_RANK_ALIAS, "6".parse().unwrap());
let result = apply_header_routing_overrides(None, &headers).unwrap();
assert_eq!(result.backend_instance_id, Some(123));
assert_eq!(result.decode_worker_id, Some(123));
assert_eq!(result.prefill_worker_id, Some(456));
assert_eq!(result.dp_rank, Some(3));
assert_eq!(result.prefill_dp_rank, Some(5));
}
#[test]
fn apply_header_routing_overrides_sets_priorities() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_REQUEST_PRIORITY, "-3".parse().unwrap());
headers.insert(HEADER_REQUEST_STRICT_PRIORITY, "7".parse().unwrap());
let hints = apply_header_routing_overrides(None, &headers)
.unwrap()
.agent_hints
.unwrap();
assert_eq!(hints.priority, Some(-3));
assert_eq!(hints.strict_priority, Some(7));
headers.remove(HEADER_REQUEST_STRICT_PRIORITY);
let nvext = NvExt {
agent_hints: Some(AgentHints {
priority: Some(1),
strict_priority: Some(2),
osl: Some(99),
..Default::default()
}),
..Default::default()
};
let hints = apply_header_routing_overrides(Some(nvext), &headers)
.unwrap()
.agent_hints
.unwrap();
assert_eq!(hints.priority, Some(-3));
assert_eq!(hints.strict_priority, Some(2));
assert_eq!(hints.osl, Some(99));
}
#[test]
fn resolve_request_priority_header_over_body_with_independent_fallback() {
let hints = AgentHints {
priority: Some(5),
strict_priority: Some(2),
latency_sensitivity: Some(1.5),
..Default::default()
};
let r = resolve_request_priority(Some(&hints), Some("-3"), Some("7"));
assert_eq!(r.priority, Some(-3));
assert_eq!(r.strict_priority, Some(7));
assert_eq!(r.priority_jump, Some(-3.0));
let r = resolve_request_priority(Some(&hints), Some("0"), None);
assert_eq!(r.priority, Some(0));
assert_eq!(r.priority_jump, Some(0.0));
assert_eq!(r.strict_priority, Some(2));
let r = resolve_request_priority(Some(&hints), Some("abc"), None);
assert_eq!(r.priority, Some(5));
assert_eq!(r.strict_priority, Some(2));
assert_eq!(r.priority_jump, Some(5.0));
let ls_only = AgentHints {
latency_sensitivity: Some(2.5),
..Default::default()
};
let r = resolve_request_priority(Some(&ls_only), None, None);
assert_eq!(r.priority, None);
assert_eq!(r.priority_jump, Some(2.5));
let both = AgentHints {
priority: Some(4),
latency_sensitivity: Some(9.0),
..Default::default()
};
assert_eq!(
resolve_request_priority(Some(&both), None, None).priority_jump,
Some(4.0)
);
assert_eq!(
resolve_request_priority(None, None, None),
ResolvedPriority::default()
);
}
#[test]
fn cache_salt_header_override_has_precedence() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_TENANT_ID, "tenant-a".parse().unwrap());
let nvext = apply_cache_salt_header_override(None, &headers).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-a"));
let mut headers = HeaderMap::new();
headers.insert(HEADER_TENANT_ID, "tenant-header".parse().unwrap());
let nvext = NvExt {
cache_salt: Some("tenant-body".to_string()),
..Default::default()
};
let nvext = apply_cache_salt_header_override(Some(nvext), &headers).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-header"));
headers.append(HEADER_TENANT_ID, " ".parse().unwrap());
headers.append(HEADER_TENANT_ID, " tenant-gateway ".parse().unwrap());
let nvext = apply_cache_salt_header_override(Some(nvext), &headers).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-gateway"));
}
#[test]
fn cache_salt_header_override_ignores_only_empty_values() {
let mut headers = HeaderMap::new();
headers.append(HEADER_TENANT_ID, "".parse().unwrap());
headers.append(HEADER_TENANT_ID, " ".parse().unwrap());
let nvext = NvExt {
cache_salt: Some("tenant-body".to_string()),
..Default::default()
};
let nvext = apply_cache_salt_header_override(Some(nvext), &headers).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-body"));
}
#[test]
fn frontend_nvext_policy_preserves_cache_salt_when_disabled() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_WORKER_INSTANCE_ID, "42".parse().unwrap());
headers.insert(HEADER_REQUEST_PRIORITY, "7".parse().unwrap());
let body = NvExt {
cache_salt: Some("tenant-body".to_string()),
backend_instance_id: Some(99),
extra_fields: Some(vec!["worker_id".to_string()]),
..Default::default()
};
let nvext = apply_frontend_nvext_policy(Some(body), &headers, false).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-body"));
assert!(!nvext.has_non_cache_salt_fields());
}
#[test]
fn frontend_nvext_policy_resolves_header_body_and_empty_values() {
let body = || NvExt {
cache_salt: Some("tenant-body".to_string()),
..Default::default()
};
let mut headers = HeaderMap::new();
headers.insert(HEADER_TENANT_ID, "tenant-header".parse().unwrap());
let nvext = apply_frontend_nvext_policy(Some(body()), &headers, false).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-header"));
headers.insert(HEADER_TENANT_ID, "".parse().unwrap());
let nvext = apply_frontend_nvext_policy(Some(body()), &headers, false).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-body"));
let empty_body = NvExt {
cache_salt: Some(String::new()),
..Default::default()
};
let mut request = CacheSaltRequest {
nvext: apply_frontend_nvext_policy(Some(empty_body), &headers, false),
..Default::default()
};
request
.unsupported_fields
.insert("cache_salt".to_string(), serde_json::json!("tenant-legacy"));
assert_eq!(request_cache_salt(&request), Some("tenant-legacy"));
assert!(apply_frontend_nvext_policy(None, &HeaderMap::new(), false).is_none());
}
#[test]
fn frontend_nvext_policy_keeps_full_enabled_behavior() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_TENANT_ID, "tenant-header".parse().unwrap());
headers.insert(HEADER_WORKER_INSTANCE_ID, "42".parse().unwrap());
let body = NvExt {
cache_salt: Some("tenant-body".to_string()),
extra_fields: Some(vec!["worker_id".to_string()]),
..Default::default()
};
let nvext = apply_frontend_nvext_policy(Some(body), &headers, true).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-header"));
assert_eq!(nvext.backend_instance_id, Some(42));
assert_eq!(nvext.decode_worker_id, Some(42));
assert_eq!(nvext.extra_fields, Some(vec!["worker_id".to_string()]));
}
#[test]
fn apply_header_routing_overrides_supports_unprefixed_aliases() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_WORKER_INSTANCE_ID_ALIAS, "123".parse().unwrap());
headers.insert(HEADER_PREFILL_INSTANCE_ID_ALIAS, "456".parse().unwrap());
headers.insert(HEADER_DP_RANK_ALIAS, "3".parse().unwrap());
headers.insert(HEADER_PREFILL_DP_RANK_ALIAS, "5".parse().unwrap());
let result = apply_header_routing_overrides(None, &headers).unwrap();
assert_eq!(result.backend_instance_id, Some(123));
assert_eq!(result.decode_worker_id, Some(123));
assert_eq!(result.prefill_worker_id, Some(456));
assert_eq!(result.dp_rank, Some(3));
assert_eq!(result.prefill_dp_rank, Some(5));
headers.remove(HEADER_DP_RANK_ALIAS);
headers.insert(HEADER_DATA_PARALLEL_RANK_ALIAS, "4".parse().unwrap());
assert_eq!(
apply_header_routing_overrides(None, &headers)
.unwrap()
.dp_rank,
Some(4)
);
}
#[test]
fn routing_overrides_do_not_apply_tenant_header() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_TENANT_ID, "tenant-header".parse().unwrap());
let nvext = NvExt {
cache_salt: Some("tenant-body".to_string()),
..Default::default()
};
let nvext = apply_header_routing_overrides(Some(nvext), &headers).unwrap();
assert_eq!(nvext.cache_salt.as_deref(), Some("tenant-body"));
}
#[test]
fn agent_context_from_headers_derives_agent_context_table() {
use axum::http::{HeaderMap, HeaderName};
let cases = [
(HEADER_CLAUDE_CODE_SESSION_ID, "claude-run-1", None, None),
(HEADER_CODEX_THREAD_ID, "codex-root", None, None),
(
HEADER_OPENCODE_SESSION_ID,
"opencode-run-1",
Some("parent-run-1"),
Some("parent-run-1"),
),
(HEADER_DYNAMO_SESSION_ID, "generic-run-1", None, None),
];
for (header_name, header_value, parent_header_value, expected_parent_session_id) in cases {
let mut headers = HeaderMap::new();
headers.insert(
header_name.parse::<HeaderName>().unwrap(),
header_value.parse().unwrap(),
);
if let Some(parent) = parent_header_value {
headers.insert(HEADER_OPENCODE_PARENT_SESSION_ID, parent.parse().unwrap());
}
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id.as_str(), header_value);
assert_eq!(
agent_context.parent_session_id.as_deref(),
expected_parent_session_id
);
assert_eq!(agent_context.session_final, None);
assert_eq!(agent_context.kv_hints, None);
}
}
#[test]
fn agent_context_from_codex_compaction_header_preserves_metadata() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_CODEX_THREAD_ID, "codex-thread".parse().unwrap());
headers.insert(
HEADER_CODEX_TURN_METADATA,
r#"{"request_kind":"compaction","compaction":{"trigger":"manual","reason":"user_requested","implementation":"responses_compact","phase":"standalone_turn","strategy":"memento"}}"#
.parse()
.unwrap(),
);
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(
agent_context.compaction,
Some(AgentCompaction {
trigger: Some("manual".to_string()),
reason: Some("user_requested".to_string()),
implementation: Some("responses_compact".to_string()),
phase: Some("standalone_turn".to_string()),
strategy: Some("memento".to_string()),
})
);
headers.insert(HEADER_DYNAMO_SESSION_ID, "canonical".parse().unwrap());
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "canonical");
assert_eq!(
agent_context
.compaction
.and_then(|compaction| compaction.strategy),
Some("memento".to_string())
);
}
#[test]
fn agent_context_ignores_invalid_or_non_compaction_codex_metadata() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_CODEX_THREAD_ID, "codex-thread".parse().unwrap());
headers.insert(HEADER_CODEX_TURN_METADATA, "{".parse().unwrap());
assert_eq!(
agent_context_from_headers(&headers).unwrap().compaction,
None
);
headers.insert(
HEADER_CODEX_TURN_METADATA,
r#"{"request_kind":"turn","compaction":{"trigger":"manual"}}"#
.parse()
.unwrap(),
);
assert_eq!(
agent_context_from_headers(&headers).unwrap().compaction,
None
);
headers.insert(
HEADER_CODEX_TURN_METADATA,
r#"{"request_kind":"compaction"}"#.parse().unwrap(),
);
assert_eq!(
agent_context_from_headers(&headers).unwrap().compaction,
Some(AgentCompaction::default())
);
}
#[test]
fn agent_context_from_codex_thread_headers_preserves_subagent_lineage() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_CODEX_THREAD_ID, "codex-child".parse().unwrap());
headers.insert(HEADER_CODEX_PARENT_THREAD_ID, "codex-root".parse().unwrap());
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "codex-child");
assert_eq!(
agent_context.parent_session_id.as_deref(),
Some("codex-root")
);
assert_eq!(
session_affinity_from_headers(&headers).unwrap().as_str(),
"codex-child"
);
}
#[test]
fn agent_context_ignores_self_parent_header() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_CODEX_THREAD_ID, "codex-thread".parse().unwrap());
headers.insert(
HEADER_CODEX_PARENT_THREAD_ID,
"codex-thread".parse().unwrap(),
);
assert_eq!(
agent_context_from_headers(&headers)
.unwrap()
.parent_session_id,
None
);
}
#[test]
fn codex_session_id_is_ignored_without_thread_id() {
let mut headers = HeaderMap::new();
headers.insert("session-id", "codex-run".parse().unwrap());
assert!(agent_context_from_headers(&headers).is_none());
assert!(session_affinity_from_headers(&headers).is_none());
headers.insert(HEADER_CODEX_THREAD_ID, "codex-thread".parse().unwrap());
assert_eq!(
agent_context_from_headers(&headers).unwrap().session_id,
"codex-thread"
);
}
#[test]
fn session_affinity_prefers_dynamo_header_over_agent_mappings() {
let mut headers = HeaderMap::new();
headers.insert(
HEADER_CLAUDE_CODE_SESSION_ID,
"claude-session".parse().unwrap(),
);
headers.insert(HEADER_CODEX_THREAD_ID, "codex-thread".parse().unwrap());
headers.insert(
HEADER_OPENCODE_SESSION_ID,
"opencode-session".parse().unwrap(),
);
assert_eq!(
session_affinity_from_headers(&headers).unwrap().as_str(),
"claude-session"
);
headers.insert(HEADER_DYNAMO_SESSION_ID, "canonical".parse().unwrap());
assert_eq!(
session_affinity_from_headers(&headers).unwrap().as_str(),
"canonical"
);
headers.insert(HEADER_DYNAMO_SESSION_ID, " ".parse().unwrap());
assert_eq!(
session_affinity_from_headers(&headers).unwrap().as_str(),
"claude-session"
);
}
#[test]
fn session_affinity_uses_agent_child_session_when_present() {
let mut headers = HeaderMap::new();
headers.insert(
HEADER_CLAUDE_CODE_SESSION_ID,
"claude-session".parse().unwrap(),
);
headers.insert(HEADER_CLAUDE_CODE_AGENT_ID, "claude-agent".parse().unwrap());
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "claude-agent");
assert_eq!(
session_affinity_from_headers(&headers).unwrap().as_str(),
"claude-agent"
);
headers.insert(
HEADER_DYNAMO_SESSION_ID,
"affinity-session".parse().unwrap(),
);
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "affinity-session");
assert_eq!(
session_affinity_from_headers(&headers).unwrap().as_str(),
"affinity-session"
);
}
#[test]
fn session_affinity_absent_without_any_session_header() {
let mut headers = HeaderMap::new();
assert!(session_affinity_from_headers(&headers).is_none());
headers.insert(HEADER_DYNAMO_SESSION_ID, " ".parse().unwrap());
assert!(session_affinity_from_headers(&headers).is_none());
}
#[test]
fn agent_context_from_headers_uses_claude_agent_lineage() {
let mut headers = HeaderMap::new();
headers.insert(
HEADER_CLAUDE_CODE_SESSION_ID,
"claude-session".parse().unwrap(),
);
headers.insert(HEADER_CLAUDE_CODE_AGENT_ID, "claude-agent".parse().unwrap());
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "claude-agent");
assert_eq!(
agent_context.parent_session_id.as_deref(),
Some("claude-session")
);
headers.insert(
HEADER_CLAUDE_CODE_PARENT_AGENT_ID,
"claude-parent-agent".parse().unwrap(),
);
assert_eq!(
agent_context_from_headers(&headers)
.unwrap()
.parent_session_id
.as_deref(),
Some("claude-parent-agent")
);
headers.remove(HEADER_CLAUDE_CODE_AGENT_ID);
let root_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(root_context.session_id, "claude-session");
assert_eq!(root_context.parent_session_id, None);
}
#[test]
fn agent_context_from_headers_reads_dynamo_parent_and_final() {
let mut headers = HeaderMap::new();
headers.insert(HEADER_DYNAMO_SESSION_ID, "generic-run".parse().unwrap());
headers.insert(
HEADER_DYNAMO_PARENT_SESSION_ID,
"generic-parent".parse().unwrap(),
);
headers.insert(HEADER_DYNAMO_SESSION_FINAL, "true".parse().unwrap());
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "generic-run");
assert_eq!(
agent_context.parent_session_id.as_deref(),
Some("generic-parent")
);
assert_eq!(agent_context.session_final, Some(true));
assert_eq!(
agent_context.kv_hints,
Some(KvHints {
evict_session: true
})
);
headers.insert(HEADER_DYNAMO_SESSION_FINAL, "false".parse().unwrap());
assert_eq!(agent_context_from_headers(&headers).unwrap().kv_hints, None);
}
#[test]
fn dynamo_session_headers_override_agent_native_headers() {
let mut headers = HeaderMap::new();
headers.insert(
HEADER_CLAUDE_CODE_SESSION_ID,
"claude-session".parse().unwrap(),
);
headers.insert(HEADER_CLAUDE_CODE_AGENT_ID, "claude-agent".parse().unwrap());
headers.insert(HEADER_DYNAMO_SESSION_ID, "dynamo-session".parse().unwrap());
headers.insert(
HEADER_DYNAMO_PARENT_SESSION_ID,
"dynamo-parent".parse().unwrap(),
);
let agent_context = agent_context_from_headers(&headers).unwrap();
assert_eq!(agent_context.session_id, "dynamo-session");
assert_eq!(
agent_context.parent_session_id.as_deref(),
Some("dynamo-parent")
);
}
#[test]
fn apply_header_routing_overrides_ignores_session_identity_headers() {
use axum::http::{HeaderMap, HeaderName};
for header_name in [
HEADER_DYNAMO_SESSION_ID,
HEADER_DYNAMO_PARENT_SESSION_ID,
HEADER_DYNAMO_SESSION_FINAL,
] {
let mut headers = HeaderMap::new();
headers.insert(
header_name.parse::<HeaderName>().unwrap(),
"session-value".parse().unwrap(),
);
assert!(apply_header_routing_overrides(None, &headers).is_none());
}
}
#[test]
fn query_instance_annotation_detection_is_exact_prefix() {
let nvext = NvExt::builder()
.annotations(vec![
"query_instance_id_extra:bad".to_string(),
"query_instance_id:good".to_string(),
])
.build()
.unwrap();
assert!(nvext.has_query_instance_id_annotation());
let nvext = NvExt::builder()
.annotations(vec!["query_instance_id_extra:bad".to_string()])
.build()
.unwrap();
assert!(!nvext.has_query_instance_id_annotation());
}
#[test]
fn response_field_selection_respects_extra_fields() {
let nvext = NvExt::builder()
.extra_fields(vec![
"worker_id".to_string(),
"routed_experts".to_string(),
"prompt_token_ids".to_string(),
"detailed_finish_reason".to_string(),
])
.build()
.unwrap();
assert_eq!(
NvExtResponseFieldSelection::from_nvext(Some(&nvext)),
NvExtResponseFieldSelection {
worker_id: true,
routed_experts: true,
prompt_token_ids: true,
detailed_finish_reason: true,
..Default::default()
}
);
}
#[test]
fn response_field_selection_query_instance_id_exception() {
let nvext = NvExt::builder()
.annotations(vec!["query_instance_id:".to_string()])
.build()
.unwrap();
let selection = NvExtResponseFieldSelection::from_nvext(Some(&nvext));
assert!(selection.worker_id);
assert!(selection.token_ids);
assert!(!selection.timing);
assert!(!selection.routed_experts);
}
#[test]
fn response_field_selection_multiple_extra_fields() {
let nvext = NvExt::builder()
.extra_fields(vec![
"worker_id".to_string(),
"timing".to_string(),
"routed_experts".to_string(),
])
.build()
.unwrap();
assert_eq!(
NvExtResponseFieldSelection::from_nvext(Some(&nvext)),
NvExtResponseFieldSelection {
worker_id: true,
timing: true,
routed_experts: true,
..Default::default()
}
);
}
fn tracker_with_prefill_worker()
-> std::sync::Arc<crate::protocols::common::timing::RequestTracker> {
use crate::protocols::common::timing::{RequestTracker, WORKER_TYPE_PREFILL};
let tracker = std::sync::Arc::new(RequestTracker::new());
tracker.record_worker(42, Some(0), WORKER_TYPE_PREFILL);
tracker
}
fn tracker_with_query_token_ids()
-> std::sync::Arc<crate::protocols::common::timing::RequestTracker> {
use crate::protocols::common::timing::RequestTracker;
let tracker = std::sync::Arc::new(RequestTracker::new());
tracker.set_external_query_token_ids(vec![11u32, 22, 33]);
tracker
}
fn tracker_with_prompt_token_ids()
-> std::sync::Arc<crate::protocols::common::timing::RequestTracker> {
use crate::protocols::common::timing::RequestTracker;
let tracker = std::sync::Arc::new(RequestTracker::new());
tracker.set_prompt_token_ids(vec![101u32, 102, 103]);
tracker
}
fn tracker_with_forwarded_worker_info()
-> std::sync::Arc<crate::protocols::common::timing::RequestTracker> {
use crate::protocols::common::timing::RequestTracker;
let tracker = std::sync::Arc::new(RequestTracker::new());
tracker.set_external_worker_info(WorkerIdInfo {
prefill_worker_id: Some(7),
prefill_dp_rank: Some(1),
decode_worker_id: Some(9),
decode_dp_rank: Some(2),
});
tracker
}
#[test]
fn build_response_nvext_all_false_returns_none() {
let finish_reason = FinishReason::Cancelled;
assert!(
NvExtResponseFieldSelection::default()
.build_response_nvext(NvExtResponseInput {
finish_reason: Some(&finish_reason),
..Default::default()
})
.is_none()
);
}
#[test]
fn build_response_nvext_worker_id_only_without_finish() {
let selection = NvExtResponseFieldSelection {
worker_id: true,
..Default::default()
};
let tracker = tracker_with_prefill_worker();
let out = selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
..Default::default()
})
.expect("worker_id should emit regardless of finish_reason");
assert!(out.worker_id.is_some());
assert!(out.timing.is_none());
assert!(out.token_ids.is_none());
assert!(out.routed_experts.is_none());
}
#[test]
fn build_response_nvext_surfaces_forwarded_split_router_worker_id() {
let selection = NvExtResponseFieldSelection {
worker_id: true,
..Default::default()
};
let tracker = tracker_with_forwarded_worker_info();
let out = selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
..Default::default()
})
.expect("forwarded worker_id should surface in nvext");
assert_eq!(
out.worker_id,
Some(WorkerIdInfo {
prefill_worker_id: Some(7),
prefill_dp_rank: Some(1),
decode_worker_id: Some(9),
decode_dp_rank: Some(2),
})
);
}
#[test]
fn build_response_nvext_timing_is_final_only() {
let selection = NvExtResponseFieldSelection {
timing: true,
..Default::default()
};
let tracker = tracker_with_prefill_worker();
assert!(
selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
..Default::default()
})
.is_none()
);
let finish_reason = FinishReason::Stop;
let out = selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
finish_reason: Some(&finish_reason),
..Default::default()
})
.expect("timing should emit on finish");
assert!(out.timing.is_some());
}
#[test]
fn build_response_nvext_token_ids_from_tracker() {
let selection = NvExtResponseFieldSelection {
token_ids: true,
..Default::default()
};
let tracker = tracker_with_query_token_ids();
let out = selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
..Default::default()
})
.expect("token_ids should emit when present");
assert_eq!(out.token_ids, Some(vec![11u32, 22, 33]));
}
#[test]
fn build_response_nvext_routed_experts_from_engine_data() {
let selection = NvExtResponseFieldSelection {
routed_experts: true,
..Default::default()
};
let engine_data = serde_json::json!({ "routed_experts": {"layer_0": [1, 3]} });
let out = selection
.build_response_nvext(NvExtResponseInput {
engine_data: Some(engine_data),
..Default::default()
})
.expect("routed_experts should emit when present");
assert_eq!(
out.routed_experts,
Some(serde_json::json!({"layer_0": [1, 3]}))
);
}
#[test]
fn build_response_nvext_completion_token_ids_pass_through() {
let selection = NvExtResponseFieldSelection {
completion_token_ids: true,
..Default::default()
};
let out = selection
.build_response_nvext(NvExtResponseInput {
completion_token_ids: Some(&[101u32, 102, 103]),
..Default::default()
})
.expect("completion_token_ids should emit when requested and present");
assert_eq!(out.completion_token_ids, Some(vec![101u32, 102, 103]));
assert!(out.prompt_logprobs.is_none());
}
#[test]
fn build_response_nvext_detailed_finish_reason_pass_through() {
let selection = NvExtResponseFieldSelection {
detailed_finish_reason: true,
..Default::default()
};
let finish_reason = FinishReason::Cancelled;
let out = selection
.build_response_nvext(NvExtResponseInput {
finish_reason: Some(&finish_reason),
..Default::default()
})
.expect("detailed_finish_reason should emit when requested and present");
assert_eq!(out.detailed_finish_reason.as_deref(), Some("cancelled"));
}
#[test]
fn build_response_nvext_prompt_token_ids_final_chunk_only() {
let selection = NvExtResponseFieldSelection {
prompt_token_ids: true,
..Default::default()
};
let tracker = tracker_with_prompt_token_ids();
assert!(
selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
..Default::default()
})
.is_none()
);
let finish_reason = FinishReason::Stop;
let out = selection
.build_response_nvext(NvExtResponseInput {
tracker: Some(&tracker),
finish_reason: Some(&finish_reason),
..Default::default()
})
.expect("prompt_token_ids should emit on the final chunk");
assert_eq!(out.prompt_token_ids, Some(vec![101u32, 102, 103]));
}
#[test]
fn build_response_nvext_prompt_logprobs_final_chunk_only() {
let selection = NvExtResponseFieldSelection {
prompt_logprobs: true,
..Default::default()
};
let mut entry = std::collections::HashMap::new();
entry.insert(
42u32,
crate::protocols::common::llm_backend::PromptLogprobEntry {
logprob: -1.234,
rank: Some(1),
decoded_token: None,
},
);
let payload: PromptLogprobs = vec![None, Some(entry)];
assert!(
selection
.build_response_nvext(NvExtResponseInput {
prompt_logprobs: Some(payload.clone()),
..Default::default()
})
.is_none()
);
let finish_reason = FinishReason::Stop;
let out = selection
.build_response_nvext(NvExtResponseInput {
finish_reason: Some(&finish_reason),
prompt_logprobs: Some(payload),
..Default::default()
})
.expect("prompt_logprobs should emit on the final chunk");
let got = out.prompt_logprobs.expect("prompt_logprobs payload");
assert_eq!(got.len(), 2);
assert!(got[0].is_none());
assert_eq!(
got[1].as_ref().unwrap().get(&42u32).unwrap().logprob,
-1.234
);
}
#[test]
fn merge_response_nvext_concatenates_completion_token_ids() {
let mut target: Option<serde_json::Value> = None;
merge_response_nvext(
&mut target,
Some(serde_json::json!({ "completion_token_ids": [10, 11, 12] })),
);
merge_response_nvext(
&mut target,
Some(serde_json::json!({ "completion_token_ids": [13, 14] })),
);
merge_response_nvext(
&mut target,
Some(serde_json::json!({
"completion_token_ids": [15],
"worker_id": { "decode_worker_id": 7 }
})),
);
let aggregated = target.expect("aggregator state");
assert_eq!(
aggregated["completion_token_ids"],
serde_json::json!([10, 11, 12, 13, 14, 15])
);
assert_eq!(aggregated["worker_id"]["decode_worker_id"], 7);
}
}