1//! Choice over more than 255 labels by chunk and shortlist (contract
2//! *SQL surface*, "More than 255 options", chunk-and-shortlist mode), for
3//! label sets with no meaningful hierarchy. This is the vendor's shape in
4//! the skill-suggestion and line-by-line-search cookbooks.
5//!
6//! Two rounds at any N:
7//!
8//! 1. ⌈N/254⌉ chunk Choices, all in one request, each over a contiguous
9//!    run of the labels in the caller's order, labels only. The top
10//!    [`Params::keep`] of each chunk go through.
11//! 2. One Choice over those finalists, in the caller's order, with the
12//!    caller's full descriptions.
13//!
14//! The chunks are as even as they can be (sizes differ by at most one),
15//! so no chunk is a one-option non-question. Nothing is ever reordered:
16//! order is part of the question (measured order bias), so a finalist
17//! keeps its place in the caller's list, never its rank in round 1.
18//!
19//! What is returned is exact about what it is: the finalist Choice's top
20//! label and its probability *among the finalists*. That is not a
21//! probability over all N labels, which no request here measures.
22//!
23//! Sans-I/O: [`Search::step`] says which Choices to ask,
24//! [`Shortlist::choice`] builds each one, and [`Search::answer`] takes its
25//! probabilities in the order sent. Sending them is the caller's.
26
27use std::collections::HashSet;
28use std::num::NonZeroUsize;
29use std::ops::Range;
30
31use jev_protocol::{Choice, Json, MAX_CHOICE_OPTIONS, Noul};
32
33/// Labels per chunk at most: the contract's ⌈N/254⌉.
34pub const CHUNK: usize = 254;
35
36/// Finalists kept from each chunk. Ours and unmeasured; the label-tree
37/// release gate (*Testing*) measures it against the tree.
38pub const DEFAULT_KEEP: NonZeroUsize = NonZeroUsize::new(3).unwrap();
39
40#[derive(Clone, Copy, Debug, PartialEq, Eq)]
41pub struct Params {
42    /// The top few of each chunk that reach the finalist Choice.
43    pub keep: NonZeroUsize,
44}
45
46impl Default for Params {
47    fn default() -> Self {
48        Params { keep: DEFAULT_KEEP }
49    }
50}
51
52/// Options that cannot be asked this way. Every one is the caller's
53/// input, so every one is `invalid_parameter_value`.
54#[derive(Clone, Debug, PartialEq, Eq)]
55pub enum ShortlistError {
56    NoOptions,
57    EmptyLabel { index: usize },
58    DuplicateLabel { label: String },
59    /// The finalists would not make one Choice of 2 to 255 options.
60    Finalists { count: usize },
61}
62
63impl ShortlistError {
64    /// Raised before anything is sent.
65    pub fn sqlstate(&self) -> &'static str {
66        match self {
67            ShortlistError::NoOptions
68            | ShortlistError::EmptyLabel { .. }
69            | ShortlistError::DuplicateLabel { .. }
70            | ShortlistError::Finalists { .. } => "22023",
71        }
72    }
73}
74
75impl std::fmt::Display for ShortlistError {
76    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
77        match self {
78            ShortlistError::NoOptions => write!(f, "a Choice needs at least one option"),
79            ShortlistError::EmptyLabel { index } => write!(f, "option {} has an empty label", index + 1),
80            ShortlistError::DuplicateLabel { label } => write!(f, "the option {label:?} appears twice"),
81            ShortlistError::Finalists { count } => write!(
82                f,
83                "the shortlist would have {count} finalists; a Choice takes 2 to {MAX_CHOICE_OPTIONS}"
84            ),
85        }
86    }
87}
88
89impl std::error::Error for ShortlistError {}
90
91/// A validated option list, cut into chunks.
92#[derive(Clone, Debug)]
93pub struct Shortlist {
94    /// `(label, description)` in the caller's order.
95    options: Vec<(String, Option<Json>)>,
96    chunks: Vec<Range<usize>>,
97    params: Params,
98}
99
100impl Shortlist {
101    pub fn new(
102        options: impl IntoIterator<Item = (String, Option<Json>)>,
103        params: Params,
104    ) -> Result<Self, ShortlistError> {
105        let options: Vec<_> = options.into_iter().collect();
106        if options.is_empty() {
107            return Err(ShortlistError::NoOptions);
108        }
109        let mut seen = HashSet::new();
110        for (index, (label, _)) in options.iter().enumerate() {
111            if label.is_empty() {
112                return Err(ShortlistError::EmptyLabel { index });
113            }
114            if !seen.insert(label.as_str()) {
115                return Err(ShortlistError::DuplicateLabel { label: label.clone() });
116            }
117        }
118        let n = options.len();
119        let count = n.div_ceil(CHUNK);
120        let (base, extra) = (n / count, n % count);
121        let mut chunks = Vec::with_capacity(count);
122        let mut start = 0;
123        for i in 0..count {
124            let len = base + usize::from(i < extra);
125            chunks.push(start..start + len);
126            start += len;
127        }
128        let s = Shortlist { options, chunks, params };
129        if n > 1 {
130            let finalists: usize = s.chunks.iter().map(|c| s.kept(c)).sum();
131            if !(2..=MAX_CHOICE_OPTIONS).contains(&finalists) {
132                return Err(ShortlistError::Finalists { count: finalists });
133            }
134        }
135        Ok(s)
136    }
137
138    fn kept(&self, chunk: &Range<usize>) -> usize {
139        chunk.len().min(self.params.keep.get())
140    }
141
142    pub fn len(&self) -> usize {
143        self.options.len()
144    }
145
146    pub fn is_empty(&self) -> bool {
147        self.options.is_empty()
148    }
149
150    pub fn label(&self, index: usize) -> &str {
151        &self.options[index].0
152    }
153
154    /// The chunks, as index ranges into the caller's list.
155    pub fn chunks(&self) -> &[Range<usize>] {
156        &self.chunks
157    }
158
159    /// `ask`'s Choice for `search`'s current round. Every piece is fixed,
160    /// so equal asks give equal bytes and one cache key: the instructions
161    /// are `question` alone, a chunk sends labels only, and the finalists
162    /// send the caller's descriptions.
163    pub fn choice(&self, search: &Search<'_>, ask: Ask, question: &str) -> Choice {
164        let options: Vec<_> = match ask {
165            Ask::Chunk(i) => self.chunks[i].clone().map(|o| (self.options[o].0.clone(), None)).collect(),
166            Ask::Finalists => search.finalists.iter().map(|&o| self.options[o].clone()).collect(),
167        };
168        Choice::new(Json::text(question), options).expect("a Shortlist's asks make valid Choices")
169    }
170
171    /// The "does any label fit" Noul ([`crate::gate`]) over every label, in
172    /// the caller's order, without descriptions (as a chunk sends them).
173    /// `None` for a single label, which the search answers without a
174    /// request.
175    pub fn gate(&self, question: &str) -> Option<Noul> {
176        if self.options.len() < 2 {
177            return None;
178        }
179        let labels = self.options.iter().map(|(l, _)| l.as_str()).collect::<Vec<_>>().join(", ");
180        Some(crate::gate::gate(question, &labels))
181    }
182}
183
184/// One Choice a round asks.
185#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
186pub enum Ask {
187    /// Round 1: the chunk at this index of [`Shortlist::chunks`].
188    Chunk(usize),
189    /// Round 2.
190    Finalists,
191}
192
193#[derive(Clone, Debug, PartialEq)]
194pub enum Step {
195    /// Ask these Choices, in one request, and [`Search::answer`] each.
196    Ask(Vec<Ask>),
197    Done(Outcome),
198}
199
200#[derive(Clone, Debug, PartialEq)]
201pub struct Outcome {
202    /// Index into the caller's list.
203    pub choice: usize,
204    /// Among the finalists, not over all N. 1 for a single option, which
205    /// is never asked.
206    pub probability: f64,
207    /// The finalists as `(index, probability)` in the caller's order;
208    /// empty for a single option.
209    pub finalists: Vec<(usize, f64)>,
210}
211
212#[derive(Clone, Debug, PartialEq)]
213pub enum AnswerError {
214    /// Not asked in the current round, or already answered.
215    NotAsked(Ask),
216    /// One probability per option, in the order sent.
217    Count { ask: Ask, expected: usize, got: usize },
218    NotAProbability { ask: Ask, value: f64 },
219}
220
221impl std::fmt::Display for AnswerError {
222    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
223        match self {
224            AnswerError::NotAsked(a) => write!(f, "shortlist Choice {a:?} was not asked"),
225            AnswerError::Count { ask, expected, got } => {
226                write!(f, "shortlist Choice {ask:?} has {expected} options, answered with {got}")
227            }
228            AnswerError::NotAProbability { ask, value } => {
229                write!(f, "shortlist Choice {ask:?} was answered with {value}, not a probability")
230            }
231        }
232    }
233}
234
235impl std::error::Error for AnswerError {}
236
237#[derive(Debug)]
238enum Phase {
239    /// Round 1; the chunks still unanswered.
240    Chunks(Vec<usize>),
241    /// Round 2.
242    Finalists,
243    Done(Outcome),
244}
245
246/// One chunk-and-shortlist search.
247#[derive(Debug)]
248pub struct Search<'s> {
249    list: &'s Shortlist,
250    phase: Phase,
251    /// Kept so far, sorted into the caller's order before round 2.
252    finalists: Vec<usize>,
253}
254
255impl<'s> Search<'s> {
256    pub fn new(list: &'s Shortlist) -> Self {
257        let phase = if list.len() == 1 {
258            Phase::Done(Outcome { choice: 0, probability: 1.0, finalists: Vec::new() })
259        } else {
260            Phase::Chunks((0..list.chunks.len()).collect())
261        };
262        Search { list, phase, finalists: Vec::new() }
263    }
264
265    /// The round to ask, or the result. Asking again before every asked
266    /// Choice is answered returns the same round.
267    pub fn step(&self) -> Step {
268        match &self.phase {
269            Phase::Chunks(left) => Step::Ask(left.iter().map(|&i| Ask::Chunk(i)).collect()),
270            Phase::Finalists => Step::Ask(vec![Ask::Finalists]),
271            Phase::Done(o) => Step::Done(o.clone()),
272        }
273    }
274
275    /// A Choice's answer: one probability per option, in the order sent.
276    pub fn answer(&mut self, ask: Ask, probabilities: &[f64]) -> Result<(), AnswerError> {
277        let expected = match (&self.phase, ask) {
278            (Phase::Chunks(left), Ask::Chunk(i)) if left.contains(&i) => self.list.chunks[i].len(),
279            (Phase::Finalists, Ask::Finalists) => self.finalists.len(),
280            _ => return Err(AnswerError::NotAsked(ask)),
281        };
282        if probabilities.len() != expected {
283            return Err(AnswerError::Count { ask, expected, got: probabilities.len() });
284        }
285        if let Some(&value) = probabilities.iter().find(|p| !(0.0..=1.0).contains(*p)) {
286            return Err(AnswerError::NotAProbability { ask, value });
287        }
288        match ask {
289            Ask::Chunk(i) => {
290                let chunk = self.list.chunks[i].clone();
291                let mut ranked: Vec<usize> = (0..chunk.len()).collect();
292                // Stable, so ties keep the caller's order.
293                ranked.sort_by(|&a, &b| probabilities[b].total_cmp(&probabilities[a]));
294                self.finalists.extend(ranked.into_iter().take(self.list.kept(&chunk)).map(|o| chunk.start + o));
295                let Phase::Chunks(left) = &mut self.phase else { unreachable!("matched above") };
296                left.retain(|&c| c != i);
297                if left.is_empty() {
298                    self.finalists.sort_unstable();
299                    self.phase = Phase::Finalists;
300                }
301            }
302            Ask::Finalists => {
303                let mut top = 0;
304                for (i, p) in probabilities.iter().enumerate() {
305                    if *p > probabilities[top] {
306                        top = i;
307                    }
308                }
309                self.phase = Phase::Done(Outcome {
310                    choice: self.finalists[top],
311                    probability: probabilities[top],
312                    finalists: self.finalists.iter().copied().zip(probabilities.iter().copied()).collect(),
313                });
314            }
315        }
316        Ok(())
317    }
318}
319
320#[cfg(test)]
321mod tests {
322    use super::*;
323
324    fn labels(n: usize) -> Vec<(String, Option<Json>)> {
325        (0..n).map(|i| (format!("label {i}"), Some(Json::text(&format!("about {i}"))))).collect()
326    }
327
328    /// The labels a Choice sends, in order.
329    fn sent(c: &Choice) -> Vec<String> {
330        c.labels().map(str::to_owned).collect()
331    }
332
333    /// Runs a search, answering each Choice with `score(label)` per option
334    /// normalized, and returns the outcome and the rounds' Choices.
335    fn run(list: &Shortlist, score: impl Fn(usize) -> f64) -> (Outcome, Vec<Vec<Choice>>) {
336        let mut s = Search::new(list);
337        let mut rounds = Vec::new();
338        loop {
339            match s.step() {
340                Step::Done(o) => return (o, rounds),
341                Step::Ask(round) => {
342                    let choices: Vec<Choice> = round.iter().map(|&a| list.choice(&s, a, "Which fits?")).collect();
343                    for (&a, c) in round.iter().zip(&choices) {
344                        let raw: Vec<f64> = c
345                            .labels()
346                            .map(|l| score(l.strip_prefix("label ").unwrap().parse().unwrap()))
347                            .collect();
348                        let total: f64 = raw.iter().sum();
349                        s.answer(a, &raw.iter().map(|r| r / total).collect::<Vec<_>>()).unwrap();
350                    }
351                    rounds.push(choices);
352                }
353            }
354        }
355    }
356
357    #[test]
358    fn exactly_two_rounds_up_to_a_thousand() {
359        for n in 2..=1000 {
360            let list = Shortlist::new(labels(n), Params::default()).unwrap();
361            let want = n * 37 / 100;
362            let (o, rounds) = run(&list, |i| if i == want { 10.0 } else { 1.0 + (i % 5) as f64 });
363            assert_eq!(rounds.len(), 2, "n = {n}");
364            assert_eq!(rounds[0].len(), n.div_ceil(CHUNK), "n = {n}");
365            assert_eq!(rounds[1].len(), 1, "n = {n}");
366            assert_eq!(o.choice, want, "n = {n}");
367            // Every chunk is a real question, and the chunks cover the list
368            // once, in order.
369            let chunks = list.chunks();
370            assert!(chunks.iter().all(|c| (2..=CHUNK).contains(&c.len())), "n = {n}");
371            assert_eq!(chunks.first().unwrap().start, 0);
372            assert_eq!(chunks.last().unwrap().end, n);
373            assert!(chunks.windows(2).all(|w| w[0].end == w[1].start));
374        }
375    }
376
377    #[test]
378    fn caller_order_is_kept_end_to_end() {
379        let list = Shortlist::new(labels(600), Params::default()).unwrap();
380        // Likeliest last, so rank order is the reverse of caller order.
381        let (o, rounds) = run(&list, |i| 1.0 + i as f64);
382        let round1: Vec<String> = rounds[0].iter().flat_map(sent).collect();
383        let caller: Vec<String> = (0..600).map(|i| format!("label {i}")).collect();
384        assert_eq!(round1, caller);
385        // Each chunk's top 3 are its last 3, sent in the caller's order.
386        let mut want = Vec::new();
387        for c in list.chunks() {
388            want.extend((c.end - 3..c.end).map(|i| format!("label {i}")));
389        }
390        assert_eq!(sent(&rounds[1][0]), want);
391        assert_eq!(o.choice, 599);
392        assert!(o.finalists.windows(2).all(|w| w[0].0 < w[1].0));
393    }
394
395    #[test]
396    fn chunks_send_labels_and_finalists_their_descriptions() {
397        let list = Shortlist::new(labels(300), Params::default()).unwrap();
398        let s = Search::new(&list);
399        let chunk = list.choice(&s, Ask::Chunk(0), "Which fits?");
400        let bytes = String::from_utf8(jev_protocol::question_bytes(&jev_protocol::Question::Choice(chunk)).unwrap()).unwrap();
401        assert!(!bytes.contains("about"), "{bytes}");
402        let (_, rounds) = run(&list, |i| 1.0 + i as f64);
403        let bytes = String::from_utf8(jev_protocol::question_bytes(&jev_protocol::Question::Choice(rounds[1][0].clone())).unwrap()).unwrap();
404        assert!(bytes.contains("about 299"), "{bytes}");
405    }
406
407    #[test]
408    fn ties_keep_the_caller_order() {
409        let list = Shortlist::new(labels(10), Params::default()).unwrap();
410        let (o, rounds) = run(&list, |_| 1.0);
411        assert_eq!(sent(&rounds[1][0]), ["label 0", "label 1", "label 2"]);
412        assert_eq!(o.choice, 0);
413    }
414
415    #[test]
416    fn keep_is_a_parameter() {
417        let keep = NonZeroUsize::new(7).unwrap();
418        let list = Shortlist::new(labels(1000), Params { keep }).unwrap();
419        let (_, rounds) = run(&list, |i| 1.0 + i as f64);
420        assert_eq!(rounds[1][0].labels().len(), 4 * 7);
421    }
422
423    #[test]
424    fn a_single_option_is_not_asked() {
425        let list = Shortlist::new(labels(1), Params::default()).unwrap();
426        let (o, rounds) = run(&list, |_| 1.0);
427        assert!(rounds.is_empty());
428        assert_eq!(o, Outcome { choice: 0, probability: 1.0, finalists: Vec::new() });
429    }
430
431    #[test]
432    fn the_gate_lists_every_label_in_the_caller_order() {
433        let list = Shortlist::new(labels(300), Params::default()).unwrap();
434        let every = (0..300).map(|i| list.label(i)).collect::<Vec<_>>().join(", ");
435        assert_eq!(list.gate("q"), Some(crate::gate::gate("q", &every)));
436        assert_eq!(Shortlist::new(labels(1), Params::default()).unwrap().gate("q"), None);
437    }
438
439    #[test]
440    fn bad_input_is_22023() {
441        let p = Params::default();
442        let one = |k| Params { keep: NonZeroUsize::new(k).unwrap() };
443        let cases = [
444            Shortlist::new(Vec::new(), p).unwrap_err(),
445            Shortlist::new([("a".into(), None), (String::new(), None)], p).unwrap_err(),
446            Shortlist::new([("a".into(), None), ("b".into(), None), ("a".into(), None)], p).unwrap_err(),
447            // One chunk keeping one: no finalist Choice to ask.
448            Shortlist::new(labels(10), one(1)).unwrap_err(),
449            // Four chunks keeping 64 each: 256 finalists.
450            Shortlist::new(labels(1000), one(64)).unwrap_err(),
451        ];
452        assert_eq!(
453            cases,
454            [
455                ShortlistError::NoOptions,
456                ShortlistError::EmptyLabel { index: 1 },
457                ShortlistError::DuplicateLabel { label: "a".into() },
458                ShortlistError::Finalists { count: 1 },
459                ShortlistError::Finalists { count: 256 },
460            ]
461        );
462        assert!(cases.iter().all(|e| e.sqlstate() == "22023"));
463    }
464
465    #[test]
466    fn answers_are_checked() {
467        let list = Shortlist::new(labels(300), Params::default()).unwrap();
468        let mut s = Search::new(&list);
469        assert_eq!(s.answer(Ask::Finalists, &[]), Err(AnswerError::NotAsked(Ask::Finalists)));
470        assert_eq!(s.answer(Ask::Chunk(2), &[]), Err(AnswerError::NotAsked(Ask::Chunk(2))));
471        assert_eq!(
472            s.answer(Ask::Chunk(0), &[0.5]),
473            Err(AnswerError::Count { ask: Ask::Chunk(0), expected: 150, got: 1 })
474        );
475        let mut bad = vec![0.0; 150];
476        bad[3] = 1.5;
477        assert_eq!(
478            s.answer(Ask::Chunk(0), &bad),
479            Err(AnswerError::NotAProbability { ask: Ask::Chunk(0), value: 1.5 })
480        );
481        s.answer(Ask::Chunk(0), &vec![0.0; 150]).unwrap();
482        assert_eq!(s.answer(Ask::Chunk(0), &vec![0.0; 150]), Err(AnswerError::NotAsked(Ask::Chunk(0))));
483        assert_eq!(s.step(), Step::Ask(vec![Ask::Chunk(1)]));
484    }
485}