use crate::datatypes::Value;
use crate::graph::core::filtering::{json_single_element_string, values_equal};
use chrono::{NaiveDate, NaiveDateTime};
use rustc_hash::FxHashSet;
use std::sync::Arc;
const LINEAR_MAX: usize = 8;
const EXACT_INT_LIMIT: i64 = 1i64 << 53;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
enum ScalarKey {
Int(i64),
Float(u64),
Bool(bool),
Date(NaiveDate),
Timestamp(NaiveDateTime),
}
#[derive(Debug, Clone, Default)]
struct MembershipIndex {
scalars: FxHashSet<ScalarKey>,
strings: FxHashSet<Box<str>>,
residual: Vec<Value>,
}
#[derive(Debug, Clone, Default)]
pub struct MembershipSet {
values: Vec<Value>,
has_null: bool,
index: Option<Arc<MembershipIndex>>,
}
impl MembershipSet {
pub fn new(values: Vec<Value>) -> Self {
let has_null = values.iter().any(|v| matches!(v, Value::Null));
let index = (values.len() > LINEAR_MAX).then(|| Arc::new(build_index(&values)));
Self {
values,
has_null,
index,
}
}
#[inline]
pub fn matches(&self, value: &Value) -> bool {
match &self.index {
Some(index) => index.matches(value),
None => self.values.iter().any(|v| values_equal(value, v)),
}
}
#[inline]
pub fn has_null(&self) -> bool {
self.has_null
}
#[inline]
pub fn values(&self) -> &[Value] {
&self.values
}
pub fn into_values(self) -> Vec<Value> {
self.values
}
#[inline]
pub fn kleene_contains(&self, value: &Value) -> Option<bool> {
if matches!(value, Value::Null) {
return None;
}
if self.matches(value) {
return Some(true);
}
if self.has_null {
return None;
}
Some(false)
}
}
#[inline]
pub fn kleene_contains_linear(value: &Value, items: &[Value]) -> Option<bool> {
if matches!(value, Value::Null) {
return None;
}
let mut saw_null = false;
for item in items {
match probe_element(value, item) {
Some(true) => return Some(true),
Some(false) => {}
None => saw_null = true,
}
}
if saw_null {
None
} else {
Some(false)
}
}
#[inline]
pub fn probe_element(value: &Value, element: &Value) -> Option<bool> {
if matches!(element, Value::Null) {
return None;
}
Some(values_equal(value, element))
}
impl std::ops::Deref for MembershipSet {
type Target = [Value];
fn deref(&self) -> &Self::Target {
&self.values
}
}
impl<'a> IntoIterator for &'a MembershipSet {
type Item = &'a Value;
type IntoIter = std::slice::Iter<'a, Value>;
fn into_iter(self) -> Self::IntoIter {
self.values.iter()
}
}
impl From<Vec<Value>> for MembershipSet {
fn from(values: Vec<Value>) -> Self {
Self::new(values)
}
}
impl FromIterator<Value> for MembershipSet {
fn from_iter<I: IntoIterator<Item = Value>>(iter: I) -> Self {
Self::new(iter.into_iter().collect())
}
}
impl MembershipIndex {
#[inline]
fn matches(&self, value: &Value) -> bool {
if let Value::String(s) = value {
if self.strings.contains(s.as_str())
|| json_single_element_string(s).is_some_and(|inner| self.strings.contains(inner))
{
return true;
}
} else if let Some(key) = scalar_key(value) {
if self.scalars.contains(&key) {
return true;
}
}
!self.residual.is_empty() && self.residual.iter().any(|v| values_equal(value, v))
}
}
fn build_index(values: &[Value]) -> MembershipIndex {
let mut index = MembershipIndex {
scalars: FxHashSet::with_capacity_and_hasher(values.len(), Default::default()),
strings: FxHashSet::default(),
residual: Vec::new(),
};
for value in values {
match value {
Value::Null => {}
Value::Float64(f) if f.is_nan() => {}
Value::String(s) => {
index.strings.insert(s.as_str().into());
if let Some(inner) = json_single_element_string(s) {
index.strings.insert(inner.into());
}
}
other => match scalar_key(other) {
Some(key) => {
index.scalars.insert(key);
if beyond_exact_int_range(other) {
index.residual.push(other.clone());
}
}
None => index.residual.push(other.clone()),
},
}
}
index
}
#[inline]
fn scalar_key(value: &Value) -> Option<ScalarKey> {
match value {
Value::Int64(i) => Some(ScalarKey::Int(*i)),
Value::UniqueId(u) => Some(ScalarKey::Int(*u as i64)),
Value::Float64(f) => {
if f.is_nan() {
None
} else if f.fract() == 0.0 && f.abs() < EXACT_INT_LIMIT as f64 {
Some(ScalarKey::Int(*f as i64))
} else {
let canonical = if *f == 0.0 { 0.0f64 } else { *f };
Some(ScalarKey::Float(canonical.to_bits()))
}
}
Value::Boolean(b) => Some(ScalarKey::Bool(*b)),
Value::DateTime(d) => Some(ScalarKey::Date(*d)),
Value::Timestamp(t) => Some(ScalarKey::Timestamp(*t)),
_ => None,
}
}
#[inline]
fn beyond_exact_int_range(value: &Value) -> bool {
match value {
Value::Int64(i) => i.unsigned_abs() >= EXACT_INT_LIMIT as u64,
Value::Float64(f) => f.abs() >= EXACT_INT_LIMIT as f64,
_ => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_agrees(list: &[Value], probes: &[Value]) {
let set = MembershipSet::new(list.to_vec());
for probe in probes {
let linear = list.iter().any(|v| values_equal(probe, v));
assert_eq!(
set.matches(probe),
linear,
"membership disagreed with values_equal for {probe:?} in {list:?}"
);
}
}
fn padded(list: &[Value]) -> Vec<Value> {
let mut out = list.to_vec();
out.extend((0..LINEAR_MAX + 2).map(|i| Value::String(format!("__pad_{i}"))));
out
}
#[test]
fn numeric_family_coerces_like_values_equal() {
let list = [
Value::Int64(5),
Value::Float64(7.5),
Value::UniqueId(9),
Value::Float64(11.0),
];
let probes = [
Value::Int64(5),
Value::Float64(5.0),
Value::UniqueId(5),
Value::Float64(5.5),
Value::Int64(7),
Value::Float64(7.5),
Value::Int64(9),
Value::Float64(9.0),
Value::Int64(11),
Value::UniqueId(11),
Value::Int64(-5),
Value::Float64(-0.0),
Value::Int64(0),
];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
}
#[test]
fn nan_never_matches_on_either_side() {
let list = [Value::Float64(f64::NAN), Value::Int64(1)];
let probes = [Value::Float64(f64::NAN), Value::Int64(1)];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
assert!(!MembershipSet::new(padded(&list)).matches(&Value::Float64(f64::NAN)));
}
#[test]
fn signed_zero_shares_a_key() {
let list = [Value::Float64(-0.0)];
let probes = [Value::Float64(0.0), Value::Int64(0), Value::UniqueId(0)];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
}
#[test]
fn json_single_element_strings_match_their_inner_value() {
let list = [
Value::String("[\"Oslo\"]".to_string()),
Value::String("Bergen".to_string()),
];
let probes = [
Value::String("Oslo".to_string()),
Value::String("[\"Oslo\"]".to_string()),
Value::String("[\"Bergen\"]".to_string()),
Value::String("Bergen".to_string()),
Value::String("Tromso".to_string()),
Value::String("[\"]".to_string()),
Value::String("[\"\"]".to_string()),
];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
}
#[test]
fn huge_integers_fall_back_to_values_equal() {
let big = EXACT_INT_LIMIT + 1;
let list = [Value::Int64(big), Value::Float64(EXACT_INT_LIMIT as f64)];
let probes = [
Value::Int64(big),
Value::Float64(big as f64),
Value::Int64(EXACT_INT_LIMIT),
Value::Float64(EXACT_INT_LIMIT as f64),
];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
}
#[test]
fn null_is_reported_not_matched() {
let set = MembershipSet::new(padded(&[Value::Null, Value::Int64(1)]));
assert!(set.has_null());
assert!(!set.matches(&Value::Null));
assert!(set.matches(&Value::Int64(1)));
assert!(!MembershipSet::new(vec![Value::Int64(1)]).has_null());
}
#[test]
fn non_scalar_values_compare_structurally() {
let list = [
Value::List(vec![Value::Int64(1), Value::Int64(2)]),
Value::Point { lat: 1.0, lon: 2.0 },
];
let probes = [
Value::List(vec![Value::Int64(1), Value::Int64(2)]),
Value::List(vec![Value::Int64(1)]),
Value::Point { lat: 1.0, lon: 2.0 },
Value::Point { lat: 9.0, lon: 2.0 },
];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
}
#[test]
fn cross_type_probes_stay_disjoint() {
let list = [
Value::Boolean(true),
Value::Int64(1),
Value::String("1".to_string()),
];
let probes = [
Value::Boolean(true),
Value::Boolean(false),
Value::Int64(1),
Value::String("1".to_string()),
Value::String("true".to_string()),
];
assert_agrees(&list, &probes);
assert_agrees(&padded(&list), &probes);
}
#[test]
fn threshold_crossing_preserves_answers() {
let mut list = Vec::new();
for i in 0..(LINEAR_MAX * 3) as i64 {
list.push(Value::Int64(i * 2));
let set = MembershipSet::new(list.clone());
for probe in 0..(LINEAR_MAX * 6) as i64 {
let expected = list.iter().any(|v| values_equal(&Value::Int64(probe), v));
assert_eq!(set.matches(&Value::Int64(probe)), expected, "n={i}");
}
}
}
}