use core::fmt;
use std::collections::{BTreeMap, BTreeSet};
use mongreldb_types::ids::{MetadataVersion, SchemaVersion, TableId};
#[repr(transparent)]
#[derive(
Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
pub struct StatementId(pub u64);
impl StatementId {
pub const ZERO: Self = Self(0);
pub const fn new(value: u64) -> Self {
Self(value)
}
pub const fn get(self) -> u64 {
self.0
}
}
impl fmt::Display for StatementId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
impl fmt::Debug for StatementId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "StatementId({})", self.0)
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PreparedStatementBinding {
pub statement_id: StatementId,
pub sql: String,
pub parameter_types: Vec<String>,
pub catalog_version: MetadataVersion,
pub schema_versions: BTreeMap<TableId, SchemaVersion>,
pub feature_set: BTreeSet<String>,
}
impl PreparedStatementBinding {
pub fn is_compatible(
&self,
catalog_version: MetadataVersion,
schema_versions: &BTreeMap<TableId, SchemaVersion>,
) -> bool {
self.catalog_version == catalog_version
&& self
.schema_versions
.iter()
.all(|(table, version)| schema_versions.get(table) == Some(version))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::assert_serde_round_trip;
fn binding() -> PreparedStatementBinding {
let mut schema_versions = BTreeMap::new();
schema_versions.insert(TableId::new(1), SchemaVersion::new(10));
schema_versions.insert(TableId::new(2), SchemaVersion::new(20));
let mut feature_set = BTreeSet::new();
feature_set.insert("ann-index".to_owned());
feature_set.insert("cdc".to_owned());
PreparedStatementBinding {
statement_id: StatementId::new(7),
sql: "SELECT * FROM events WHERE tenant = ?".to_owned(),
parameter_types: vec!["INT64".to_owned()],
catalog_version: MetadataVersion::new(100),
schema_versions,
feature_set,
}
}
#[test]
fn statement_id_basics_and_serde() {
assert_eq!(StatementId::ZERO.get(), 0);
let id = StatementId::new(42);
assert_eq!(id.get(), 42);
assert_eq!(id.to_string(), "42");
assert_eq!(format!("{id:?}"), "StatementId(42)");
assert_serde_round_trip(&id);
assert_serde_round_trip(&StatementId::ZERO);
}
#[test]
fn binding_serde_round_trip() {
assert_serde_round_trip(&binding());
let empty = PreparedStatementBinding {
statement_id: StatementId::ZERO,
sql: String::new(),
parameter_types: vec![],
catalog_version: MetadataVersion::ZERO,
schema_versions: BTreeMap::new(),
feature_set: BTreeSet::new(),
};
assert_serde_round_trip(&empty);
}
#[test]
fn invalidation_matrix() {
let binding = binding();
let current = binding.schema_versions.clone();
assert!(binding.is_compatible(MetadataVersion::new(100), ¤t));
assert!(!binding.is_compatible(MetadataVersion::new(101), ¤t));
let mut altered = current.clone();
altered.insert(TableId::new(1), SchemaVersion::new(11));
assert!(!binding.is_compatible(MetadataVersion::new(100), &altered));
let mut dropped = current.clone();
dropped.remove(&TableId::new(2));
assert!(!binding.is_compatible(MetadataVersion::new(100), &dropped));
let mut unrelated = current.clone();
unrelated.insert(TableId::new(3), SchemaVersion::new(1));
assert!(binding.is_compatible(MetadataVersion::new(100), &unrelated));
let no_tables = PreparedStatementBinding {
schema_versions: BTreeMap::new(),
..binding.clone()
};
assert!(no_tables.is_compatible(MetadataVersion::new(100), &BTreeMap::new()));
assert!(!no_tables.is_compatible(MetadataVersion::new(99), &BTreeMap::new()));
}
}