use super::{Policy, Resource};
use crate::morphology::filter::{AllocatedStream, ComposingToken, Work};
use crate::token::allocation::{TokenBatchAllocation, TokenBatchInput, TokenBuffer};
use crate::{AnalysisError, AnalysisResult};
use uqa_core::memory::{Budgeted, BudgetedVec, MemoryBudget, MemoryReservation};
struct OwnedToken<T> {
token: T,
memory: MemoryReservation,
}
impl<T> From<Budgeted<T>> for OwnedToken<T> {
fn from(input: Budgeted<T>) -> Self {
let (token, memory) = input.into_parts();
Self { token, memory }
}
}
impl<T> OwnedToken<T> {
fn into_budgeted(self) -> Budgeted<T> {
Budgeted::new(self.token, self.memory)
}
}
impl<T> std::ops::Deref for OwnedToken<T> {
type Target = T;
fn deref(&self) -> &T {
&self.token
}
}
struct State<'a, T: ComposingToken, P: Policy<T>> {
context: T::Context,
input: Option<TokenBatchInput<T>>,
current: Option<OwnedToken<T>>,
changed: bool,
saved: Option<OwnedToken<T>>,
numeral: BudgetedVec<u16>,
fall_through: u32,
exhausted: bool,
output: TokenBuffer<T>,
output_units: usize,
final_position_increment: u32,
budget: MemoryBudget,
policy: P,
work: Work<'a>,
}
pub(crate) fn filter<T: ComposingToken>(
input: AllocatedStream<T>,
policy: impl Policy<T>,
poll: &mut dyn FnMut() -> AnalysisResult<()>,
) -> AnalysisResult<AllocatedStream<T>> {
let work = Work::new(poll)?;
policy.check(Resource::Tokens, input.batch.tokens().len())?;
policy.check(Resource::InputUnits, input.final_offset_utf16)?;
let budget = input.batch.budget().clone();
let mut state = State {
context: input.context,
input: Some(input.batch.into_input()),
current: None,
changed: false,
saved: None,
numeral: BudgetedVec::new(&budget),
fall_through: 0,
exhausted: false,
output: TokenBuffer::new(&budget),
output_units: 0,
final_position_increment: 0,
budget,
policy,
work,
};
while state.next()? {
let token = state.current.take().expect("emitted attributes");
let length = token.term_len(&mut state.work)?;
state.output_units = state.total_units(&token, length)?;
state
.policy
.check(Resource::Tokens, state.output.len() + 1)?;
state.output.push(token.into_budgeted())?;
state.changed = false;
}
state.numeral = BudgetedVec::new(&state.budget);
state.saved = None;
if state.changed {
if let Some(token) = state.current.take() {
let length = token.term_len(&mut state.work)?;
state.total_units(&token, length)?;
state.output.set_terminal_token(token.into_budgeted())?;
}
}
state.work.finish()?;
Ok(AllocatedStream {
context: state.context,
batch: TokenBatchAllocation::from_budgeted(
state.output.into_batch(state.final_position_increment),
),
final_offset_utf16: input.final_offset_utf16,
})
}
impl<T: ComposingToken, P: Policy<T>> State<'_, T, P> {
fn total_units(&mut self, token: &T, term_units: usize) -> AnalysisResult<usize> {
let total = self
.output_units
.checked_add(self.policy.token_units(token, term_units, &mut self.work)?)
.ok_or_else(|| self.policy.invalid("attribute size overflow"))?;
self.policy.check(Resource::OutputUnits, total)?;
Ok(total)
}
fn read(&mut self) -> AnalysisResult<bool> {
self.work.tick()?;
let Some(input) = self.input.take() else {
return Ok(false);
};
let (token, input) = input.next(self.work.poll)?;
let (token, more) = if let Some(token) = token {
self.input = Some(input);
(Some(token), true)
} else {
let (terminal, increment) = input.finish();
self.final_position_increment = increment;
let token = terminal.map(|terminal| {
let (terminal, mut memory) = terminal.into_parts();
let token = {
let allocation = terminal;
*allocation
};
drop(memory.split(size_of::<T>()));
Budgeted::new(token, memory)
});
(token, false)
};
if let Some(token) = token {
let length = token.term_len(&mut self.work)?;
self.total_units(&token, length)?;
self.current = Some(token.into());
self.changed = true;
}
Ok(more)
}
fn next(&mut self) -> AnalysisResult<bool> {
self.work.tick()?;
if let Some(saved) = self.saved.take() {
self.current = Some(saved);
self.changed = true;
return Ok(true);
}
if self.exhausted {
return Ok(false);
}
if !self.read()? {
self.exhausted = true;
return Ok(false);
}
let current = self.current.as_ref().expect("read attributes");
if current.keyword() {
return Ok(true);
}
if self.fall_through > 0 {
self.fall_through -= 1;
return Ok(true);
}
if current.increment() == 0 {
self.fall_through = current
.position_length()
.checked_sub(1)
.ok_or(AnalysisError::InvalidTokenPosition)?;
return Ok(true);
}
if !super::numeral::<P::Symbols>(current.term(), &mut self.work)? {
return Ok(true);
}
self.compose()
}
fn compose(&mut self) -> AnalysisResult<bool> {
let original = self
.current
.as_ref()
.expect("numeric attributes")
.clone_reserved(&self.budget, &mut self.work)?;
let mut term = original.copy_term(&self.budget, &mut self.work)?;
let first = original.span();
let mut last;
let more = loop {
self.work.tick()?;
last = self.current.as_ref().expect("numeric attributes").span();
let more = self.read()?;
if !more {
self.exhausted = true;
}
let current = self
.current
.as_ref()
.expect("last successful or explicit terminal attributes");
if current.increment() == 0 {
self.fall_through = current
.position_length()
.checked_sub(1)
.ok_or(AnalysisError::InvalidTokenPosition)?;
self.saved = Some(current.clone_reserved(&self.budget, &mut self.work)?.into());
self.current = Some(original.into());
self.changed = true;
return Ok(more);
}
let required = self
.numeral
.len()
.checked_add(term.len())
.ok_or_else(|| self.policy.invalid("numeral size overflow"))?;
self.policy.check(Resource::NumericUnits, required)?;
self.numeral.reserve(term.len())?;
for unit in term.iter() {
self.work.tick()?;
self.numeral.push(*unit)?;
}
if !more {
break false;
}
term.clear();
term.reserve(current.term_len(&mut self.work)?)?;
for unit in current.term() {
self.work.tick()?;
term.push(unit)?;
}
if !super::numeral::<P::Symbols>(term.iter().copied(), &mut self.work)?
&& !super::punctuation::<P::Symbols>(term.iter().copied(), &mut self.work)?
{
break true;
}
};
drop(original);
drop(term);
if more {
self.saved = Some(
self.current
.as_ref()
.expect("lookahead attributes")
.clone_reserved(&self.budget, &mut self.work)?
.into(),
);
}
let mut current = self.current.take().expect("composed attributes");
current
.token
.cover(&first, &last, &self.context, &mut self.work)?;
let metadata = self.total_units(¤t, 0)?;
let maximum = self.policy.maximum_output() - metadata;
let numeral = std::mem::replace(&mut self.numeral, BudgetedVec::new(&self.budget));
let normalized = self
.policy
.normalize(&numeral, maximum, &self.budget, &mut self.work)?;
current.token.replace_term(
normalized,
&mut current.memory,
&self.context,
&mut self.work,
)?;
self.current = Some(current);
self.changed = true;
Ok(true)
}
}