use crate::{
context::project_text_tokens,
output::{ContextUsageSource, OutputEvent},
};
pub(super) struct ContextProjectionTracker<'a> {
provider_id: &'a str,
model: &'a str,
request_sequence: u64,
request_input_tokens: usize,
max_tokens: usize,
has_output: bool,
pending_delta_text: String,
projected_output_tokens: usize,
projection_source: ContextUsageSource,
last_emitted_tokens: usize,
}
impl<'a> ContextProjectionTracker<'a> {
const FLUSH_BYTES: usize = 256;
const EMIT_TOKEN_DELTA: usize = 16;
pub(super) fn new(
provider_id: &'a str,
model: &'a str,
request_sequence: u64,
request_input_tokens: usize,
max_tokens: usize,
) -> Self {
Self {
provider_id,
model,
request_sequence,
request_input_tokens,
max_tokens,
has_output: false,
pending_delta_text: String::new(),
projected_output_tokens: 0,
projection_source: ContextUsageSource::FallbackProjection,
last_emitted_tokens: request_input_tokens,
}
}
pub(super) fn update_after_delta(&mut self, delta: &str) -> Option<OutputEvent> {
self.has_output = true;
self.pending_delta_text.push_str(delta);
if self.pending_delta_text.len() < Self::FLUSH_BYTES {
return None;
}
let pending_delta = std::mem::take(&mut self.pending_delta_text);
let projection = project_text_tokens(self.provider_id, self.model, &pending_delta);
self.projected_output_tokens = self
.projected_output_tokens
.saturating_add(projection.tokens);
self.projection_source = projection.source;
if self.projected_output_tokens.abs_diff(
self.last_emitted_tokens
.saturating_sub(self.request_input_tokens),
) < Self::EMIT_TOKEN_DELTA
{
return None;
}
let event = self.project_event_from_incremental();
let current_tokens = event_current_tokens(&event);
(current_tokens.abs_diff(self.last_emitted_tokens) >= Self::EMIT_TOKEN_DELTA).then(|| {
self.last_emitted_tokens = current_tokens;
event
})
}
pub(super) fn force_event_if_changed(&mut self) -> Option<OutputEvent> {
if !self.has_output {
return None;
}
if !self.pending_delta_text.is_empty() {
let pending_delta = std::mem::take(&mut self.pending_delta_text);
let projection = project_text_tokens(self.provider_id, self.model, &pending_delta);
self.projected_output_tokens = self
.projected_output_tokens
.saturating_add(projection.tokens);
self.projection_source = projection.source;
}
let event = self.project_event_from_incremental();
let current_tokens = event_current_tokens(&event);
(current_tokens != self.last_emitted_tokens).then(|| {
self.last_emitted_tokens = current_tokens;
event
})
}
fn project_event_from_incremental(&self) -> OutputEvent {
OutputEvent::ContextUsage {
current_tokens: self
.request_input_tokens
.saturating_add(self.projected_output_tokens),
max_tokens: self.max_tokens,
reasoning_tokens: None,
source: self.projection_source,
request_sequence: self.request_sequence,
}
}
}
fn event_current_tokens(event: &OutputEvent) -> usize {
match event {
OutputEvent::ContextUsage { current_tokens, .. } => *current_tokens,
_ => unreachable!("context projection tracker only emits context usage"),
}
}