1//! Choice over more than 255 labels by walking a label tree (contract
2//! *SQL surface*, "More than 255 options", label tree mode).
3//!
4//! Each internal node is one Choice over its children. A node's score is
5//! the product of the edge probabilities from the root, kept as
6//! `Σ ln max(p, 1e-9)`. That is a real probability, and it can only fall
7//! with depth, so a scored node bounds every leaf below it. The
8//! cookbook's geometric mean `exp(mean(ln p))` is not used: it rises with
9//! depth and prefers deep paths of confident edges over a shallow likely
10//! one.
11//!
12//! The search is best-first, K nodes per round (K = 3 by default): each
13//! round asks the Choices of the K best unexpanded nodes, which the
14//! caller sends in one request. Nothing is dropped for being outside the
15//! K; a node is dropped only when its bound shows it cannot change the
16//! result. So the answer is exact: the best leaf, and the runner-up the
17//! separation `exp(top − second)` is measured against.
18//!
19//! With `tau`, the search returns the deepest node whose path probability
20//! is at least τ (ties at one depth to the likelier), and never expands a
21//! node below τ. Since sibling probabilities share their parent's, at
22//! most `1/τ` nodes per depth qualify.
23//!
24//! A node with one child is not a question (`jev_protocol::Choice`): its
25//! edge is ×1 and it is passed through without a request.
26//!
27//! Two things ride in a round beside its K nodes, and neither changes
28//! the result:
29//!
30//! - **Lookahead** ([`Lookahead`]): the Choices of each asked node's
31//!   likeliest children, within a token budget, so one round trip can
32//!   score two levels. An answer that arrives before its node is scored
33//!   is held, and applied without a request once it is.
34//! - **The order twin** ([`TwinBelow`]): a node whose top two children
35//!   come back close is not expanded; it is queued again and asked in
36//!   reverse order, and the two answers' log-probabilities are averaged.
37//!   Closeness is only known from the answer, so the twin goes in a
38//!   later round, in its turn by score like any node: it costs a round
39//!   trip only when nothing else is left to ask, and nothing when its
40//!   node is pruned first.
41//!
42//! Sans-I/O: [`Search::step`] says which Choices to ask, [`Tree::choice`]
43//! builds each one (options describing their subtree, instructions
44//! naming the path, in [`Ask::order`]), and [`Search::answer`] takes its
45//! probabilities in the order sent. Sending them is the caller's.
46
47use std::collections::{HashMap, HashSet};
48use std::num::NonZeroUsize;
49
50use jev_protocol::{Choice, Json, MAX_CHOICE_OPTIONS, Noul};
51
52/// The floor under an edge probability, so a 0 costs `ln 1e-9` rather
53/// than sending a path to −∞ and making every such path tie.
54pub const PROBABILITY_FLOOR: f64 = 1e-9;
55
56/// The vendor's beam width (hierarchical-classification cookbook); not
57/// measured by us.
58pub const DEFAULT_K: NonZeroUsize = NonZeroUsize::new(3).unwrap();
59
60/// A node of a [`Tree`], valid only for the tree that issued it.
61#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
62pub struct NodeId(u32);
63
64/// A label hierarchy as the caller writes it.
65#[derive(Clone, Debug, PartialEq)]
66pub struct Branch {
67    pub label: String,
68    pub children: Vec<Branch>,
69}
70
71impl Branch {
72    pub fn leaf(label: impl Into<String>) -> Self {
73        Branch { label: label.into(), children: Vec::new() }
74    }
75
76    pub fn node(label: impl Into<String>, children: impl IntoIterator<Item = Branch>) -> Self {
77        Branch { label: label.into(), children: children.into_iter().collect() }
78    }
79}
80
81impl Branch {
82    /// A root over `paths`, each a sequence of labels from the top down.
83    /// Siblings keep the order in which their label first appears, since
84    /// a node's options are sent in that order. Every path ends at a leaf,
85    /// so a path that is a prefix of another is refused rather than
86    /// silently dropped from the answers a search can return.
87    pub fn from_paths<P, L>(paths: P) -> Result<Branch, PathError>
88    where
89        P: IntoIterator<Item = L>,
90        L: IntoIterator,
91        L::Item: Into<String>,
92    {
93        let mut root = Branch::leaf("");
94        let mut ends: HashSet<Vec<String>> = HashSet::new();
95        for path in paths {
96            let path: Vec<String> = path.into_iter().map(Into::into).collect();
97            if path.is_empty() {
98                return Err(PathError::Empty);
99            }
100            let mut at = &mut root;
101            for (i, label) in path.iter().enumerate() {
102                if ends.contains(&path[..i]) && i > 0 {
103                    return Err(PathError::LeafAndGroup { path: path[..i].to_vec() });
104                }
105                let c = match at.children.iter().position(|c| c.label == *label) {
106                    Some(c) => c,
107                    None => {
108                        at.children.push(Branch::leaf(label.clone()));
109                        at.children.len() - 1
110                    }
111                };
112                at = &mut at.children[c];
113            }
114            if !at.children.is_empty() {
115                return Err(PathError::LeafAndGroup { path });
116            }
117            if !ends.insert(path.clone()) {
118                return Err(PathError::Repeated { path });
119            }
120        }
121        Ok(root)
122    }
123}
124
125/// Paths that do not describe a tree of leaves.
126#[derive(Clone, Debug, PartialEq, Eq)]
127pub enum PathError {
128    Empty,
129    Repeated { path: Vec<String> },
130    /// A path that also groups other paths below it.
131    LeafAndGroup { path: Vec<String> },
132}
133
134impl std::fmt::Display for PathError {
135    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
136        match self {
137            PathError::Empty => write!(f, "a label path is empty"),
138            PathError::Repeated { path } => write!(f, "the label path {path:?} is given twice"),
139            PathError::LeafAndGroup { path } => {
140                write!(f, "the label path {path:?} is both a label and the parent of other labels")
141            }
142        }
143    }
144}
145
146impl std::error::Error for PathError {}
147
148#[derive(Clone, Debug, PartialEq, Eq)]
149pub enum TreeError {
150    /// The root has no children, so there is nothing to choose.
151    NoChoice,
152    /// More children than one Choice can carry.
153    Fanout { path: Vec<String>, children: usize },
154    EmptyLabel { path: Vec<String> },
155    DuplicateLabel { path: Vec<String>, label: String },
156}
157
158impl std::fmt::Display for TreeError {
159    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
160        match self {
161            TreeError::NoChoice => write!(f, "the label tree has no labels under its root"),
162            TreeError::Fanout { path, children } => write!(
163                f,
164                "the label tree node {path:?} has {children} children; a Choice takes at most {MAX_CHOICE_OPTIONS}"
165            ),
166            TreeError::EmptyLabel { path } => write!(f, "a label under {path:?} is empty"),
167            TreeError::DuplicateLabel { path, label } => {
168                write!(f, "the label {label:?} appears twice under {path:?}")
169            }
170        }
171    }
172}
173
174impl std::error::Error for TreeError {}
175
176#[derive(Clone, Debug)]
177struct Node {
178    label: String,
179    parent: Option<NodeId>,
180    children: Vec<NodeId>,
181    depth: u32,
182    /// Edges from this node down to its deepest leaf.
183    height: u32,
184    /// Leaves under this node; a leaf counts itself.
185    leaves: u32,
186}
187
188/// A validated label hierarchy: every node has at most 255 children with
189/// unique, non-empty labels, so every node's Choice can be built.
190#[derive(Clone, Debug)]
191pub struct Tree {
192    nodes: Vec<Node>,
193}
194
195impl Tree {
196    /// The root's own label is never sent; its children are the first
197    /// Choice.
198    pub fn new(root: Branch) -> Result<Self, TreeError> {
199        if root.children.is_empty() {
200            return Err(TreeError::NoChoice);
201        }
202        let mut tree = Tree { nodes: Vec::new() };
203        tree.add(root, None, 0)?;
204        Ok(tree)
205    }
206
207    fn add(&mut self, b: Branch, parent: Option<NodeId>, depth: u32) -> Result<NodeId, TreeError> {
208        let id = NodeId(self.nodes.len() as u32);
209        self.nodes.push(Node { label: b.label, parent, children: Vec::new(), depth, height: 0, leaves: 1 });
210        if b.children.len() > MAX_CHOICE_OPTIONS {
211            return Err(TreeError::Fanout { path: self.path_owned(id), children: b.children.len() });
212        }
213        let mut seen = HashSet::new();
214        for c in &b.children {
215            if c.label.is_empty() {
216                return Err(TreeError::EmptyLabel { path: self.path_owned(id) });
217            }
218            if !seen.insert(c.label.clone()) {
219                return Err(TreeError::DuplicateLabel { path: self.path_owned(id), label: c.label.clone() });
220            }
221        }
222        let (mut height, mut leaves) = (0, 0);
223        for c in b.children {
224            let child = self.add(c, Some(id), depth + 1)?;
225            self.nodes[id.0 as usize].children.push(child);
226            height = height.max(self.nodes[child.0 as usize].height + 1);
227            leaves += self.nodes[child.0 as usize].leaves;
228        }
229        self.nodes[id.0 as usize].height = height;
230        self.nodes[id.0 as usize].leaves = leaves.max(1);
231        Ok(id)
232    }
233
234    fn node(&self, n: NodeId) -> &Node {
235        &self.nodes[n.0 as usize]
236    }
237
238    pub fn root(&self) -> NodeId {
239        NodeId(0)
240    }
241
242    pub fn label(&self, n: NodeId) -> &str {
243        &self.node(n).label
244    }
245
246    /// In the caller's order, which is the order the Choice sends.
247    pub fn children(&self, n: NodeId) -> &[NodeId] {
248        &self.node(n).children
249    }
250
251    pub fn is_leaf(&self, n: NodeId) -> bool {
252        self.node(n).children.is_empty()
253    }
254
255    /// Edges from the root; the root is 0.
256    pub fn depth(&self, n: NodeId) -> u32 {
257        self.node(n).depth
258    }
259
260    /// The labels from below the root down to `n`.
261    pub fn path(&self, n: NodeId) -> Vec<&str> {
262        let mut path = Vec::new();
263        let mut at = n;
264        while let Some(parent) = self.node(at).parent {
265            path.push(self.label(at));
266            at = parent;
267        }
268        path.reverse();
269        path
270    }
271
272    fn path_owned(&self, n: NodeId) -> Vec<String> {
273        self.path(n).into_iter().map(str::to_owned).collect()
274    }
275
276    /// The leaves under `n`, depth first in the caller's order.
277    fn leaves(&self, n: NodeId) -> Vec<NodeId> {
278        let mut out = Vec::new();
279        let mut stack = vec![n];
280        while let Some(at) = stack.pop() {
281            if self.is_leaf(at) {
282                out.push(at);
283            }
284            stack.extend(self.children(at).iter().rev());
285        }
286        out
287    }
288
289    /// `n`'s Choice over its children, in `order` (the caller's, or
290    /// reversed for an order-debiasing twin), each label the option key
291    /// (contract *SQL surface*). The instructions are
292    /// `question` and, below the root, the path to `n`; an internal
293    /// child's description lists its direct children and a sample of the
294    /// leaves under it, so the model sees what lives under a branch
295    /// (vendor guidance; its hierarchical cookbook does neither). A leaf
296    /// child has no description. `None` for a node with fewer than two
297    /// children, which [`Search`] never asks.
298    ///
299    /// Every piece is written in a fixed form, so equal nodes give equal
300    /// bytes and one cache key.
301    pub fn choice(&self, n: NodeId, order: Order, question: &str, describe: Describe) -> Option<Choice> {
302        let children = self.children(n);
303        if children.len() < 2 {
304            return None;
305        }
306        let path = self.path(n);
307        let instructions = if path.is_empty() {
308            question.to_owned()
309        } else {
310            format!("{question}\n\nWithin: {}", path.join(" > "))
311        };
312        let option = |&c: &NodeId| (self.label(c).to_owned(), self.describe(c, describe).map(|d| Json::text(&d)));
313        let options: Vec<_> = match order {
314            Order::Caller => children.iter().map(option).collect(),
315            Order::Reversed => children.iter().rev().map(option).collect(),
316        };
317        Some(Choice::new(Json::text(&instructions), options).expect("a Tree node's children make a valid Choice"))
318    }
319
320    /// The "does any label fit" Noul ([`crate::gate`]), with the tree shown
321    /// as the root is described: its top-level labels and a sample of
322    /// leaves, so its size does not grow with the tree. `None` for a tree
323    /// with one leaf, which the search answers without a request.
324    pub fn gate(&self, question: &str, describe: Describe) -> Option<Noul> {
325        let root = self.root();
326        if self.leaves(root).len() < 2 {
327            return None;
328        }
329        let labels = self.describe(root, describe).expect("a root with children is described");
330        Some(crate::gate::gate(question, &labels))
331    }
332
333    fn describe(&self, n: NodeId, describe: Describe) -> Option<String> {
334        let children = self.children(n);
335        if children.is_empty() {
336            return None;
337        }
338        let listed = children.len().min(describe.children);
339        let mut text = format!(
340            "Contains: {}",
341            children[..listed].iter().map(|&c| self.label(c)).collect::<Vec<_>>().join(", ")
342        );
343        if listed < children.len() {
344            text.push_str(&format!(", and {} more", children.len() - listed));
345        }
346        // Leaves that are not already listed as direct children.
347        let sample: Vec<&str> = self
348            .leaves(n)
349            .into_iter()
350            .filter(|l| self.node(*l).parent != Some(n))
351            .take(describe.leaves)
352            .map(|l| self.label(l))
353            .collect();
354        if !sample.is_empty() {
355            text.push_str(&format!(". For example: {}", sample.join(", ")));
356        }
357        Some(text)
358    }
359}
360
361/// How much of a subtree an option's description shows. The defaults
362/// are ours and unmeasured; the label-tree release gate (*Testing*)
363/// measures them with K.
364#[derive(Clone, Copy, Debug, PartialEq, Eq)]
365pub struct Describe {
366    /// Direct children listed, in the caller's order.
367    pub children: usize,
368    /// Leaves below those children, depth first.
369    pub leaves: usize,
370}
371
372impl Default for Describe {
373    fn default() -> Self {
374        Describe { children: 20, leaves: 5 }
375    }
376}
377
378/// The order a Choice's options are sent in.
379#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
380pub enum Order {
381    Caller,
382    /// The order-debiasing twin's.
383    Reversed,
384}
385
386/// A probability threshold for early stop, in (0, 1].
387#[derive(Clone, Copy, Debug, PartialEq)]
388pub struct Tau(f64);
389
390impl Tau {
391    pub fn new(tau: f64) -> Option<Self> {
392        (tau > 0.0 && tau <= 1.0).then_some(Tau(tau))
393    }
394
395    pub fn get(self) -> f64 {
396        self.0
397    }
398}
399
400/// Speculative Choices carried in a round: for each node asked, the
401/// Choices of its `m` likeliest children, while the round's estimated
402/// tokens stay within `budget`. So one round trip scores two levels.
403///
404/// "Likeliest" is judged before the parent is answered, since both ride
405/// in one request, so it is a prior: the children with the most leaves
406/// under them (a uniform prior over leaves), ties in the caller's order.
407/// `m` and the prior are unmeasured; the label-tree release gate
408/// (*Testing*) measures them.
409#[derive(Clone, Copy, Debug, PartialEq, Eq)]
410pub struct Lookahead {
411    pub m: usize,
412    /// Estimated input tokens for the whole round, the asked nodes'
413    /// Choices included. The asked nodes are sent even when they alone
414    /// exceed it.
415    pub budget: usize,
416}
417
418impl Default for Lookahead {
419    /// Below jev-1.13.0's 64k tokens per request, leaving room for the
420    /// request's own overhead.
421    fn default() -> Self {
422        Lookahead { m: 2, budget: 60_000 }
423    }
424}
425
426/// A node whose top two children are within this separation
427/// (`p_top / p_second`, floored) is a close call, and its Choice is asked
428/// again in reverse order; the two answers' log-probabilities are
429/// averaged (permutation self-consistency, arXiv 2310.07712), since the
430/// measured order bias favours whatever is listed last. Greater than 1.
431#[derive(Clone, Copy, Debug, PartialEq)]
432pub struct TwinBelow(f64);
433
434impl TwinBelow {
435    pub fn new(separation: f64) -> Option<Self> {
436        (separation > 1.0 && separation.is_finite()).then_some(TwinBelow(separation))
437    }
438
439    pub fn get(self) -> f64 {
440        self.0
441    }
442}
443
444#[derive(Clone, Copy, Debug, PartialEq)]
445pub struct Params {
446    /// Scored nodes asked per round, all in one request. A twin takes a
447    /// slot, since it is its node's own question; lookahead Choices ride
448    /// outside it.
449    pub k: NonZeroUsize,
450    /// Unset: always return a leaf.
451    pub tau: Option<Tau>,
452    /// Unset: no speculative Choices.
453    pub lookahead: Option<Lookahead>,
454    /// Unset: no twins.
455    pub twin: Option<TwinBelow>,
456}
457
458impl Default for Params {
459    /// The twin's 1.5× is ours and unmeasured, like [`Lookahead`]'s
460    /// defaults.
461    fn default() -> Self {
462        Params { k: DEFAULT_K, tau: None, lookahead: Some(Lookahead::default()), twin: TwinBelow::new(1.5) }
463    }
464}
465
466/// One Choice a round asks. Build it with [`Tree::choice`] in
467/// [`Ask::order`], and answer it with [`Search::answer`].
468#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
469pub enum Ask {
470    /// A scored node, best first.
471    Node(NodeId),
472    /// A close call's second asking, in reverse order.
473    Twin(NodeId),
474    /// A child of a node asked this round, before its score is known.
475    Lookahead(NodeId),
476}
477
478impl Ask {
479    pub fn node(self) -> NodeId {
480        match self {
481            Ask::Node(n) | Ask::Twin(n) | Ask::Lookahead(n) => n,
482        }
483    }
484
485    pub fn order(self) -> Order {
486        match self {
487            Ask::Twin(_) => Order::Reversed,
488            Ask::Node(_) | Ask::Lookahead(_) => Order::Caller,
489        }
490    }
491}
492
493/// What the search needs next.
494#[derive(Clone, Debug, PartialEq)]
495pub enum Step {
496    /// Ask these Choices, in one request, and [`Search::answer`] each.
497    Ask(Vec<Ask>),
498    Done(Outcome),
499}
500
501#[derive(Clone, Debug, PartialEq)]
502pub enum Outcome {
503    /// The most probable leaf.
504    Leaf {
505        leaf: NodeId,
506        probability: f64,
507        /// `exp(score_top − score_second)`: near 1 is ambiguous. `None`
508        /// when the tree admits one leaf only.
509        separation: Option<f64>,
510    },
511    /// Under `tau`: the deepest node at or above τ.
512    Stop { node: NodeId, depth: u32, leaf: bool, probability: f64 },
513}
514
515#[derive(Clone, Debug, PartialEq)]
516pub enum AnswerError {
517    /// Not asked in the current round, or already answered.
518    NotAsked(Ask),
519    /// One probability per option, in the order sent.
520    Count { node: NodeId, expected: usize, got: usize },
521    NotAProbability { node: NodeId, value: f64 },
522}
523
524impl std::fmt::Display for AnswerError {
525    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
526        match self {
527            AnswerError::NotAsked(a) => write!(f, "label tree Choice {a:?} was not asked"),
528            AnswerError::Count { node, expected, got } => {
529                write!(f, "label tree node {node:?} has {expected} options, answered with {got}")
530            }
531            AnswerError::NotAProbability { node, value } => {
532                write!(f, "label tree node {node:?} was answered with {value}, not a probability")
533            }
534        }
535    }
536}
537
538impl std::error::Error for AnswerError {}
539
540#[derive(Clone, Copy, Debug)]
541struct Entry {
542    score: f64,
543    node: NodeId,
544}
545
546/// One label-tree search: best-first, exact, K nodes per round.
547///
548/// Lookahead and twins change what a round carries, never the result:
549/// an answer that arrives early is held until its node is scored, and
550/// the node is then expanded without a request. The pruning is the same
551/// bound either way.
552#[derive(Debug)]
553pub struct Search<'t> {
554    tree: &'t Tree,
555    params: Params,
556    /// Scored and not yet expanded or finalized, best first. A queued
557    /// node with an entry in `answers` is a close call awaiting its twin.
558    queue: Vec<Entry>,
559    /// Asked this round; `score` is `None` for a lookahead, whose node is
560    /// not scored yet.
561    pending: Vec<(Ask, Option<f64>)>,
562    /// Answers not yet applied, in child order: a lookahead's, until its
563    /// node is scored, and a close call's first, until its twin.
564    answers: HashMap<NodeId, Vec<f64>>,
565    /// Leaves proven best, in order (no `tau`).
566    finalized: Vec<Entry>,
567    /// The deepest node at or above τ so far (`tau`).
568    stop: Option<Entry>,
569}
570
571impl<'t> Search<'t> {
572    pub fn new(tree: &'t Tree, params: Params) -> Self {
573        let mut s = Search {
574            tree,
575            params,
576            queue: Vec::new(),
577            pending: Vec::new(),
578            answers: HashMap::new(),
579            finalized: Vec::new(),
580            stop: None,
581        };
582        s.push(Entry { score: 0.0, node: tree.root() });
583        s
584    }
585
586    fn ln_tau(&self) -> Option<f64> {
587        self.params.tau.map(|t| t.get().ln())
588    }
589
590    /// Scores a node, passing through one-child nodes at ×1, and expands
591    /// it at once if its answer is already held.
592    fn push(&mut self, mut e: Entry) {
593        loop {
594            if let Some(ln_tau) = self.ln_tau() {
595                if e.score < ln_tau {
596                    return;
597                }
598                if self.better_stop(e) {
599                    self.stop = Some(e);
600                }
601            }
602            match self.tree.children(e.node) {
603                [only] => e.node = *only,
604                _ => break,
605            }
606        }
607        if self.params.tau.is_some() && self.tree.is_leaf(e.node) {
608            return;
609        }
610        if let Some(first) = self.answers.get(&e.node)
611            && !self.is_close(first)
612        {
613            let first = self.answers.remove(&e.node).expect("just read");
614            let ln: Vec<f64> = first.iter().map(|p| ln_floored(*p)).collect();
615            return self.expand(e, &ln);
616        }
617        self.enqueue(e);
618    }
619
620    fn enqueue(&mut self, e: Entry) {
621        let at = self.queue.partition_point(|q| rank(q, &e).is_lt());
622        self.queue.insert(at, e);
623    }
624
625    fn expand(&mut self, parent: Entry, ln: &[f64]) {
626        let children = self.tree.children(parent.node);
627        for (&child, &l) in children.iter().zip(ln) {
628            self.push(Entry { score: parent.score + l, node: child });
629        }
630    }
631
632    /// Whether this first answer, in child order, needs a twin.
633    fn is_close(&self, probabilities: &[f64]) -> bool {
634        let Some(twin) = self.params.twin else { return false };
635        let (mut top, mut second) = (0.0f64, 0.0f64);
636        for &p in probabilities {
637            let p = p.max(PROBABILITY_FLOOR);
638            if p > top {
639                (top, second) = (p, top);
640            } else if p > second {
641                second = p;
642            }
643        }
644        top / second <= twin.get()
645    }
646
647    fn better_stop(&self, e: Entry) -> bool {
648        self.stop.is_none_or(|s| {
649            let (d, sd) = (self.tree.depth(e.node), self.tree.depth(s.node));
650            d > sd || (d == sd && e.score > s.score)
651        })
652    }
653
654    /// Whether anything under `e` could replace the current stop.
655    fn could_beat_stop(&self, e: &Entry) -> bool {
656        self.stop.is_none_or(|s| {
657            let reach = self.tree.depth(e.node) + self.tree.node(e.node).height;
658            let sd = self.tree.depth(s.node);
659            reach > sd || (reach == sd && e.score > s.score)
660        })
661    }
662
663    fn round(&self) -> Step {
664        Step::Ask(self.pending.iter().map(|(a, _)| *a).collect())
665    }
666
667    /// The next round to ask, or the result. Asking again before every
668    /// asked Choice is answered returns the same round.
669    ///
670    /// `tokens` estimates one Choice's input tokens, for the lookahead
671    /// budget; it is not called without [`Params::lookahead`].
672    pub fn step(&mut self, tokens: impl FnMut(Ask) -> usize) -> Step {
673        if !self.pending.is_empty() {
674            return self.round();
675        }
676        if self.params.tau.is_some() {
677            return self.step_tau(tokens);
678        }
679        // A leaf at the head of the queue beats every leaf still unscored,
680        // because scores only fall with depth.
681        while let Some(head) = self.queue.first() {
682            if !self.tree.is_leaf(head.node) {
683                break;
684            }
685            self.finalized.push(self.queue.remove(0));
686            if self.finalized.len() == 2 {
687                return Step::Done(self.leaf_outcome());
688            }
689        }
690        if self.queue.is_empty() {
691            return Step::Done(self.leaf_outcome());
692        }
693        // Only the best two leaves matter: a node no likelier than the
694        // second-best known leaf cannot change them.
695        let needed = 2 - self.finalized.len();
696        let bound = self.queue.iter().filter(|e| self.tree.is_leaf(e.node)).nth(needed - 1).map(|e| e.score);
697        if let Some(bound) = bound {
698            let tree = self.tree;
699            self.queue.retain(|e| tree.is_leaf(e.node) || e.score > bound);
700        }
701        self.take_round(|tree, e| !tree.is_leaf(e.node), tokens)
702    }
703
704    fn step_tau(&mut self, tokens: impl FnMut(Ask) -> usize) -> Step {
705        let queue = std::mem::take(&mut self.queue);
706        self.queue = queue.into_iter().filter(|e| self.could_beat_stop(e)).collect();
707        if self.queue.is_empty() {
708            let s = self.stop.expect("the root is always at or above τ");
709            return Step::Done(Outcome::Stop {
710                node: s.node,
711                depth: self.tree.depth(s.node),
712                leaf: self.tree.is_leaf(s.node),
713                probability: s.score.exp(),
714            });
715        }
716        self.take_round(|_, _| true, tokens)
717    }
718
719    fn take_round(&mut self, askable: impl Fn(&Tree, &Entry) -> bool, mut tokens: impl FnMut(Ask) -> usize) -> Step {
720        let k = self.params.k.get();
721        let mut i = 0;
722        while i < self.queue.len() && self.pending.len() < k {
723            if askable(self.tree, &self.queue[i]) {
724                let e = self.queue.remove(i);
725                let ask = if self.answers.contains_key(&e.node) { Ask::Twin(e.node) } else { Ask::Node(e.node) };
726                self.pending.push((ask, Some(e.score)));
727            } else {
728                i += 1;
729            }
730        }
731        if let Some(look) = self.params.lookahead {
732            self.add_lookahead(look, &mut tokens);
733        }
734        self.round()
735    }
736
737    /// Fills the round with the asked nodes' likeliest children's Choices,
738    /// best-ranked asked node first, within the budget.
739    fn add_lookahead(&mut self, look: Lookahead, tokens: &mut impl FnMut(Ask) -> usize) {
740        let mut used: usize = self.pending.iter().map(|(a, _)| tokens(*a)).sum();
741        let parents: Vec<NodeId> = self.pending.iter().map(|(a, _)| a.node()).collect();
742        for parent in parents {
743            let mut children = self.tree.children(parent).to_vec();
744            // Stable, so ties keep the caller's order.
745            children.sort_by_key(|&c| std::cmp::Reverse(self.tree.node(c).leaves));
746            for c in children.into_iter().take(look.m) {
747                let Some(n) = self.askable_below(c) else { continue };
748                if self.answers.contains_key(&n) || self.pending.iter().any(|(a, _)| a.node() == n) {
749                    continue;
750                }
751                let ask = Ask::Lookahead(n);
752                let cost = tokens(ask);
753                if used + cost > look.budget {
754                    continue;
755                }
756                used += cost;
757                self.pending.push((ask, None));
758            }
759        }
760    }
761
762    /// The node a Choice would be asked for once `n` is scored: `n`
763    /// passed through its one-child chain, or `None` at a leaf.
764    fn askable_below(&self, mut n: NodeId) -> Option<NodeId> {
765        loop {
766            match self.tree.children(n) {
767                [] => return None,
768                [only] => n = *only,
769                _ => return Some(n),
770            }
771        }
772    }
773
774    fn leaf_outcome(&self) -> Outcome {
775        let top = self.finalized[0];
776        Outcome::Leaf {
777            leaf: top.node,
778            probability: top.score.exp(),
779            separation: self.finalized.get(1).map(|second| (top.score - second.score).exp()),
780        }
781    }
782
783    /// A Choice's answer: one probability per option, in the order it was
784    /// sent ([`Ask::order`]).
785    pub fn answer(&mut self, ask: Ask, probabilities: &[f64]) -> Result<(), AnswerError> {
786        let at = self.pending.iter().position(|(a, _)| *a == ask).ok_or(AnswerError::NotAsked(ask))?;
787        let node = ask.node();
788        let children = self.tree.children(node);
789        if probabilities.len() != children.len() {
790            return Err(AnswerError::Count { node, expected: children.len(), got: probabilities.len() });
791        }
792        if let Some(&value) = probabilities.iter().find(|p| !(0.0..=1.0).contains(*p)) {
793            return Err(AnswerError::NotAProbability { node, value });
794        }
795        let mut in_child_order = probabilities.to_vec();
796        if ask.order() == Order::Reversed {
797            in_child_order.reverse();
798        }
799        let (_, score) = self.pending.remove(at);
800        match ask {
801            Ask::Node(_) => {
802                let e = Entry { score: score.expect("an asked node is scored"), node };
803                if self.is_close(&in_child_order) {
804                    self.answers.insert(node, in_child_order);
805                    self.enqueue(e);
806                } else {
807                    let ln: Vec<f64> = in_child_order.iter().map(|p| ln_floored(*p)).collect();
808                    self.expand(e, &ln);
809                }
810            }
811            Ask::Twin(_) => {
812                let e = Entry { score: score.expect("a twin's node is scored"), node };
813                let first = self.answers.remove(&node).expect("a twin follows a first answer");
814                let ln: Vec<f64> =
815                    first.iter().zip(&in_child_order).map(|(a, b)| (ln_floored(*a) + ln_floored(*b)) / 2.0).collect();
816                self.expand(e, &ln);
817            }
818            Ask::Lookahead(_) => {
819                self.answers.insert(node, in_child_order);
820                // Its parent may have been answered first in this round.
821                if let Some(i) = self.queue.iter().position(|e| e.node == node) {
822                    let e = self.queue.remove(i);
823                    self.push(e);
824                }
825            }
826        }
827        Ok(())
828    }
829}
830
831fn ln_floored(p: f64) -> f64 {
832    p.max(PROBABILITY_FLOOR).ln()
833}
834
835/// Best first; equal scores in tree order, so a search is deterministic.
836fn rank(a: &Entry, b: &Entry) -> std::cmp::Ordering {
837    b.score.total_cmp(&a.score).then(a.node.cmp(&b.node))
838}
839
840#[cfg(test)]
841mod tests {
842    use super::*;
843
844    /// Runs a search against a fixed answer per internal node, by path
845    /// and in child order (reversed for a twin), and returns the rounds.
846    fn run_rounds(
847        tree: &Tree,
848        params: Params,
849        answer: impl Fn(&[&str], Order) -> Vec<f64>,
850    ) -> (Outcome, Vec<Vec<Ask>>) {
851        let mut s = Search::new(tree, params);
852        let mut rounds = Vec::new();
853        loop {
854            match s.step(|_| 100) {
855                Step::Done(o) => return (o, rounds),
856                Step::Ask(round) => {
857                    assert!(!round.is_empty());
858                    assert!(round.iter().filter(|a| !matches!(a, Ask::Lookahead(_))).count() <= params.k.get());
859                    // Lookahead first, so a held answer is applied when its
860                    // parent's lands, as well as the other way round.
861                    for &a in round.iter().rev() {
862                        s.answer(a, &answer(&tree.path(a.node()), a.order())).unwrap();
863                    }
864                    rounds.push(round);
865                }
866            }
867        }
868    }
869
870    /// An answer by path that is the same in either order.
871    fn fixed(answer: impl Fn(&[&str]) -> Vec<f64>) -> impl Fn(&[&str], Order) -> Vec<f64> {
872        move |path, order| {
873            let mut p = answer(path);
874            if order == Order::Reversed {
875                p.reverse();
876            }
877            p
878        }
879    }
880
881    /// The search alone: no lookahead, no twins.
882    fn plain() -> Params {
883        Params { k: DEFAULT_K, tau: None, lookahead: None, twin: None }
884    }
885
886    #[test]
887    fn the_gate_shows_the_tree_as_the_root_is_described() {
888        let tree = Tree::new(Branch::from_paths([vec!["A", "a1"], vec!["A", "a2"], vec!["B"]]).unwrap()).unwrap();
889        let want = crate::gate::gate("q", "Contains: A, B. For example: a1, a2");
890        assert_eq!(tree.gate("q", Describe::default()), Some(want));
891        let one = Tree::new(Branch::from_paths([vec!["A"]]).unwrap()).unwrap();
892        assert_eq!(one.gate("q", Describe::default()), None);
893    }
894
895    #[test]
896    fn from_paths_groups_by_prefix_in_first_order() {
897        let root = Branch::from_paths([vec!["B", "y"], vec!["A"], vec!["B", "x"], vec!["C", "z", "1"]]).unwrap();
898        let want = Branch::node(
899            "",
900            [
901                Branch::node("B", [Branch::leaf("y"), Branch::leaf("x")]),
902                Branch::leaf("A"),
903                Branch::node("C", [Branch::node("z", [Branch::leaf("1")])]),
904            ],
905        );
906        assert_eq!(root, want);
907        assert!(Tree::new(root).is_ok());
908    }
909
910    #[test]
911    fn from_paths_refuses_what_is_not_a_tree_of_leaves() {
912        let none: [Vec<&str>; 1] = [vec![]];
913        assert_eq!(Branch::from_paths(none), Err(PathError::Empty));
914        assert_eq!(
915            Branch::from_paths([vec!["A", "x"], vec!["A", "x"]]),
916            Err(PathError::Repeated { path: vec!["A".into(), "x".into()] })
917        );
918        // Either order: the group first, or the leaf first.
919        assert_eq!(
920            Branch::from_paths([vec!["A", "x"], vec!["A"]]),
921            Err(PathError::LeafAndGroup { path: vec!["A".into()] })
922        );
923        assert_eq!(
924            Branch::from_paths([vec!["A"], vec!["A", "x"]]),
925            Err(PathError::LeafAndGroup { path: vec!["A".into()] })
926        );
927    }
928
929    /// With [`plain`] params, the Choices asked.
930    fn run(tree: &Tree, params: Params, answer: impl Fn(&[&str]) -> Vec<f64>) -> (Outcome, usize) {
931        let (o, rounds) = run_rounds(tree, params, fixed(answer));
932        (o, rounds.iter().map(Vec::len).sum())
933    }
934
935    fn chain(label: &str, depth: usize, leaf: &str) -> Branch {
936        (0..depth).rev().fold(Branch::leaf(leaf), |below, i| {
937            Branch::node(format!("{label}{i}"), [below, Branch::leaf(format!("{label}{i}-other"))])
938        })
939    }
940
941    fn leaf_of(o: &Outcome) -> NodeId {
942        match o {
943            Outcome::Leaf { leaf, .. } => *leaf,
944            other => panic!("{other:?}"),
945        }
946    }
947
948    fn sent(c: &Choice) -> String {
949        jev_protocol::question_bytes(&jev_protocol::Question::Choice(c.clone()))
950            .map(|b| String::from_utf8(b).unwrap())
951            .unwrap()
952    }
953
954    /// Options describe their subtree; instructions name the path; a leaf
955    /// has no description; a one-child node has no Choice.
956    #[test]
957    fn node_choices() {
958        let tree = Tree::new(Branch::node(
959            "root",
960            [
961                Branch::node("Animals", [
962                    Branch::node("Birds", [Branch::leaf("Owl"), Branch::leaf("Wren")]),
963                    Branch::leaf("Fish"),
964                ]),
965                Branch::leaf("Plants"),
966                Branch::node("Rocks", [Branch::leaf("Granite")]),
967            ],
968        ))
969        .unwrap();
970        let root = tree.root();
971        let c = tree.choice(root, Order::Caller, "What is it?", Describe::default()).unwrap();
972        assert_eq!(
973            sent(&c),
974            r#"{"type":"choice","instructions":"What is it?","criteria":{"Animals":"Contains: Birds, Fish. For example: Owl, Wren","Plants":null,"Rocks":"Contains: Granite"}}"#
975        );
976        let twin = tree.choice(root, Order::Reversed, "What is it?", Describe::default()).unwrap();
977        assert_eq!(
978            sent(&twin),
979            r#"{"type":"choice","instructions":"What is it?","criteria":{"Rocks":"Contains: Granite","Plants":null,"Animals":"Contains: Birds, Fish. For example: Owl, Wren"}}"#
980        );
981        let animals = tree.children(root)[0];
982        let birds = tree.children(animals)[0];
983        let c = tree.choice(birds, Order::Caller, "What is it?", Describe::default()).unwrap();
984        assert_eq!(
985            sent(&c),
986            r#"{"type":"choice","instructions":"What is it?\n\nWithin: Animals > Birds","criteria":{"Owl":null,"Wren":null}}"#
987        );
988        let rocks = tree.children(root)[2];
989        assert!(tree.choice(rocks, Order::Caller, "q", Describe::default()).is_none());
990        let short = tree.choice(root, Order::Caller, "q", Describe { children: 1, leaves: 1 }).unwrap();
991        assert!(sent(&short).contains(r#""Animals":"Contains: Birds, and 1 more. For example: Owl""#), "{}", sent(&short));
992    }
993
994    /// Contract: five 0.9 edges are a 0.59 path and lose to one 0.8 edge,
995    /// which the geometric mean (0.90 against 0.8) would get backwards.
996    #[test]
997    fn five_confident_edges_lose_to_one_likelier_edge() {
998        // root → deep (0.9) → d1 (0.9) → d2 (0.9) → d3 (0.9) → deep-leaf (0.9);
999        // root → shallow (0.8, a leaf).
1000        let deep = Branch::node(
1001            "deep",
1002            [
1003                Branch::node(
1004                    "d1",
1005                    [
1006                        Branch::node(
1007                            "d2",
1008                            [
1009                                Branch::node("d3", [Branch::leaf("deep-leaf"), Branch::leaf("d3-x")]),
1010                                Branch::leaf("d2-x"),
1011                            ],
1012                        ),
1013                        Branch::leaf("d1-x"),
1014                    ],
1015                ),
1016                Branch::leaf("deep-x"),
1017            ],
1018        );
1019        let tree = Tree::new(Branch::node("", [deep, Branch::leaf("shallow")])).unwrap();
1020        let (o, _) = run(&tree, plain(), |path| match path {
1021            [] => vec![0.9, 0.8],
1022            _ => vec![0.9, 0.1],
1023        });
1024        let Outcome::Leaf { leaf, probability, separation } = o else { panic!("{o:?}") };
1025        assert_eq!(tree.path(leaf), ["shallow"]);
1026        assert!((probability - 0.8).abs() < 1e-12, "{probability}");
1027        // The runner-up is `deep-leaf` itself, at 0.9⁵ ≈ 0.59.
1028        assert!((0.9f64.powi(5) - 0.59).abs() < 0.001);
1029        assert!((separation.unwrap() - 0.8 / 0.9f64.powi(5)).abs() < 1e-9, "{separation:?}");
1030    }
1031
1032    /// Exactness: the result equals exhaustive enumeration of every
1033    /// leaf, for every K, on a tree where greedy descent is wrong.
1034    #[test]
1035    fn exact_against_exhaustive() {
1036        let tree = Tree::new(Branch::node(
1037            "",
1038            [
1039                Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2"), Branch::leaf("a3")]),
1040                Branch::node("b", [Branch::leaf("b1"), Branch::node("b2", [Branch::leaf("b2x"), Branch::leaf("b2y")])]),
1041                Branch::node("c", [Branch::leaf("c1"), Branch::leaf("c2")]),
1042            ],
1043        ))
1044        .unwrap();
1045        // Greedy takes `a` (0.4) but its leaves split three ways; `b1` is
1046        // 0.35 × 0.95.
1047        let answer = |path: &[&str]| -> Vec<f64> {
1048            match path {
1049                [] => vec![0.4, 0.35, 0.25],
1050                ["a"] => vec![0.34, 0.33, 0.33],
1051                ["b"] => vec![0.95, 0.05],
1052                ["b", "b2"] => vec![0.5, 0.5],
1053                ["c"] => vec![0.6, 0.4],
1054                other => panic!("{other:?} asked"),
1055            }
1056        };
1057        let mut leaves: Vec<(f64, Vec<&str>)> = Vec::new();
1058        let mut stack = vec![(tree.root(), 0.0)];
1059        while let Some((n, s)) = stack.pop() {
1060            if tree.is_leaf(n) {
1061                leaves.push((s, tree.path(n)));
1062                continue;
1063            }
1064            for (&c, p) in tree.children(n).iter().zip(answer(&tree.path(n))) {
1065                stack.push((c, s + p.max(PROBABILITY_FLOOR).ln()));
1066            }
1067        }
1068        leaves.sort_by(|a, b| b.0.total_cmp(&a.0));
1069        // Lookahead and twins change the rounds, never the result.
1070        let extras = [
1071            (None, None),
1072            (Some(Lookahead::default()), None),
1073            (Some(Lookahead { m: 1, budget: 250 }), TwinBelow::new(1.5)),
1074            (Some(Lookahead::default()), TwinBelow::new(1.5)),
1075        ];
1076        for k in 1..=5 {
1077            for (lookahead, twin) in extras {
1078                let params = Params { k: NonZeroUsize::new(k).unwrap(), tau: None, lookahead, twin };
1079                let (o, _) = run_rounds(&tree, params, fixed(answer));
1080                let Outcome::Leaf { leaf, separation, .. } = o else { panic!() };
1081                assert_eq!(tree.path(leaf), leaves[0].1, "K={k} {params:?}");
1082                assert!((separation.unwrap() - (leaves[0].0 - leaves[1].0).exp()).abs() < 1e-12, "{params:?}");
1083            }
1084        }
1085    }
1086
1087    /// A confident path never asks the unlikely subtrees below it.
1088    #[test]
1089    fn prunes_what_cannot_win() {
1090        let wide = |l: &str| Branch::node(l, (0..50).map(|i| chain(&format!("{l}{i}-"), 3, &format!("{l}{i}!"))));
1091        let tree = Tree::new(Branch::node("", [wide("x"), wide("y"), wide("z")])).unwrap();
1092        let (o, asked) = run(&tree, plain(), |path| match path.len() {
1093            0 => vec![0.98, 0.01, 0.01],
1094            1 => std::iter::once(0.9).chain(std::iter::repeat_n(0.1 / 49.0, 49)).collect(),
1095            _ => vec![0.9, 0.1],
1096        });
1097        assert_eq!(tree.path(leaf_of(&o)).last(), Some(&"x0!"));
1098        // The best path, and its runner-ups' first Choices; nowhere near
1099        // the 3 + 150 + 450 internal nodes.
1100        assert!(asked <= 12, "{asked}");
1101    }
1102
1103    #[test]
1104    fn one_child_is_not_asked() {
1105        let tree = Tree::new(Branch::node(
1106            "",
1107            [Branch::node("only", [Branch::node("still", [Branch::leaf("p"), Branch::leaf("q")])])],
1108        ))
1109        .unwrap();
1110        let (o, asked) = run(&tree, plain(), |path| {
1111            assert_eq!(path, ["only", "still"]);
1112            vec![0.3, 0.7]
1113        });
1114        let Outcome::Leaf { leaf, probability, .. } = o else { panic!() };
1115        assert_eq!(tree.path(leaf), ["only", "still", "q"]);
1116        assert!((probability - 0.7).abs() < 1e-12);
1117        assert_eq!(asked, 1);
1118    }
1119
1120    #[test]
1121    fn a_single_leaf_has_no_separation() {
1122        let tree = Tree::new(Branch::node("", [Branch::leaf("sole")])).unwrap();
1123        let (o, asked) = run(&tree, plain(), |_| unreachable!());
1124        assert_eq!(o, Outcome::Leaf { leaf: NodeId(1), probability: 1.0, separation: None });
1125        assert_eq!(asked, 0);
1126    }
1127
1128    #[test]
1129    fn a_zero_is_floored_not_minus_infinity() {
1130        let tree = Tree::new(Branch::node(
1131            "",
1132            [Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2")]), Branch::leaf("b")],
1133        ))
1134        .unwrap();
1135        let (o, _) = run(&tree, plain(), |path| match path {
1136            [] => vec![1.0, 0.0],
1137            _ => vec![0.0, 0.0],
1138        });
1139        let Outcome::Leaf { leaf, separation, .. } = o else { panic!() };
1140        // a1 and a2 tie at 1e-9; `b` (also 1e-9) is not ranked below −∞.
1141        assert_eq!(tree.path(leaf), ["a", "a1"]);
1142        assert_eq!(separation, Some(1.0));
1143    }
1144
1145    /// Lookahead scores two levels per round trip on a chain, and asks
1146    /// nothing it was not going to use there.
1147    #[test]
1148    fn lookahead_halves_the_round_trips() {
1149        let tree = Tree::new(chain("n", 8, "end")).unwrap();
1150        let answer = |_: &[&str]| vec![0.9, 0.1];
1151        let k1 = Params { k: NonZeroUsize::MIN, ..plain() };
1152        let (o, rounds) = run_rounds(&tree, k1, fixed(answer));
1153        assert_eq!(tree.path(leaf_of(&o)).last(), Some(&"end"));
1154        assert_eq!(rounds.len(), 8);
1155        let look = Params { lookahead: Some(Lookahead { m: 1, budget: 1_000 }), ..k1 };
1156        let (o2, rounds) = run_rounds(&tree, look, fixed(answer));
1157        assert_eq!(o2, o);
1158        assert_eq!(rounds.len(), 4, "{rounds:?}");
1159        assert!(rounds.iter().all(|r| matches!(r[..], [Ask::Node(_), Ask::Lookahead(_)])), "{rounds:?}");
1160    }
1161
1162    /// Lookahead fills the round up to the budget, likeliest child
1163    /// (most leaves) first, and never past it.
1164    #[test]
1165    fn lookahead_stays_within_the_budget() {
1166        let tree = Tree::new(Branch::node(
1167            "",
1168            [
1169                Branch::node("small", [Branch::leaf("s1"), Branch::leaf("s2")]),
1170                Branch::node("big", [Branch::leaf("b1"), Branch::leaf("b2"), Branch::leaf("b3")]),
1171                Branch::node("mid", [Branch::leaf("m1"), Branch::leaf("m2")]),
1172            ],
1173        ))
1174        .unwrap();
1175        let [small, big, mid] = tree.children(tree.root()) else { panic!() };
1176        let first = |budget| {
1177            let mut s = Search::new(&tree, Params { lookahead: Some(Lookahead { m: 3, budget }), ..plain() });
1178            let Step::Ask(round) = s.step(|_| 100) else { panic!() };
1179            round
1180        };
1181        let root = Ask::Node(tree.root());
1182        assert_eq!(first(99), [root]);
1183        assert_eq!(first(250), [root, Ask::Lookahead(*big)]);
1184        // Equal leaf counts keep the caller's order.
1185        assert_eq!(first(400), [root, Ask::Lookahead(*big), Ask::Lookahead(*small), Ask::Lookahead(*mid)]);
1186    }
1187
1188    /// A close call is asked again in reverse, and the two answers'
1189    /// log-probabilities are averaged, which undoes a bias toward the
1190    /// first-listed option.
1191    #[test]
1192    fn a_close_call_gets_an_order_twin() {
1193        let tree = Tree::new(Branch::node("", [Branch::leaf("a"), Branch::leaf("b")])).unwrap();
1194        // Listed first gains 0.1: truly a 0.45, b 0.55.
1195        let biased = |_: &[&str], order: Order| match order {
1196            Order::Caller => vec![0.55, 0.45],
1197            Order::Reversed => vec![0.65, 0.35],
1198        };
1199        let (o, rounds) = run_rounds(&tree, plain(), biased);
1200        assert_eq!(tree.path(leaf_of(&o)), ["a"]);
1201        assert_eq!(rounds.len(), 1);
1202
1203        let twin = Params { twin: TwinBelow::new(1.5), ..plain() };
1204        let (o, rounds) = run_rounds(&tree, twin, biased);
1205        assert_eq!(rounds, [vec![Ask::Node(tree.root())], vec![Ask::Twin(tree.root())]]);
1206        let Outcome::Leaf { leaf, probability, separation } = o else { panic!() };
1207        assert_eq!(tree.path(leaf), ["b"]);
1208        let (a, b) = ((0.55f64 * 0.35).sqrt(), (0.45f64 * 0.65).sqrt());
1209        assert!((probability - b).abs() < 1e-12, "{probability}");
1210        assert!((separation.unwrap() - b / a).abs() < 1e-12);
1211
1212        // Not close: no twin.
1213        let (_, rounds) = run_rounds(&tree, twin, fixed(|_| vec![0.2, 0.8]));
1214        assert_eq!(rounds.len(), 1);
1215        assert!(TwinBelow::new(1.0).is_none());
1216        assert!(TwinBelow::new(f64::INFINITY).is_none());
1217    }
1218
1219    /// A close call found by lookahead waits for its twin too.
1220    #[test]
1221    fn a_close_lookahead_answer_is_twinned() {
1222        let tree = Tree::new(Branch::node(
1223            "",
1224            [Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2")]), Branch::leaf("b")],
1225        ))
1226        .unwrap();
1227        let a = tree.children(tree.root())[0];
1228        let (o, rounds) = run_rounds(&tree, Params::default(), fixed(|path| match path {
1229            [] => vec![0.9, 0.1],
1230            _ => vec![0.5, 0.5],
1231        }));
1232        assert_eq!(rounds, [vec![Ask::Node(tree.root()), Ask::Lookahead(a)], vec![Ask::Twin(a)]]);
1233        assert_eq!(tree.path(leaf_of(&o)), ["a", "a1"]);
1234    }
1235
1236    fn tau(t: f64) -> Params {
1237        Params { tau: Tau::new(t), ..plain() }
1238    }
1239
1240    #[test]
1241    fn tau_stops_where_the_path_falls_below_it() {
1242        let tree = Tree::new(chain("n", 4, "end")).unwrap();
1243        // Root is n0; each level keeps 0.8 on the chain: 0.8, 0.64, 0.51, 0.41.
1244        let answer = |_: &[&str]| vec![0.8, 0.2];
1245        let (o, asked) = run(&tree, tau(0.5), answer);
1246        let Outcome::Stop { node, depth, leaf, probability } = o else { panic!("{o:?}") };
1247        assert_eq!(tree.path(node), ["n1", "n2", "n3"]);
1248        assert_eq!(depth, 3);
1249        assert!(!leaf);
1250        assert!((probability - 0.512).abs() < 1e-12);
1251        // n3 is asked, since only its answer shows its children fall
1252        // below 0.5; nothing below τ is asked.
1253        assert_eq!(asked, 4);
1254
1255        let (o, _) = run(&tree, tau(0.4), answer);
1256        let Outcome::Stop { node, leaf, .. } = o else { panic!() };
1257        assert_eq!(tree.path(node), ["n1", "n2", "n3", "end"]);
1258        assert!(leaf);
1259    }
1260
1261    #[test]
1262    fn tau_prefers_depth_then_probability() {
1263        let tree = Tree::new(Branch::node(
1264            "",
1265            [
1266                Branch::node("a", [Branch::leaf("a1"), Branch::leaf("a2")]),
1267                Branch::node("b", [Branch::node("b1", [Branch::leaf("x"), Branch::leaf("y")]), Branch::leaf("b2")]),
1268            ],
1269        ))
1270        .unwrap();
1271        let (o, _) = run(&tree, tau(0.2), |path| match path {
1272            [] => vec![0.6, 0.4],
1273            ["a"] => vec![0.5, 0.5],
1274            ["b"] => vec![0.6, 0.4],
1275            ["b", "b1"] => vec![0.5, 0.5],
1276            other => panic!("{other:?}"),
1277        });
1278        // a1/a2 are 0.3 at depth 2; b1 is 0.24 at depth 2, and x/y 0.12 < τ.
1279        // So depth 2, likeliest: a1.
1280        let Outcome::Stop { node, depth, .. } = o else { panic!() };
1281        assert_eq!((tree.path(node), depth), (vec!["a", "a1"], 2));
1282    }
1283
1284    #[test]
1285    fn tau_is_in_zero_to_one() {
1286        assert!(Tau::new(0.0).is_none());
1287        assert!(Tau::new(1.5).is_none());
1288        assert!(Tau::new(f64::NAN).is_none());
1289        assert!(Tau::new(1.0).is_some());
1290    }
1291
1292    #[test]
1293    fn trees_that_cannot_be_asked_are_refused() {
1294        assert_eq!(Tree::new(Branch::leaf("")).unwrap_err(), TreeError::NoChoice);
1295        let wide = Branch::node("", (0..256).map(|i| Branch::leaf(i.to_string())));
1296        assert!(matches!(Tree::new(wide), Err(TreeError::Fanout { children: 256, .. })));
1297        assert!(Tree::new(Branch::node("", (0..255).map(|i| Branch::leaf(i.to_string())))).is_ok());
1298        let dup = Branch::node("", [Branch::node("a", [Branch::leaf("x"), Branch::leaf("x")])]);
1299        assert_eq!(
1300            Tree::new(dup).unwrap_err(),
1301            TreeError::DuplicateLabel { path: vec!["a".into()], label: "x".into() }
1302        );
1303        let empty = Branch::node("", [Branch::leaf("")]);
1304        assert_eq!(Tree::new(empty).unwrap_err(), TreeError::EmptyLabel { path: vec![] });
1305    }
1306
1307    #[test]
1308    fn answers_are_checked() {
1309        let tree = Tree::new(Branch::node("", [Branch::leaf("a"), Branch::leaf("b")])).unwrap();
1310        let mut s = Search::new(&tree, Params::default());
1311        let root = Ask::Node(tree.root());
1312        let Step::Ask(round) = s.step(|_| 100) else { panic!() };
1313        assert_eq!(round, [root]);
1314        assert_eq!(s.step(|_| 100), Step::Ask(round.clone()));
1315        assert_eq!(s.answer(Ask::Node(NodeId(1)), &[1.0]), Err(AnswerError::NotAsked(Ask::Node(NodeId(1)))));
1316        assert_eq!(s.answer(Ask::Twin(tree.root()), &[0.5, 0.5]), Err(AnswerError::NotAsked(Ask::Twin(tree.root()))));
1317        assert!(matches!(s.answer(root, &[1.0]), Err(AnswerError::Count { expected: 2, got: 1, .. })));
1318        assert!(matches!(s.answer(root, &[0.5, f64::NAN]), Err(AnswerError::NotAProbability { .. })));
1319        assert!(matches!(s.answer(root, &[0.5, 1.2]), Err(AnswerError::NotAProbability { .. })));
1320        s.answer(root, &[0.25, 0.75]).unwrap();
1321        let Step::Done(o) = s.step(|_| 100) else { panic!() };
1322        assert_eq!(tree.path(leaf_of(&o)), ["b"]);
1323    }
1324}