use super::*;
#[test]
fn removing_a_row_never_reassigns_another_rows_state() {
let mut rows = Rows::default();
let (a, _ra) = test_slot(3, 11);
let (b, _rb) = test_slot(5, 22);
let (c, _rc) = test_slot(7, 33);
let a = rows.insert(a);
let b = rows.insert(b);
let c = rows.insert(c);
assert_eq!(rows.order, vec![a, b, c]);
let removed = rows.remove(b).expect("b was present");
assert_eq!(removed.max_tokens, 5);
assert!(
rows.get(b).is_none(),
"a stale uid must resolve to nothing, never to another request's row"
);
assert_eq!(rows.get(a).expect("a still in flight").max_tokens, 3);
assert_eq!(
rows.get(c).expect("c still in flight").max_tokens,
7,
"c must still be c after b left"
);
assert_eq!(rows.order, vec![a, c], "admission order is preserved");
assert_eq!(rows.len(), 2);
let mut positional = vec![3usize, 5, 7];
let c_index = 2;
positional.swap_remove(1);
assert_eq!(positional[1], 7, "C moved into B's index");
assert!(
positional.get(c_index).is_none(),
"C's index now names nothing"
);
}
#[test]
fn uids_are_unique_and_insertion_does_not_disturb_existing_rows() {
let mut rows = Rows::default();
let (a, _ra) = test_slot(3, 11);
let a = rows.insert(a);
let (b, _rb) = test_slot(5, 22);
let b = rows.insert(b);
rows.remove(a);
let (c, _rc) = test_slot(7, 33);
let c = rows.insert(c);
assert_ne!(c, a, "a uid is never reused after its row leaves");
assert_ne!(c, b);
assert_eq!(rows.get(b).expect("b untouched").max_tokens, 5);
assert_eq!(rows.get(c).expect("c inserted").max_tokens, 7);
}
#[test]
fn flush_replies_on_each_rows_own_channel() {
let mut rows = Rows::default();
let (a, ra) = test_slot(3, 11);
let (mut b, rb) = test_slot(5, 22);
b.finish = Some(FinishReason::Stop);
b.visible.push_str("bee");
b.generated_ids.push(7);
let a = rows.insert(a);
let b = rows.insert(b);
let (c, _rc) = test_slot(7, 33);
let c = rows.insert(c);
assert_eq!(rows.ready(), vec![a, c], "a finished row takes no step");
rows.flush_finished(&no_budget());
assert!(rows.get(b).is_none());
assert_eq!(rows.order, vec![a, c]);
let (finish, ids, text, usage) =
finished_result(rb.try_recv().expect("b's caller got a reply")).expect("ok");
assert_eq!(finish, FinishReason::Stop);
assert_eq!(ids, vec![7]);
assert_eq!(text, "bee");
assert_eq!(usage.completion_tokens, 1);
assert!(
ra.try_recv().is_err(),
"an unfinished row's caller must not be replied to"
);
}
#[test]
fn a_row_leaving_mid_batch_does_not_shift_its_neighbours_output() {
let decoder = tiny_decoder();
let prompts = [vec![1usize, 2, 3], vec![4usize, 5], vec![6usize]];
let budgets = [25usize, 25, 20];
let refs: Vec<Vec<usize>> = prompts
.iter()
.zip(budgets.iter())
.map(|(p, &n)| sequential_ids(&decoder, p, &greedy_params(n, 4)))
.collect();
let letter = |id: &usize| char::from_u32(65 + (*id as u32 % 26)).unwrap_or('?');
let middle_text: String = refs[1].iter().map(letter).collect();
assert!(middle_text.len() >= 4);
let stop = middle_text[2..4].to_string();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let barrier = Arc::new(Barrier::new(prompts.len()));
let handles: Vec<_> = (0..prompts.len())
.map(|i| {
let batcher = batcher.clone();
let barrier = Arc::clone(&barrier);
let prompt = prompts[i].clone();
let mut params = greedy_params(budgets[i], 4);
if i == 1 {
params.stop = vec![stop.clone()];
}
thread::spawn(move || {
barrier.wait();
batcher
.generate(prompt, params, StopTokens::default())
.expect("generate")
.1
})
})
.collect();
let got: Vec<Vec<usize>> = handles.into_iter().map(|h| h.join().unwrap()).collect();
assert_eq!(got[0], refs[0], "row 0 received another row's output");
assert_eq!(got[2], refs[2], "row 2 received another row's output");
assert!(
got[1].len() < refs[1].len() && refs[1].starts_with(&got[1]),
"the stopped row must be a strict prefix of its own stream"
);
}