use std::collections::BTreeMap;
use std::future::Future;
use std::marker::PhantomData;
use std::num::NonZeroUsize;
pub const MAX_MATERIALIZATION_CANDIDATES: usize = 4_096;
pub const MAX_MATERIALIZATION_LOADER_BATCH: usize = 256;
pub const MAX_MATERIALIZATION_OUTPUTS: usize = 4_096;
pub const MAX_MATERIALIZATION_DIAGNOSTICS: usize = 4_096;
pub const MAX_MATERIALIZATION_DROP_REASONS: usize = 32;
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum MaterializationLimitError {
#[error("materialization {field} limit {requested} exceeds v1 maximum {maximum}")]
AboveV1Maximum {
field: &'static str,
requested: usize,
maximum: usize,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct MaterializationLimits {
max_candidates: usize,
max_loader_batch_size: NonZeroUsize,
max_output_rows: usize,
max_diagnostic_details: usize,
}
impl MaterializationLimits {
pub fn try_new(
max_candidates: usize,
max_loader_batch_size: NonZeroUsize,
max_output_rows: usize,
max_diagnostic_details: usize,
) -> Result<Self, MaterializationLimitError> {
for (field, requested, maximum) in [
("candidates", max_candidates, MAX_MATERIALIZATION_CANDIDATES),
(
"loader_batch_size",
max_loader_batch_size.get(),
MAX_MATERIALIZATION_LOADER_BATCH,
),
("output_rows", max_output_rows, MAX_MATERIALIZATION_OUTPUTS),
(
"diagnostic_details",
max_diagnostic_details,
MAX_MATERIALIZATION_DIAGNOSTICS,
),
] {
if requested > maximum {
return Err(MaterializationLimitError::AboveV1Maximum {
field,
requested,
maximum,
});
}
}
Ok(Self {
max_candidates,
max_loader_batch_size,
max_output_rows,
max_diagnostic_details,
})
}
pub const fn max_candidates(self) -> usize {
self.max_candidates
}
pub const fn max_loader_batch_size(self) -> NonZeroUsize {
self.max_loader_batch_size
}
pub const fn max_output_rows(self) -> usize {
self.max_output_rows
}
pub const fn max_diagnostic_details(self) -> usize {
self.max_diagnostic_details
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RankedCandidate<Key, Score> {
pub key: Key,
pub score: Score,
}
pub trait DropReason: Copy + Eq + 'static {
const ALL: &'static [Self];
fn ordinal(self) -> usize;
}
#[derive(Debug, PartialEq, Eq)]
pub enum MaterializationDecision<Output, Reason, Error> {
Keep(Output),
Drop(Reason),
Fatal(Error),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MaterializedItem<Key, Score, Output> {
pub candidate: RankedCandidate<Key, Score>,
pub rank: usize,
pub output: Output,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DropDiagnostic<Key, Score, Reason> {
pub candidate: RankedCandidate<Key, Score>,
pub reason: Reason,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DropCounts<Reason> {
counts: [usize; MAX_MATERIALIZATION_DROP_REASONS],
marker: PhantomData<Reason>,
}
impl<Reason: DropReason> DropCounts<Reason> {
pub fn count(&self, reason: Reason) -> Option<usize> {
let ordinal = reason.ordinal();
Reason::ALL
.get(ordinal)
.filter(|declared| **declared == reason)
.map(|_| self.counts[ordinal])
}
pub fn total(&self) -> usize {
self.counts.iter().sum()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MaterializedPrefix<Key, Score, Output, Reason> {
pub accepted: Vec<MaterializedItem<Key, Score, Output>>,
pub drop_counts: DropCounts<Reason>,
pub diagnostic_details: Vec<DropDiagnostic<Key, Score, Reason>>,
pub diagnostics_truncated: bool,
}
#[derive(Debug, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum MaterializationError<Error> {
#[error("materialization request {field} {requested} exceeds configured limit {limit}")]
RequestExceedsLimit {
field: &'static str,
requested: usize,
limit: usize,
},
#[error("duplicate materialization candidate at indexes {first_index} and {duplicate_index}")]
DuplicateCandidate {
first_index: usize,
duplicate_index: usize,
},
#[error(
"materialization candidate order is not strict at indexes {previous_index} and {index}"
)]
NonMonotonicOrder {
previous_index: usize,
index: usize,
},
#[error("invalid materialization drop taxonomy: {message}")]
InvalidDropTaxonomy {
message: &'static str,
},
#[error("materialization loader returned an unexpected key")]
UnexpectedLoaderKey,
#[error("materialization loader returned a duplicate key")]
DuplicateLoaderKey,
#[error(
"materialization loader returned {returned} rows for a batch of {requested} candidates"
)]
LoaderReturnedTooManyRows {
returned: usize,
requested: usize,
},
#[error("materialization caller callback failed")]
Caller(Error),
}
#[allow(clippy::too_many_arguments)]
pub async fn materialize_ranked_prefix<
Key,
Score,
OrderKey,
Row,
Output,
Reason,
Error,
OrderFn,
Validator,
Loader,
LoaderFuture,
Classifier,
>(
candidates: Vec<RankedCandidate<Key, Score>>,
output_limit: usize,
loader_batch_size: NonZeroUsize,
limits: MaterializationLimits,
mut order_key: OrderFn,
mut candidate_validator: Validator,
mut batch_loader: Loader,
mut classifier: Classifier,
) -> Result<MaterializedPrefix<Key, Score, Output, Reason>, MaterializationError<Error>>
where
Key: Clone + Ord,
OrderKey: Ord,
Reason: DropReason,
OrderFn: FnMut(&RankedCandidate<Key, Score>) -> OrderKey,
Validator: FnMut(&RankedCandidate<Key, Score>) -> Result<(), Error>,
Loader: FnMut(Vec<Key>) -> LoaderFuture,
LoaderFuture: Future<Output = Result<Vec<(Key, Row)>, Error>>,
Classifier: FnMut(
&RankedCandidate<Key, Score>,
Option<Row>,
) -> MaterializationDecision<Output, Reason, Error>,
{
for (field, requested, limit) in [
("candidates", candidates.len(), limits.max_candidates()),
(
"loader_batch_size",
loader_batch_size.get(),
limits.max_loader_batch_size().get(),
),
("output_rows", output_limit, limits.max_output_rows()),
] {
if requested > limit {
return Err(MaterializationError::RequestExceedsLimit {
field,
requested,
limit,
});
}
}
if Reason::ALL.len() > MAX_MATERIALIZATION_DROP_REASONS {
return Err(MaterializationError::InvalidDropTaxonomy {
message: "drop taxonomy exceeds 32 variants",
});
}
for (ordinal, reason) in Reason::ALL.iter().copied().enumerate() {
if reason.ordinal() != ordinal {
return Err(MaterializationError::InvalidDropTaxonomy {
message: "drop taxonomy ordinals are not contiguous and ordered",
});
}
}
{
let mut first_indexes = BTreeMap::<&Key, usize>::new();
let mut previous_order = None;
for (index, candidate) in candidates.iter().enumerate() {
if let Some(first_index) = first_indexes.insert(&candidate.key, index) {
return Err(MaterializationError::DuplicateCandidate {
first_index,
duplicate_index: index,
});
}
let current_order = order_key(candidate);
if previous_order
.as_ref()
.is_some_and(|previous| previous >= ¤t_order)
{
return Err(MaterializationError::NonMonotonicOrder {
previous_index: index - 1,
index,
});
}
previous_order = Some(current_order);
}
}
let accepted_capacity = output_limit.min(candidates.len());
let diagnostic_capacity = limits.max_diagnostic_details().min(candidates.len());
let mut accepted = Vec::with_capacity(accepted_capacity);
let mut diagnostic_details = Vec::with_capacity(diagnostic_capacity);
let mut counts = [0_usize; MAX_MATERIALIZATION_DROP_REASONS];
let mut diagnostics_truncated = false;
let mut remaining = candidates.into_iter();
while accepted.len() < output_limit {
let batch = remaining
.by_ref()
.take(loader_batch_size.get())
.collect::<Vec<_>>();
if batch.is_empty() {
break;
}
for candidate in &batch {
candidate_validator(candidate).map_err(MaterializationError::Caller)?;
}
let loader_keys = batch
.iter()
.map(|candidate| candidate.key.clone())
.collect::<Vec<_>>();
let rows = batch_loader(loader_keys)
.await
.map_err(MaterializationError::Caller)?;
if rows.len() > batch.len() {
return Err(MaterializationError::LoaderReturnedTooManyRows {
returned: rows.len(),
requested: batch.len(),
});
}
let mut indexes = BTreeMap::<&Key, usize>::new();
for (index, candidate) in batch.iter().enumerate() {
indexes.insert(&candidate.key, index);
}
let mut row_slots = (0..batch.len()).map(|_| None).collect::<Vec<Option<Row>>>();
let mut saw_unexpected = false;
let mut saw_duplicate = false;
for (key, row) in rows {
match indexes.get(&key).copied() {
Some(index) if row_slots[index].is_none() => row_slots[index] = Some(row),
Some(_) => saw_duplicate = true,
None => saw_unexpected = true,
}
}
drop(indexes);
if saw_unexpected {
return Err(MaterializationError::UnexpectedLoaderKey);
}
if saw_duplicate {
return Err(MaterializationError::DuplicateLoaderKey);
}
for (candidate, row) in batch.into_iter().zip(row_slots) {
if accepted.len() == output_limit {
break;
}
match classifier(&candidate, row) {
MaterializationDecision::Keep(output) => {
accepted.push(MaterializedItem {
candidate,
rank: accepted.len() + 1,
output,
});
}
MaterializationDecision::Drop(reason) => {
let ordinal = reason.ordinal();
if Reason::ALL.get(ordinal).copied() != Some(reason) {
return Err(MaterializationError::InvalidDropTaxonomy {
message: "classifier returned an undeclared drop reason",
});
}
counts[ordinal] += 1;
if diagnostic_details.len() < limits.max_diagnostic_details() {
diagnostic_details.push(DropDiagnostic { candidate, reason });
} else {
diagnostics_truncated = true;
}
}
MaterializationDecision::Fatal(error) => {
return Err(MaterializationError::Caller(error));
}
}
}
}
for candidate in remaining {
candidate_validator(&candidate).map_err(MaterializationError::Caller)?;
}
Ok(MaterializedPrefix {
accepted,
drop_counts: DropCounts {
counts,
marker: PhantomData,
},
diagnostic_details,
diagnostics_truncated,
})
}
#[cfg(test)]
#[path = "materialization_tests.rs"]
mod tests;