use super::*;
use crate::engine::arena::{AstNodeData, CompactRefType, DataStore};
use crate::format::FormatId;
use crate::function::FamilyKernel;
pub(super) struct MemoPlan {
pub(super) key_args: smallvec::SmallVec<[AstNodeId; 8]>,
}
impl MemoPlan {
pub(super) fn plan(
functions: &dyn crate::traits::FunctionProvider,
ds: &DataStore,
template: AstNodeId,
) -> Option<Self> {
let AstNodeData::Function { name_id, .. } = ds.get_node(template)? else {
return Self::plan_general(functions, ds, template);
};
let fun = functions.get_function("", ds.resolve_ast_string(*name_id))?;
if !matches!(
fun.family_kernel(),
Some(FamilyKernel::CriteriaAggregate | FamilyKernel::Lookup)
) {
return Self::plan_general(functions, ds, template);
}
let mut key_args = smallvec::SmallVec::new();
for &arg in ds.get_args(template)? {
match ds.get_node(arg)? {
AstNodeData::Reference { ref_type, .. } => match ref_type {
CompactRefType::Cell {
row_abs, col_abs, ..
} if *row_abs && *col_abs => {}
CompactRefType::Range {
start_row,
start_col,
end_row,
end_col,
start_row_abs,
start_col_abs,
end_row_abs,
end_col_abs,
..
} if (*start_row_abs || *start_row == 0)
&& (*start_col_abs || *start_col == 0)
&& (*end_row_abs || *end_row == u32::MAX)
&& (*end_col_abs || *end_col == u32::MAX) => {}
CompactRefType::Cell { .. } => key_args.push(arg),
_ => return None,
},
_ => key_args.push(arg),
}
}
Some(Self { key_args })
}
fn plan_general(
functions: &dyn crate::traits::FunctionProvider,
ds: &DataStore,
template: AstNodeId,
) -> Option<Self> {
fn visit(
functions: &dyn crate::traits::FunctionProvider,
ds: &DataStore,
id: AstNodeId,
as_arg: bool,
keys: &mut smallvec::SmallVec<[AstNodeId; 8]>,
worth: &mut bool,
) -> Option<()> {
match ds.get_node(id)? {
AstNodeData::Literal(vref) => {
if matches!(ds.retrieve_value(*vref), LiteralValue::Array(_)) {
return None;
}
}
AstNodeData::Reference { ref_type, .. } => match ref_type {
CompactRefType::Cell {
row_abs: true, row, ..
} if *row > 0 => {}
CompactRefType::Cell { row, col, .. } if *row > 0 && *col > 0 => {
if keys.len() >= 8 {
return None;
}
keys.push(id);
}
CompactRefType::Range {
start_row,
end_row,
start_row_abs,
end_row_abs,
..
} if as_arg
&& (*start_row_abs || *start_row == 0)
&& (*end_row_abs || *end_row == u32::MAX) => {}
_ => return None,
},
AstNodeData::UnaryOp { expr_id, .. } => {
visit(functions, ds, *expr_id, false, keys, worth)?
}
AstNodeData::BinaryOp {
left_id, right_id, ..
} => {
visit(functions, ds, *left_id, false, keys, worth)?;
visit(functions, ds, *right_id, false, keys, worth)?;
}
AstNodeData::Function { name_id, .. } => {
let name = ds.resolve_ast_string(*name_id);
if !super::lift::pure_listed_function(functions, name) {
return None;
}
*worth |= MEMO_WORTH.iter().any(|f| f.eq_ignore_ascii_case(name));
for &arg in ds.get_args(id)? {
visit(functions, ds, arg, true, keys, worth)?;
}
}
_ => return None,
}
Some(())
}
let mut key_args = smallvec::SmallVec::new();
let mut worth = false;
visit(functions, ds, template, false, &mut key_args, &mut worth)?;
(worth && !key_args.is_empty()).then_some(Self { key_args })
}
}
const MEMO_WORTH: &[&str] = &[
"INDEX",
"MATCH",
"VLOOKUP",
"HLOOKUP",
"XLOOKUP",
"XMATCH",
"SUMIF",
"SUMIFS",
"COUNTIF",
"COUNTIFS",
"AVERAGEIF",
"AVERAGEIFS",
"SUMPRODUCT",
];
#[derive(Clone, PartialEq, Eq, Hash)]
pub(super) enum KeyValue {
Number(u64),
Int(i64),
Text(String),
Boolean(bool),
Empty,
Date(chrono::NaiveDate),
DateTime(chrono::NaiveDateTime),
Time(chrono::NaiveTime),
Duration(i64, i32),
}
impl KeyValue {
pub(super) fn of(value: LiteralValue) -> Option<Self> {
Some(match value {
LiteralValue::Number(n) => KeyValue::Number(n.to_bits()),
LiteralValue::Int(i) => KeyValue::Int(i),
LiteralValue::Text(s) => KeyValue::Text(s),
LiteralValue::Boolean(b) => KeyValue::Boolean(b),
LiteralValue::Empty => KeyValue::Empty,
LiteralValue::Date(d) => KeyValue::Date(d),
LiteralValue::DateTime(d) => KeyValue::DateTime(d),
LiteralValue::Time(t) => KeyValue::Time(t),
LiteralValue::Duration(d) => KeyValue::Duration(d.num_seconds(), d.subsec_nanos()),
_ => return None,
})
}
}
pub(super) type MemoKey = smallvec::SmallVec<[(KeyValue, Option<FormatId>); 4]>;
type MemoMap = rustc_hash::FxHashMap<MemoKey, (LiteralValue, Option<FormatId>)>;
#[derive(Default)]
pub(super) struct SharedMemo {
map: std::sync::Mutex<MemoMap>,
seen: std::sync::atomic::AtomicUsize,
pub(super) criteria: std::sync::OnceLock<Option<super::criteria::CriteriaIndex>>,
pub(super) invariant: std::sync::OnceLock<crate::interpreter::InvariantValues>,
pub(super) run_len: u32,
}
const MEMO_WARMUP: usize = 256;
pub(super) struct RunMemo<'s> {
local: MemoMap,
shared: Option<&'s SharedMemo>,
seen: usize,
off: bool,
}
impl<'s> RunMemo<'s> {
pub(super) fn new(shared: Option<&'s SharedMemo>) -> Self {
Self {
local: MemoMap::default(),
shared,
seen: 0,
off: false,
}
}
fn counts(&self) -> (usize, usize) {
use std::sync::atomic::Ordering::Relaxed;
match self.shared {
Some(s) => (
s.seen.load(Relaxed),
s.map.lock().map(|m| m.len()).unwrap_or(usize::MAX),
),
None => (self.seen, self.local.len()),
}
}
pub(super) fn active(&mut self) -> bool {
if !self.off && self.seen.is_multiple_of(16) {
let (seen, distinct) = self.counts();
if seen >= MEMO_WARMUP && distinct.saturating_mul(2) > seen {
self.off = true;
self.local = MemoMap::default();
}
}
!self.off
}
pub(super) fn get(&mut self, key: &MemoKey) -> Option<(LiteralValue, Option<FormatId>)> {
use std::sync::atomic::Ordering::Relaxed;
match self.shared {
Some(shared) => {
self.seen += 1;
shared.seen.fetch_add(1, Relaxed);
shared.map.lock().ok()?.get(key).cloned()
}
None => {
self.seen += 1;
self.local.get(key).cloned()
}
}
}
pub(super) fn insert(&mut self, key: MemoKey, value: (LiteralValue, Option<FormatId>)) {
match self.shared {
Some(shared) => {
if let Ok(mut map) = shared.map.lock() {
map.insert(key, value);
}
}
None => {
self.local.insert(key, value);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn key(i: usize) -> MemoKey {
smallvec::smallvec![(KeyValue::Number((i as f64).to_bits()), None)]
}
fn drive(memo: &mut RunMemo<'_>, keys: impl Iterator<Item = usize>) -> (usize, usize, bool) {
let (mut evaluated, mut hits) = (0, 0);
for k in keys {
if !memo.active() {
evaluated += 1;
continue;
}
match memo.get(&key(k)) {
Some(_) => hits += 1,
None => {
evaluated += 1;
memo.insert(key(k), (LiteralValue::Number(k as f64), None));
}
}
}
(evaluated, hits, memo.active())
}
#[test]
fn cycling_keys_keep_the_memo_on() {
for shared in [None, Some(SharedMemo::default())] {
let mut memo = RunMemo::new(shared.as_ref());
let (evaluated, hits, on) = drive(&mut memo, (0..10_000).map(|i| i % 60));
assert!(on);
assert_eq!(evaluated, 60);
assert_eq!(hits, 10_000 - 60);
}
}
#[test]
fn distinct_keys_turn_the_memo_off_after_the_warm_up() {
for shared in [None, Some(SharedMemo::default())] {
let mut memo = RunMemo::new(shared.as_ref());
let (evaluated, hits, on) = drive(&mut memo, 0..10_000);
assert!(!on);
assert_eq!((evaluated, hits), (10_000, 0));
assert!(memo.seen <= MEMO_WARMUP + 16);
}
}
}