use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PredicateInfo {
pub factorization: Factorization,
pub monotonicity: PerAxisMap<Monotonicity>,
pub range_constraint: PerAxisMap<RangeConstraint>,
pub determinism: Determinism,
pub coords_referenced: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum Factorization {
PerAxis(PerAxisMap<String>),
Conjunctive(Vec<String>),
Disjunctive(Vec<String>),
Opaque(OpaqueReason),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum OpaqueReason {
UnknownPattern,
NonDeterministic,
CrossTupleState,
SideEffecting,
Continuous,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Monotonicity {
Increasing,
Decreasing,
None,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum RangeConstraint {
Bounded {
lo: Option<ConstValue>,
hi: Option<ConstValue>,
lo_inclusive: bool,
hi_inclusive: bool,
},
Discrete(Vec<ConstValue>),
None,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Determinism {
Deterministic,
Opaque,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ConstValue {
Int(i64),
Float(f64),
String(String),
Bool(bool),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PerAxisMap<T> {
entries: Vec<(String, T)>,
}
impl<T> PerAxisMap<T> {
pub fn new() -> Self {
Self { entries: Vec::new() }
}
pub fn insert<K: Into<String>>(&mut self, key: K, value: T) {
let key = key.into();
if let Some(slot) = self.entries.iter_mut().find(|(k, _)| *k == key) {
slot.1 = value;
} else {
self.entries.push((key, value));
}
}
pub fn get(&self, key: &str) -> Option<&T> {
self.entries.iter().find(|(k, _)| k == key).map(|(_, v)| v)
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &T)> {
self.entries.iter().map(|(k, v)| (k.as_str(), v))
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn keys(&self) -> impl Iterator<Item = &str> {
self.entries.iter().map(|(k, _)| k.as_str())
}
}
impl<T> Default for PerAxisMap<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> FromIterator<(String, T)> for PerAxisMap<T> {
fn from_iter<I: IntoIterator<Item = (String, T)>>(iter: I) -> Self {
let mut m = Self::new();
for (k, v) in iter {
m.insert(k, v);
}
m
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn per_axis_map_insertion_order() {
let mut m: PerAxisMap<i64> = PerAxisMap::new();
m.insert("k", 10);
m.insert("limit", 20);
let keys: Vec<&str> = m.keys().collect();
assert_eq!(keys, vec!["k", "limit"]);
}
#[test]
fn per_axis_map_replace_preserves_position() {
let mut m: PerAxisMap<i64> = PerAxisMap::new();
m.insert("k", 10);
m.insert("limit", 20);
m.insert("k", 100); let keys: Vec<&str> = m.keys().collect();
assert_eq!(keys, vec!["k", "limit"]);
assert_eq!(*m.get("k").unwrap(), 100);
}
#[test]
fn predicate_info_round_trip_serde() {
let info = PredicateInfo {
factorization: Factorization::PerAxis(
vec![
("k".to_string(), "{k} > 0".to_string()),
("limit".to_string(), "{limit} < 100".to_string()),
]
.into_iter()
.collect(),
),
monotonicity: PerAxisMap::new(),
range_constraint: PerAxisMap::new(),
determinism: Determinism::Deterministic,
coords_referenced: vec!["k".to_string(), "limit".to_string()],
};
let json = serde_json::to_string(&info).unwrap();
let back: PredicateInfo = serde_json::from_str(&json).unwrap();
assert_eq!(info, back);
}
#[test]
fn opaque_reason_serde() {
let info = PredicateInfo {
factorization: Factorization::Opaque(OpaqueReason::Continuous),
monotonicity: PerAxisMap::new(),
range_constraint: PerAxisMap::new(),
determinism: Determinism::Deterministic,
coords_referenced: vec!["theta".to_string()],
};
let json = serde_json::to_string(&info).unwrap();
assert!(json.contains("continuous"));
let back: PredicateInfo = serde_json::from_str(&json).unwrap();
assert_eq!(info, back);
}
}