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}