use std::time::Duration;
use super::generation_receipt::{GenerationCommit, GenerationReceiptConfig};
pub(crate) const GENERATION_COMMIT_INTERVAL: Duration = Duration::from_millis(100);
pub(crate) const PENDING_CAPACITY: usize = 256;
pub(crate) struct GenerationCommitBatcher<'config> {
config: Option<&'config GenerationReceiptConfig>,
request_id: u64,
session_id: u64,
pending: Vec<i32>,
generated_token_count: usize,
last_emit: Option<Duration>,
}
impl<'config> GenerationCommitBatcher<'config> {
#[must_use]
pub(crate) fn new(
config: Option<&'config GenerationReceiptConfig>,
request_id: u64,
session_id: u64,
) -> Self {
Self {
config,
request_id,
session_id,
pending: if config.is_some() {
Vec::with_capacity(PENDING_CAPACITY)
} else {
Vec::new()
},
generated_token_count: 0,
last_emit: None,
}
}
pub(crate) fn commit(&mut self, token_id: i32, elapsed: Duration) {
let Some(config) = self.config else {
return;
};
self.generated_token_count = self.generated_token_count.saturating_add(1);
self.pending.push(token_id);
if self.is_due(elapsed) {
emit(
config,
self.request_id,
self.session_id,
self.generated_token_count,
&mut self.pending,
);
self.last_emit = Some(elapsed);
}
}
pub(crate) fn flush(&mut self) {
let Some(config) = self.config else {
return;
};
if self.pending.is_empty() {
return;
}
emit(
config,
self.request_id,
self.session_id,
self.generated_token_count,
&mut self.pending,
);
}
#[cfg(test)]
#[must_use]
pub(crate) fn generated_token_count(&self) -> usize {
self.generated_token_count
}
fn is_due(&self, elapsed: Duration) -> bool {
if self.pending.len() >= PENDING_CAPACITY {
return true;
}
match self.last_emit {
None => true,
Some(last) => elapsed.saturating_sub(last) >= GENERATION_COMMIT_INTERVAL,
}
}
}
impl Drop for GenerationCommitBatcher<'_> {
fn drop(&mut self) {
self.flush();
}
}
fn emit(
config: &GenerationReceiptConfig,
request_id: u64,
session_id: u64,
generated_token_count: usize,
pending: &mut Vec<i32>,
) {
let token_ids: Box<[i32]> = Box::from(pending.as_slice());
pending.clear();
config.committed(GenerationCommit {
request_id,
session_id,
generated_token_count,
token_ids,
});
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use anyhow::Result;
use super::*;
use crate::frontend::generation_receipt::{
GenerationLifecycleIngress, GenerationLifecycleObservation,
};
#[derive(Default)]
struct RecordingIngress {
commits: Mutex<Vec<(usize, Vec<i32>)>>,
}
impl GenerationLifecycleIngress for RecordingIngress {
fn try_submit(&self, observation: GenerationLifecycleObservation) -> Result<()> {
if let GenerationLifecycleObservation::Committed(commit) = observation {
self.commits
.lock()
.unwrap()
.push((commit.generated_token_count, commit.token_ids.to_vec()));
}
Ok(())
}
}
impl RecordingIngress {
fn commits(&self) -> Vec<(usize, Vec<i32>)> {
self.commits.lock().unwrap().clone()
}
}
fn config(ingress: &Arc<RecordingIngress>) -> GenerationReceiptConfig {
GenerationReceiptConfig::from_lifecycle_ingress(Arc::clone(ingress) as Arc<_>)
}
fn at(millis: u64) -> Duration {
Duration::from_millis(millis)
}
#[test]
fn the_first_token_emits_immediately() {
let ingress = Arc::new(RecordingIngress::default());
let config = config(&ingress);
let mut batcher = GenerationCommitBatcher::new(Some(&config), 7, 9);
batcher.commit(11, at(0));
assert_eq!(ingress.commits(), vec![(1, vec![11])]);
}
#[test]
fn tokens_inside_one_window_accumulate_into_one_commit() {
let ingress = Arc::new(RecordingIngress::default());
let config = config(&ingress);
let mut batcher = GenerationCommitBatcher::new(Some(&config), 7, 9);
batcher.commit(1, at(0));
for (index, token) in [2, 3, 4].into_iter().enumerate() {
batcher.commit(token, at(10 * (index as u64 + 1)));
}
batcher.commit(5, at(100));
assert_eq!(
ingress.commits(),
vec![(1, vec![1]), (5, vec![2, 3, 4, 5])],
"the second commit carries the whole window and the total after it"
);
}
#[test]
fn concatenated_batches_reproduce_the_token_sequence_exactly() {
let ingress = Arc::new(RecordingIngress::default());
let config = config(&ingress);
let tokens: Vec<i32> = (0..500).collect();
{
let mut batcher = GenerationCommitBatcher::new(Some(&config), 7, 9);
for (index, token) in tokens.iter().enumerate() {
batcher.commit(*token, at(3 * index as u64));
}
batcher.flush();
assert_eq!(batcher.generated_token_count(), tokens.len());
}
let commits = ingress.commits();
let observed: Vec<i32> = commits
.iter()
.flat_map(|(_, batch)| batch.iter().copied())
.collect();
assert_eq!(observed, tokens);
assert!(
commits.len() < tokens.len(),
"batching must produce fewer commits than tokens; got {} for {}",
commits.len(),
tokens.len()
);
let totals: Vec<usize> = commits.iter().map(|(total, _)| *total).collect();
let mut running = 0usize;
let expected: Vec<usize> = commits
.iter()
.map(|(_, batch)| {
running += batch.len();
running
})
.collect();
assert_eq!(
totals, expected,
"generated_token_count must be the running total after each batch"
);
}
#[test]
fn reaching_the_buffer_capacity_commits_early() {
let ingress = Arc::new(RecordingIngress::default());
let config = config(&ingress);
let mut batcher = GenerationCommitBatcher::new(Some(&config), 7, 9);
for token in 0..=PENDING_CAPACITY as i32 {
batcher.commit(token, at(0));
}
let commits = ingress.commits();
assert_eq!(commits.len(), 2, "first token, then the capacity ceiling");
assert_eq!(commits[1].1.len(), PENDING_CAPACITY);
}
#[test]
fn flush_emits_the_residual_batch_exactly_once() {
let ingress = Arc::new(RecordingIngress::default());
let config = config(&ingress);
let mut batcher = GenerationCommitBatcher::new(Some(&config), 7, 9);
batcher.commit(1, at(0));
batcher.commit(2, at(5));
batcher.flush();
batcher.flush();
assert_eq!(ingress.commits(), vec![(1, vec![1]), (2, vec![2])]);
}
#[test]
fn dropping_flushes_the_residual_batch() {
let ingress = Arc::new(RecordingIngress::default());
let config = config(&ingress);
{
let mut batcher = GenerationCommitBatcher::new(Some(&config), 7, 9);
batcher.commit(1, at(0));
batcher.commit(2, at(5));
}
assert_eq!(ingress.commits(), vec![(1, vec![1]), (2, vec![2])]);
}
#[test]
fn no_receipt_config_retains_nothing() {
let mut batcher = GenerationCommitBatcher::new(None, 7, 9);
batcher.commit(1, at(0));
batcher.commit(2, at(500));
batcher.flush();
assert_eq!(batcher.generated_token_count(), 0);
}
}