use crate::SheetId;
use crate::engine::arena::value_ref::ValueType;
use crate::engine::arena::{AstNodeData, AstNodeId, CompactRefType, DataStore, SheetKey};
use crate::engine::template::canonical::{CanonicalExpr, LiteralSlotDescriptor};
use crate::engine::template::domain::ValueRefSlotDescriptor;
use crate::engine::template::read_summary::{ProjectionFallbackReason, ReadProjection};
use crate::reference::CellRef;
use formualizer_common::LiteralValue;
use rustc_hash::FxHashMap;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use super::arena::CanonicalLabels;
pub(crate) const MAX_SHAPES: usize = 16_384;
pub(crate) const MAX_BUCKET_SHAPES: usize = 4;
pub(crate) const MAX_SEEN: usize = 4 * MAX_SHAPES;
pub(crate) const MAX_SPECIALIZATIONS: usize = MAX_SHAPES;
pub(crate) const MAX_STORED_TOKENS: usize = 1 << 20;
pub(crate) const MAX_SHAPE_TOKENS: usize = 4_096;
pub(crate) const MAX_SHAPE_REFS: usize = 1_024;
pub(crate) const MAX_SPECIALIZATION_BYTES: usize = 256 << 10;
pub(crate) const MAX_STORED_BYTES: usize = 64 << 20;
pub(crate) const PROBE_WARMUP: u64 = 64;
pub(crate) const MAX_PROBE_STRIDE: u64 = 64;
const T_EMPTY: u64 = 1;
const T_INT: u64 = 2;
const T_NUMBER: u64 = 3;
const T_TEXT: u64 = 4;
const T_BOOL: u64 = 5;
const T_OMITTED: u64 = 6;
const T_CELL: u64 = 7;
const T_RANGE: u64 = 8;
const T_NAME: u64 = 9;
const T_UNARY: u64 = 10;
const T_BINARY: u64 = 11;
const T_FUNCTION: u64 = 12;
const T_ARRAY: u64 = 13;
pub(crate) struct Specialization {
pub(crate) canonical_hash: u64,
pub(crate) exact_canonical_hash: u64,
pub(crate) exact_canonical_key: Arc<str>,
pub(crate) parameterized_canonical_hash: u64,
pub(crate) parameterized_canonical_key: Arc<str>,
pub(crate) literal_slot_descriptors: Arc<[LiteralSlotDescriptor]>,
pub(crate) literal_bindings: Box<[LiteralValue]>,
pub(crate) value_ref_slot_descriptors: Arc<[ValueRefSlotDescriptor]>,
pub(crate) expr: CanonicalExpr,
pub(crate) labels: CanonicalLabels,
pub(crate) read_projections: Option<Vec<ReadProjection>>,
pub(crate) read_projection_fallback: Option<ProjectionFallbackReason>,
pub(crate) volatile: bool,
pub(crate) dynamic: bool,
pub(crate) visit: Box<[u32]>,
}
impl Specialization {
pub(crate) fn retained_bytes(&self, shape_tokens: usize) -> usize {
use std::mem::size_of;
let literal_text: usize = self
.literal_bindings
.iter()
.map(|value| match value {
LiteralValue::Text(text) => text.len(),
LiteralValue::Error(error) => error.message.as_ref().map_or(0, String::len) + 64,
_ => 0,
})
.sum();
let exact = self.exact_canonical_key.len();
size_of::<Self>()
+ exact
+ self.parameterized_canonical_key.len()
+ exact
+ shape_tokens * size_of::<CanonicalExpr>()
+ self.literal_bindings.len() * size_of::<LiteralValue>()
+ literal_text
+ self.literal_slot_descriptors.len() * size_of::<LiteralSlotDescriptor>()
+ self.value_ref_slot_descriptors.len() * size_of::<ValueRefSlotDescriptor>()
+ exact
+ self
.read_projections
.as_ref()
.map_or(0, |p| p.len() * size_of::<ReadProjection>())
+ self.visit.len() * size_of::<u32>()
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(crate) struct MemoCounts {
pub(crate) shape_hits: u64,
pub(crate) shape_misses: u64,
pub(crate) specialization_hits: u64,
pub(crate) specialization_misses: u64,
pub(crate) bypasses: u64,
pub(crate) first_sightings: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct MemoValidity {
pub(crate) provider_revision: Option<u64>,
pub(crate) registry_epoch: u64,
}
impl MemoValidity {
pub(crate) fn current(provider: &dyn crate::traits::FunctionProvider) -> Self {
Self {
provider_revision: provider.planning_semantic_revision(),
registry_epoch: crate::function_registry::semantic_epoch_lock_free(),
}
}
}
pub(crate) struct ShapeMemo {
pub(crate) counts: MemoCounts,
validity: Option<MemoValidity>,
seed: u64,
seen: rustc_hash::FxHashSet<u64>,
buckets: FxHashMap<u64, Vec<u32>>,
shapes: Vec<Box<[u64]>>,
specializations: FxHashMap<(u32, SheetId), Option<Specialization>>,
stored_tokens: usize,
stored_bytes: usize,
since_hit: u64,
pub(crate) tokens: Vec<u64>,
pub(crate) refs: Vec<CompactRefType>,
#[cfg(test)]
pub(crate) forced_hash: Option<u64>,
#[cfg(test)]
pub(crate) comparisons: u64,
#[cfg(test)]
pub(crate) key_walks: u64,
}
#[cfg(test)]
thread_local! {
pub(crate) static FORCED_HASH: std::cell::Cell<Option<u64>> =
const { std::cell::Cell::new(None) };
}
impl Default for ShapeMemo {
fn default() -> Self {
use std::hash::BuildHasher;
Self {
counts: MemoCounts::default(),
validity: None,
seed: std::hash::RandomState::new().hash_one(0x5a17_u64),
seen: Default::default(),
buckets: Default::default(),
shapes: Vec::new(),
specializations: Default::default(),
stored_tokens: 0,
stored_bytes: 0,
since_hit: 0,
tokens: Vec::new(),
refs: Vec::new(),
#[cfg(test)]
forced_hash: FORCED_HASH.with(std::cell::Cell::get),
#[cfg(test)]
comparisons: 0,
#[cfg(test)]
key_walks: 0,
}
}
}
pub(crate) enum ShapeLookup {
Shape(u32, bool),
FirstSighting,
Full,
}
impl ShapeMemo {
pub(crate) fn revalidate(&mut self, validity: MemoValidity) {
if self.validity != Some(validity) {
self.seen.clear();
self.buckets.clear();
self.shapes.clear();
self.specializations.clear();
self.stored_tokens = 0;
self.stored_bytes = 0;
self.validity = Some(validity);
}
}
pub(crate) fn should_probe(&mut self) -> bool {
self.since_hit += 1;
let past = self.since_hit.saturating_sub(PROBE_WARMUP);
if past == 0 {
return true;
}
let doublings = (past - 1) / PROBE_WARMUP + 1;
let stride = 1u64 << doublings.min(u64::from(MAX_PROBE_STRIDE.trailing_zeros()));
past.is_multiple_of(stride)
}
pub(crate) fn record_hit(&mut self) {
self.since_hit = 0;
}
fn bucket_hash(&self) -> u64 {
#[cfg(test)]
if let Some(hash) = self.forced_hash {
return hash;
}
let mut hasher = rustc_hash::FxHasher::default();
self.tokens.hash(&mut hasher);
fmix64(hasher.finish() ^ self.seed)
}
pub(crate) fn lookup_shape(&mut self) -> ShapeLookup {
if self.tokens.len() > MAX_SHAPE_TOKENS {
return ShapeLookup::Full;
}
let hash = self.bucket_hash();
let bucket_len = match self.buckets.get(&hash) {
Some(bucket) => {
for &index in bucket {
#[cfg(test)]
{
self.comparisons += 1;
}
if *self.shapes[index as usize] == *self.tokens {
return ShapeLookup::Shape(index, false);
}
}
bucket.len()
}
None => 0,
};
if self.shapes.len() >= MAX_SHAPES
|| bucket_len >= MAX_BUCKET_SHAPES
|| self.stored_tokens + self.tokens.len() > MAX_STORED_TOKENS
{
return ShapeLookup::Full;
}
if !self.seen.contains(&hash) {
if self.seen.len() >= MAX_SEEN {
return ShapeLookup::Full;
}
self.seen.insert(hash);
return ShapeLookup::FirstSighting;
}
let index = self.shapes.len() as u32;
self.shapes.push(self.tokens.clone().into_boxed_slice());
self.stored_tokens += self.tokens.len();
self.buckets.entry(hash).or_default().push(index);
ShapeLookup::Shape(index, true)
}
pub(crate) fn specialization(
&self,
shape: u32,
sheet: SheetId,
) -> Option<&Option<Specialization>> {
self.specializations.get(&(shape, sheet))
}
pub(crate) fn can_insert_specialization(&self, shape: u32) -> bool {
self.specializations.len() < MAX_SPECIALIZATIONS
&& self.stored_tokens + self.shapes[shape as usize].len() <= MAX_STORED_TOKENS
}
pub(crate) fn insert_specialization(
&mut self,
shape: u32,
sheet: SheetId,
specialization: Option<Specialization>,
) {
if !self.can_insert_specialization(shape) {
return;
}
let shape_tokens = self.shapes[shape as usize].len();
let specialization = specialization.filter(|specialization| {
let bytes = specialization.retained_bytes(shape_tokens);
if bytes > MAX_SPECIALIZATION_BYTES || self.stored_bytes + bytes > MAX_STORED_BYTES {
return false;
}
self.stored_bytes += bytes;
true
});
self.stored_tokens += shape_tokens;
self.specializations.insert((shape, sheet), specialization);
}
#[cfg(test)]
pub(crate) fn stored_bytes(&self) -> usize {
self.stored_bytes
}
#[cfg(test)]
pub(crate) fn retained_specializations(&self) -> usize {
self.specializations
.values()
.filter(|s| s.is_some())
.count()
}
#[cfg(test)]
pub(crate) fn footprint(&self) -> (usize, usize, usize, usize) {
(
self.shapes.len(),
self.specializations.len(),
self.seen.len(),
self.stored_tokens,
)
}
}
fn fmix64(mut h: u64) -> u64 {
h ^= h >> 33;
h = h.wrapping_mul(0xff51_afd7_ed55_8ccd);
h ^= h >> 33;
h = h.wrapping_mul(0xc4ce_b9fe_1a85_ec53);
h ^ (h >> 33)
}
pub(crate) fn shape_tokens(
data_store: &DataStore,
root: AstNodeId,
placement: CellRef,
tokens: &mut Vec<u64>,
refs: &mut Vec<CompactRefType>,
) -> bool {
tokens.clear();
refs.clear();
let anchor_row = placement.coord.row() + 1;
let anchor_col = placement.coord.col() + 1;
let eligible = walk(data_store, root, anchor_row, anchor_col, tokens, refs);
if !eligible {
tokens.clear();
refs.clear();
}
if tokens.capacity() > MAX_SHAPE_TOKENS {
tokens.shrink_to(MAX_SHAPE_TOKENS);
}
if refs.capacity() > MAX_SHAPE_REFS {
refs.shrink_to(MAX_SHAPE_REFS);
}
eligible
}
fn emit<const N: usize>(tokens: &mut Vec<u64>, items: [u64; N]) -> bool {
if tokens.len() + N > MAX_SHAPE_TOKENS {
return false;
}
tokens.extend(items);
true
}
fn sheet_token(sheet: Option<SheetKey>) -> u64 {
match sheet {
None => 0,
Some(SheetKey::Id(id)) => (1 << 32) | u64::from(id),
Some(SheetKey::Name(name)) => (2 << 32) | u64::from(name.as_u32()),
}
}
fn axis_token(value: u32, anchor: u32, absolute: bool) -> u64 {
if absolute {
u64::from(value)
} else {
(i64::from(value) - i64::from(anchor)) as u64
}
}
fn walk(
data_store: &DataStore,
id: AstNodeId,
anchor_row: u32,
anchor_col: u32,
tokens: &mut Vec<u64>,
refs: &mut Vec<CompactRefType>,
) -> bool {
let Some(node) = data_store.get_node(id) else {
return false;
};
match *node {
AstNodeData::Literal(value_ref) => match value_ref.value_type() {
ValueType::Empty => {
if !emit(tokens, [T_EMPTY]) {
return false;
}
}
ValueType::SmallInt | ValueType::LargeInt => {
let LiteralValue::Int(value) = data_store.retrieve_value(value_ref) else {
return false;
};
if !emit(tokens, [T_INT, value as u64]) {
return false;
}
}
ValueType::Number => {
let LiteralValue::Number(value) = data_store.retrieve_value(value_ref) else {
return false;
};
if !emit(tokens, [T_NUMBER, value.to_bits()]) {
return false;
}
}
ValueType::String => {
if !emit(tokens, [T_TEXT, u64::from(value_ref.as_raw())]) {
return false;
}
}
ValueType::Boolean => {
if !emit(tokens, [T_BOOL, u64::from(value_ref.as_raw())]) {
return false;
}
}
_ => return false,
},
AstNodeData::Omitted => {
if !emit(tokens, [T_OMITTED]) {
return false;
}
}
AstNodeData::Reference {
original_id,
ref_type,
} => {
if data_store
.resolve_ast_string(original_id)
.trim_end()
.ends_with('#')
{
return false;
}
match ref_type {
CompactRefType::Cell {
sheet,
row,
col,
row_abs,
col_abs,
} => {
if !emit(
tokens,
[
T_CELL,
sheet_token(sheet),
u64::from(row_abs) | (u64::from(col_abs) << 1),
axis_token(row, anchor_row, row_abs),
axis_token(col, anchor_col, col_abs),
],
) {
return false;
}
}
CompactRefType::Range {
sheet,
start_row,
start_col,
end_row,
end_col,
start_row_abs,
start_col_abs,
end_row_abs,
end_col_abs,
} => {
let rows_open = (start_row == 0, end_row == u32::MAX);
let cols_open = (start_col == 0, end_col == u32::MAX);
if rows_open.0 != rows_open.1 || cols_open.0 != cols_open.1 {
return false;
}
let flags = u64::from(start_row_abs)
| (u64::from(start_col_abs) << 1)
| (u64::from(end_row_abs) << 2)
| (u64::from(end_col_abs) << 3)
| (u64::from(rows_open.0) << 4)
| (u64::from(cols_open.0) << 5);
let (sr, er) = if rows_open.0 {
(0, 0)
} else {
(
axis_token(start_row, anchor_row, start_row_abs),
axis_token(end_row, anchor_row, end_row_abs),
)
};
let (sc, ec) = if cols_open.0 {
(0, 0)
} else {
(
axis_token(start_col, anchor_col, start_col_abs),
axis_token(end_col, anchor_col, end_col_abs),
)
};
if !emit(tokens, [T_RANGE, sheet_token(sheet), flags, sr, sc, er, ec]) {
return false;
}
}
CompactRefType::NamedRange(name) => {
if !emit(tokens, [T_NAME, u64::from(name.as_u32())]) {
return false;
}
}
CompactRefType::External { .. }
| CompactRefType::Table { .. }
| CompactRefType::Cell3D { .. }
| CompactRefType::Range3D { .. } => return false,
}
if refs.len() >= MAX_SHAPE_REFS {
return false;
}
refs.push(ref_type);
}
AstNodeData::UnaryOp { op_id, expr_id } => {
if !emit(tokens, [T_UNARY, u64::from(op_id.as_u32())]) {
return false;
}
return walk(data_store, expr_id, anchor_row, anchor_col, tokens, refs);
}
AstNodeData::BinaryOp {
op_id,
left_id,
right_id,
} => {
if !emit(tokens, [T_BINARY, u64::from(op_id.as_u32())]) {
return false;
}
return walk(data_store, left_id, anchor_row, anchor_col, tokens, refs)
&& walk(data_store, right_id, anchor_row, anchor_col, tokens, refs);
}
AstNodeData::Function { name_id, .. } => {
let Some(args) = data_store.get_args(id) else {
return false;
};
if !emit(
tokens,
[T_FUNCTION, u64::from(name_id.as_u32()), args.len() as u64],
) {
return false;
}
return args
.iter()
.all(|&arg| walk(data_store, arg, anchor_row, anchor_col, tokens, refs));
}
AstNodeData::Array { rows, cols, .. } => {
let Some((_, _, elements)) = data_store.get_array_elems(id) else {
return false;
};
if !emit(
tokens,
[
T_ARRAY,
u64::from(rows),
u64::from(cols),
elements.len() as u64,
],
) {
return false;
}
return elements
.iter()
.all(|&element| walk(data_store, element, anchor_row, anchor_col, tokens, refs));
}
}
true
}
pub(crate) fn tree_reference_keys(ast: &formualizer_parse::parser::ASTNode, out: &mut Vec<usize>) {
use formualizer_parse::parser::{ASTNodeType, ReferenceType};
match &ast.node_type {
ASTNodeType::Reference { reference, .. } => out.push(match reference {
ReferenceType::NamedRange(name) => name.as_ptr() as usize,
other => other as *const ReferenceType as usize,
}),
ASTNodeType::UnaryOp { expr, .. } => tree_reference_keys(expr, out),
ASTNodeType::BinaryOp { left, right, .. } => {
tree_reference_keys(left, out);
tree_reference_keys(right, out);
}
ASTNodeType::Function { args, .. } => {
for arg in args {
tree_reference_keys(arg, out);
}
}
ASTNodeType::Array(rows) => {
for item in rows.iter().flatten() {
tree_reference_keys(item, out);
}
}
ASTNodeType::Call { callee, args } => {
tree_reference_keys(callee, out);
for arg in args {
tree_reference_keys(arg, out);
}
}
ASTNodeType::Literal(_) | ASTNodeType::Omitted => {}
}
}
pub(crate) fn semantic_reference_key(
reference: &crate::engine::refs::SemanticReference<'_>,
) -> Option<usize> {
use crate::engine::refs::SemanticReference;
match reference {
SemanticReference::Cell(cell) => Some(cell.original as *const _ as usize),
SemanticReference::FiniteRange(range) | SemanticReference::OpenRange(range) => {
Some(range.original as *const _ as usize)
}
SemanticReference::Name(name) => Some(name.as_ptr() as usize),
SemanticReference::Table(_)
| SemanticReference::ExternalSource(_)
| SemanticReference::ThreeDimensional(_)
| SemanticReference::Unsupported(_) => None,
}
}
pub(crate) struct VisitTrace {
index_by_key: FxHashMap<usize, u32>,
pub(crate) visit: Vec<u32>,
pub(crate) valid: bool,
}
impl VisitTrace {
pub(crate) fn new(keys: &[usize]) -> Option<Self> {
let mut index_by_key = FxHashMap::default();
for (index, key) in keys.iter().enumerate() {
if index_by_key.insert(*key, index as u32).is_some() {
return None;
}
}
Some(Self {
index_by_key,
visit: Vec::new(),
valid: true,
})
}
pub(crate) fn record(&mut self, key: Option<usize>) {
match key.and_then(|key| self.index_by_key.get(&key)) {
Some(index) => self.visit.push(*index),
None => self.valid = false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn lookup(memo: &mut ShapeMemo, tokens: &[u64]) -> ShapeLookup {
memo.tokens.clear();
memo.tokens.extend_from_slice(tokens);
memo.lookup_shape()
}
fn validity() -> MemoValidity {
MemoValidity {
provider_revision: Some(0),
registry_epoch: 1,
}
}
#[test]
fn colliding_shapes_are_bounded_per_lookup() {
let mut memo = ShapeMemo {
forced_hash: Some(7),
..ShapeMemo::default()
};
memo.revalidate(validity());
let shapes = 1_000u64;
let mut materialized = 0;
let mut bypassed = 0;
for round in 0..3 {
for shape in 0..shapes {
let before = memo.comparisons;
match lookup(&mut memo, &[T_INT, shape, T_CELL, 0, 0, 1, 1]) {
ShapeLookup::Shape(_, true) => materialized += 1,
ShapeLookup::Shape(_, false) => {
assert!((1..=MAX_BUCKET_SHAPES as u64).contains(&shape))
}
ShapeLookup::FirstSighting => assert_eq!((round, shape), (0, 0)),
ShapeLookup::Full => bypassed += 1,
}
assert!(memo.comparisons - before <= MAX_BUCKET_SHAPES as u64);
}
}
assert_eq!(materialized, MAX_BUCKET_SHAPES);
assert_eq!(memo.footprint().0, MAX_BUCKET_SHAPES);
assert_eq!(bypassed, 3 * shapes as usize - 1 - 3 * MAX_BUCKET_SHAPES);
assert!(memo.comparisons <= 3 * shapes * MAX_BUCKET_SHAPES as u64);
}
#[test]
fn memory_per_pipeline_is_bounded() {
let mut memo = ShapeMemo::default();
memo.revalidate(validity());
let distinct = (MAX_SEEN + MAX_SHAPES + 1_000) as u64;
for shape in 0..distinct {
for _ in 0..2 {
if let ShapeLookup::Shape(index, _) = lookup(&mut memo, &[T_INT, shape]) {
memo.insert_specialization(index, 0, None);
memo.insert_specialization(index, 1, None);
}
}
}
let (shapes, specializations, seen, stored) = memo.footprint();
assert_eq!(shapes, MAX_SHAPES);
assert!(specializations <= MAX_SPECIALIZATIONS, "{specializations}");
assert!(seen <= MAX_SEEN, "{seen}");
assert!(stored <= MAX_STORED_TOKENS, "{stored}");
assert!(matches!(
lookup(&mut memo, &[T_INT, distinct + 1]),
ShapeLookup::Full
));
let mut memo = ShapeMemo::default();
memo.revalidate(validity());
let mut full = 0;
for shape in 0..(MAX_SEEN + 100) as u64 {
if matches!(lookup(&mut memo, &[T_INT, shape]), ShapeLookup::Full) {
full += 1;
}
}
assert_eq!(memo.footprint().2, MAX_SEEN);
assert_eq!(full, 100);
let mut memo = ShapeMemo::default();
memo.revalidate(validity());
let long = MAX_SHAPE_TOKENS;
let mut tokens = vec![T_EMPTY; long];
for shape in 0..(2 * MAX_STORED_TOKENS / long) as u64 {
tokens[0] = shape;
for _ in 0..2 {
if let ShapeLookup::Shape(index, true) = lookup(&mut memo, &tokens) {
memo.insert_specialization(index, 0, None);
}
}
}
let (shapes, _, _, stored) = memo.footprint();
assert!(stored <= MAX_STORED_TOKENS, "{stored}");
assert!(shapes < MAX_STORED_TOKENS / long, "{shapes}");
let over = vec![T_EMPTY; MAX_SHAPE_TOKENS + 1];
assert!(matches!(lookup(&mut memo, &over), ShapeLookup::Full));
memo.revalidate(MemoValidity {
registry_epoch: 2,
..validity()
});
assert_eq!(memo.footprint(), (0, 0, 0, 0));
}
#[test]
fn probe_stride_backs_off_and_resets_on_hit() {
let mut memo = ShapeMemo::default();
let probes: Vec<bool> = (0..10_000).map(|_| memo.should_probe()).collect();
let warm = PROBE_WARMUP as usize;
assert!(probes[..warm].iter().all(|&p| p));
let mut window = warm;
let mut stride = 2;
while stride <= MAX_PROBE_STRIDE as usize {
let end = if stride == MAX_PROBE_STRIDE as usize {
probes.len()
} else {
window + warm
};
let probed = probes[window..end].iter().filter(|&&p| p).count();
assert_eq!(probed, (end - window) / stride, "stride {stride}");
let gaps = probes[window..end]
.split(|&p| p)
.map(<[bool]>::len)
.max()
.unwrap();
assert!(gaps < stride, "stride {stride}: gap {gaps}");
window = end;
stride *= 2;
}
memo.record_hit();
assert!((0..PROBE_WARMUP).all(|_| memo.should_probe()));
}
fn specialization_with_key(bytes: usize) -> Specialization {
let key: Arc<str> = Arc::from("k".repeat(bytes));
Specialization {
canonical_hash: 0,
exact_canonical_hash: 0,
exact_canonical_key: key.clone(),
parameterized_canonical_hash: 0,
parameterized_canonical_key: key,
literal_slot_descriptors: Arc::from(Vec::new()),
literal_bindings: vec![LiteralValue::Text("t".repeat(bytes))].into_boxed_slice(),
value_ref_slot_descriptors: Arc::from(Vec::new()),
expr: CanonicalExpr::Omitted,
labels: CanonicalLabels::default(),
read_projections: None,
read_projection_fallback: None,
volatile: false,
dynamic: false,
visit: Box::new([]),
}
}
#[test]
fn specialization_bytes_are_bounded() {
let mut memo = ShapeMemo::default();
memo.revalidate(validity());
let mut shapes = Vec::new();
for shape in 0..2_000u64 {
for _ in 0..2 {
if let ShapeLookup::Shape(index, true) = lookup(&mut memo, &[T_INT, shape]) {
shapes.push(index);
}
}
}
memo.insert_specialization(
shapes[0],
0,
Some(specialization_with_key(MAX_SPECIALIZATION_BYTES)),
);
assert_eq!(memo.retained_specializations(), 0);
assert_eq!(memo.stored_bytes(), 0);
for &shape in &shapes[1..] {
memo.insert_specialization(shape, 0, Some(specialization_with_key(40 << 10)));
}
let retained = memo.retained_specializations();
assert!(retained > 0 && retained < shapes.len() - 1, "{retained}");
assert!(
memo.stored_bytes() <= MAX_STORED_BYTES,
"{}",
memo.stored_bytes()
);
assert!(memo.stored_bytes() > MAX_STORED_BYTES - MAX_SPECIALIZATION_BYTES);
}
}