use super::dynamic_tree::ExpandedTree;
use anyhow::{ensure, Result};
pub fn walk_tree_accept(tree: &ExpandedTree, verifier_argmax: &[u32]) -> Result<Vec<usize>> {
tree.validate()?;
ensure!(
verifier_argmax.len() == tree.len(),
"walk_tree_accept: verifier_argmax len {} != tree.len() {}",
verifier_argmax.len(),
tree.len()
);
let mut accepted: Vec<usize> = Vec::with_capacity(tree.len());
accepted.push(0); let mut current: usize = 0;
loop {
let target = verifier_argmax[current];
let next: Option<usize> = (current + 1..tree.len())
.find(|&c| tree.parents[c] == Some(current) && tree.tokens[c] == target);
match next {
None => break,
Some(c) => {
accepted.push(c);
current = c;
}
}
}
Ok(accepted)
}
#[derive(Debug, Clone)]
pub struct AcceptWalk<'tree> {
accepted: Vec<usize>,
tree: &'tree ExpandedTree,
}
impl<'tree> AcceptWalk<'tree> {
pub fn accepted(&self) -> &[usize] {
&self.accepted
}
pub fn tree(&self) -> &ExpandedTree {
self.tree
}
pub fn len(&self) -> usize {
self.accepted.len()
}
pub fn is_empty(&self) -> bool {
self.accepted.is_empty()
}
pub fn drafted_accepted(&self) -> usize {
self.accepted.len().saturating_sub(1)
}
pub fn tokens(&self) -> Vec<u32> {
self.accepted.iter().map(|&i| self.tree.tokens[i]).collect()
}
pub fn max_depth(&self) -> usize {
self.accepted
.iter()
.map(|&i| self.tree.depths[i])
.max()
.unwrap_or(0)
}
}
pub fn walk_and_summarize<'a>(
tree: &'a ExpandedTree,
verifier_argmax: &[u32],
) -> Result<AcceptWalk<'a>> {
let accepted = walk_tree_accept(tree, verifier_argmax)?;
Ok(AcceptWalk { accepted, tree })
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
mod tests {
use super::*;
fn make_tree(
tokens: Vec<u32>,
parents: Vec<Option<usize>>,
depths: Vec<usize>,
) -> ExpandedTree {
let n = tokens.len();
let cum_log_probs = vec![0.0f64; n];
ExpandedTree {
tokens,
parents,
depths,
cum_log_probs,
}
}
#[test]
fn adr_037_e5a_walk_empty_accept_when_root_rejected_2026_05_22() {
let tree = make_tree(vec![100, 200], vec![None, Some(0)], vec![0, 1]);
let argmax = vec![999_u32, 0]; let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0]);
}
#[test]
fn adr_037_e5a_walk_full_chain_accept_2026_05_22() {
let tree = make_tree(
vec![10, 20, 30, 40],
vec![None, Some(0), Some(1), Some(2)],
vec![0, 1, 2, 3],
);
let argmax = vec![20_u32, 30, 40, 99];
let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0, 1, 2, 3]);
}
#[test]
fn adr_037_e5a_walk_branches_takes_matching_child_2026_05_22() {
let tree = make_tree(
vec![100, 200, 300],
vec![None, Some(0), Some(0)],
vec![0, 1, 1],
);
let argmax = vec![300_u32, 0, 0]; let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0, 2]);
}
#[test]
fn adr_037_e5a_walk_partial_chain_then_no_match_2026_05_22() {
let tree = make_tree(
vec![10, 20, 30, 40],
vec![None, Some(0), Some(1), Some(1)],
vec![0, 1, 2, 2],
);
let argmax = vec![20_u32, 999, 0, 0];
let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0, 1]);
}
#[test]
fn adr_037_e5a_walk_rejects_size_mismatch_2026_05_22() {
let tree = make_tree(vec![10, 20], vec![None, Some(0)], vec![0, 1]);
let argmax = vec![20_u32]; let err = walk_tree_accept(&tree, &argmax).unwrap_err();
assert!(
err.to_string().contains("verifier_argmax len"),
"got: {err}"
);
}
#[test]
fn adr_037_e5a_walk_asymmetric_tree_picks_deepest_matching_path_2026_05_22() {
let tree = make_tree(
vec![1, 10, 20, 100],
vec![None, Some(0), Some(0), Some(1)],
vec![0, 1, 1, 2],
);
let argmax = vec![10_u32, 100, 0, 0];
let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0, 1, 3]);
}
#[test]
fn adr_037_e5a_walk_returns_only_root_for_single_node_tree_2026_05_22() {
let tree = make_tree(vec![42], vec![None], vec![0]);
let argmax = vec![999_u32];
let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0]);
}
#[test]
fn adr_037_e5a_summary_accessor_helpers_2026_05_22() {
let tree = make_tree(
vec![10, 20, 30, 40],
vec![None, Some(0), Some(1), Some(2)],
vec![0, 1, 2, 3],
);
let argmax = vec![20_u32, 30, 40, 99];
let summary = walk_and_summarize(&tree, &argmax).expect("walk");
assert_eq!(summary.len(), 4);
assert!(!summary.is_empty());
assert_eq!(summary.drafted_accepted(), 3);
assert_eq!(summary.tokens(), vec![10, 20, 30, 40]);
assert_eq!(summary.max_depth(), 3);
}
#[test]
fn adr_037_e5a_summary_root_only_2026_05_22() {
let tree = make_tree(vec![100, 200], vec![None, Some(0)], vec![0, 1]);
let argmax = vec![999_u32, 0];
let summary = walk_and_summarize(&tree, &argmax).expect("walk");
assert_eq!(summary.len(), 1);
assert_eq!(summary.drafted_accepted(), 0); assert_eq!(summary.tokens(), vec![100]);
assert_eq!(summary.max_depth(), 0);
}
#[test]
fn adr_037_e5a_walk_rejects_corrupt_tree_2026_05_22() {
let tree = ExpandedTree {
tokens: vec![1, 2],
parents: vec![None, Some(5)],
depths: vec![0, 1],
cum_log_probs: vec![0.0, 0.0],
};
let argmax = vec![2_u32, 0];
assert!(walk_tree_accept(&tree, &argmax).is_err());
}
#[test]
fn adr_037_e5a_walk_skips_non_root_descendants_in_search_2026_05_22() {
let tree = make_tree(
vec![1, 10, 30, 20],
vec![None, Some(0), Some(1), Some(0)],
vec![0, 1, 2, 1],
);
let argmax = vec![30_u32, 0, 0, 0];
let accepted = walk_tree_accept(&tree, &argmax).expect("walk");
assert_eq!(accepted, vec![0]);
}
#[test]
fn adr_037_e5a_walk_integration_with_phase_e4a_expand_dynamic_tree_2026_05_22() {
use crate::inference::spec_decode::eagle3::drafter::{
DraftCandidate, Drafter, TreeContextView,
};
use crate::inference::spec_decode::eagle3::dynamic_tree::{
expand_dynamic_tree, DynamicTreeConfig,
};
struct ScriptedDrafter;
impl Drafter for ScriptedDrafter {
fn predict_topk(
&mut self,
_tree: TreeContextView<'_>,
node: usize,
_top_k: usize,
) -> Result<Vec<DraftCandidate>> {
Ok(vec![
DraftCandidate {
token: (node * 10 + 1) as u32,
log_prob: -0.1,
},
DraftCandidate {
token: (node * 10 + 2) as u32,
log_prob: -0.5,
},
])
}
}
let cfg = DynamicTreeConfig {
budget: 6,
max_depth: 3,
top_k: 2,
};
let mut d = ScriptedDrafter;
let tree = expand_dynamic_tree(1000_u32, &mut d, &cfg).expect("expand");
let mut verifier_argmax = vec![0_u32; tree.len()];
for i in 0..tree.len() {
if let Some(child) = (i + 1..tree.len()).find(|&c| tree.parents[c] == Some(i)) {
verifier_argmax[i] = tree.tokens[child];
}
}
let accepted = walk_tree_accept(&tree, &verifier_argmax).expect("walk");
assert_eq!(accepted[0], 0);
assert!(accepted.len() >= 2, "walk should accept beyond root");
for w in accepted.windows(2) {
let parent = w[0];
let child = w[1];
assert_eq!(
tree.parents[child],
Some(parent),
"accept walk should follow direct-child edges"
);
assert_eq!(
tree.tokens[child], verifier_argmax[parent],
"accept walk should match verifier_argmax at each parent"
);
}
let last = *accepted.last().unwrap();
let no_matching_child = (last + 1..tree.len())
.find(|&c| tree.parents[c] == Some(last) && tree.tokens[c] == verifier_argmax[last])
.is_none();
assert!(
no_matching_child,
"walk should have stopped because no direct child of {} matches argmax {}",
last, verifier_argmax[last]
);
}
}