use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MesiState {
Modified,
Exclusive,
Shared,
Invalid,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConsistencyStrategy {
WriteThrough,
WriteBehind,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum InvalidationOp {
Modify,
Delete,
}
#[derive(Debug, Clone)]
pub struct InvalidationEvent {
pub key: String,
pub instance_id: String,
pub timestamp: u64,
pub op: InvalidationOp,
}
pub trait InvalidationBroadcaster: Send + Sync {
fn broadcast(&self, event: &InvalidationEvent) -> Result<(), CoherenceError>;
}
#[derive(Debug, Clone, Default)]
pub struct CoherenceMetrics {
pub modified_count: u64,
pub exclusive_count: u64,
pub shared_count: u64,
pub invalid_count: u64,
pub invalidation_broadcasts: u64,
pub coherence_violations: u64,
pub write_behind_rollbacks: u64,
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum CoherenceError {
#[error("broadcast failed: {0}")]
BroadcastFailed(String),
#[error("write-behind failed for key: {key}")]
WriteBehindFailed {
key: String,
},
#[error("split-brain detected for key: {key}")]
SplitBrain {
key: String,
},
#[error("cache miss for key: {0}")]
CacheMiss(String),
}
pub struct CacheCoherenceProtocol {
states: RwLock<HashMap<String, MesiState>>,
broadcaster: Arc<dyn InvalidationBroadcaster>,
instance_id: String,
strategy: ConsistencyStrategy,
metrics: Arc<RwLock<CoherenceMetrics>>,
}
impl CacheCoherenceProtocol {
pub fn new(
instance_id: String,
strategy: ConsistencyStrategy,
broadcaster: Arc<dyn InvalidationBroadcaster>,
) -> Self {
Self {
states: RwLock::new(HashMap::new()),
broadcaster,
instance_id,
strategy,
metrics: Arc::new(RwLock::new(CoherenceMetrics::default())),
}
}
pub fn state(&self, key: &str) -> MesiState {
self.states
.read()
.unwrap()
.get(key)
.copied()
.unwrap_or(MesiState::Invalid)
}
pub fn read(&self, key: &str, other_instances_have: bool) -> MesiState {
let mut states = self.states.write().unwrap();
let mut metrics = self.metrics.write().unwrap();
let current = states.get(key).copied().unwrap_or(MesiState::Invalid);
let new_state = match current {
MesiState::Invalid => {
if other_instances_have {
MesiState::Shared
} else {
MesiState::Exclusive
}
}
other => other,
};
states.insert(key.to_string(), new_state);
Self::update_metrics(&mut metrics, &new_state);
new_state
}
pub fn write(&self, key: &str) -> Result<MesiState, CoherenceError> {
let event = InvalidationEvent {
key: key.to_string(),
instance_id: self.instance_id.clone(),
timestamp: SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64,
op: InvalidationOp::Modify,
};
self.broadcaster.broadcast(&event)?;
let mut states = self.states.write().unwrap();
let mut metrics = self.metrics.write().unwrap();
states.insert(key.to_string(), MesiState::Modified);
metrics.invalidation_broadcasts += 1;
Self::update_metrics(&mut metrics, &MesiState::Modified);
Ok(MesiState::Modified)
}
pub fn handle_invalidation(&self, event: &InvalidationEvent) {
if event.instance_id == self.instance_id {
return;
}
let mut states = self.states.write().unwrap();
let mut metrics = self.metrics.write().unwrap();
states.insert(event.key.clone(), MesiState::Invalid);
metrics.invalid_count += 1;
}
pub fn metrics(&self) -> CoherenceMetrics {
self.metrics.read().unwrap().clone()
}
pub fn strategy(&self) -> ConsistencyStrategy {
self.strategy
}
fn update_metrics(metrics: &mut CoherenceMetrics, state: &MesiState) {
match state {
MesiState::Modified => metrics.modified_count += 1,
MesiState::Exclusive => metrics.exclusive_count += 1,
MesiState::Shared => metrics.shared_count += 1,
MesiState::Invalid => metrics.invalid_count += 1,
}
}
}
pub struct NoopBroadcaster;
impl InvalidationBroadcaster for NoopBroadcaster {
fn broadcast(&self, _event: &InvalidationEvent) -> Result<(), CoherenceError> {
Ok(())
}
}
pub struct LocalBroadcaster {
events: RwLock<Vec<InvalidationEvent>>,
}
impl LocalBroadcaster {
pub fn new() -> Self {
Self {
events: RwLock::new(Vec::new()),
}
}
pub fn events(&self) -> Vec<InvalidationEvent> {
self.events.read().unwrap().clone()
}
}
impl Default for LocalBroadcaster {
fn default() -> Self {
Self::new()
}
}
impl InvalidationBroadcaster for LocalBroadcaster {
fn broadcast(&self, event: &InvalidationEvent) -> Result<(), CoherenceError> {
self.events.write().unwrap().push(event.clone());
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mesi_state_transitions() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
assert_eq!(protocol.state("key1"), MesiState::Invalid);
let s = protocol.read("key1", false);
assert_eq!(s, MesiState::Exclusive);
let s = protocol.read("key1", true);
assert_eq!(s, MesiState::Exclusive);
let s = protocol.write("key1").unwrap();
assert_eq!(s, MesiState::Modified);
assert_eq!(protocol.state("key1"), MesiState::Modified);
}
#[test]
fn test_invalid_to_shared() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
let s = protocol.read("key1", true);
assert_eq!(s, MesiState::Shared);
}
#[test]
fn test_invalid_to_exclusive() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
let s = protocol.read("key1", false);
assert_eq!(s, MesiState::Exclusive);
}
#[test]
fn test_write_broadcasts_invalidation() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster.clone(),
);
protocol.write("key1").unwrap();
let events = broadcaster.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].key, "key1");
assert_eq!(events[0].op, InvalidationOp::Modify);
}
#[test]
fn test_handle_invalidation_sets_invalid() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
protocol.read("key1", false);
assert_eq!(protocol.state("key1"), MesiState::Exclusive);
let event = InvalidationEvent {
key: "key1".to_string(),
instance_id: "instance-B".to_string(),
timestamp: 0,
op: InvalidationOp::Modify,
};
protocol.handle_invalidation(&event);
assert_eq!(protocol.state("key1"), MesiState::Invalid);
}
#[test]
fn test_ignore_self_invalidation() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
protocol.read("key1", false);
assert_eq!(protocol.state("key1"), MesiState::Exclusive);
let event = InvalidationEvent {
key: "key1".to_string(),
instance_id: "instance-A".to_string(),
timestamp: 0,
op: InvalidationOp::Modify,
};
protocol.handle_invalidation(&event);
assert_eq!(protocol.state("key1"), MesiState::Exclusive);
}
#[test]
fn test_metrics_tracking() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
protocol.read("key1", false);
protocol.read("key2", true);
protocol.write("key1").unwrap();
let metrics = protocol.metrics();
assert!(metrics.exclusive_count > 0);
assert!(metrics.shared_count > 0);
assert!(metrics.modified_count > 0);
assert!(metrics.invalidation_broadcasts > 0);
}
#[test]
fn test_noop_broadcaster() {
let broadcaster = Arc::new(NoopBroadcaster);
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteBehind,
broadcaster,
);
let result = protocol.write("key1");
assert!(result.is_ok());
}
#[test]
fn test_shared_to_modified_on_write() {
let broadcaster = Arc::new(LocalBroadcaster::new());
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteThrough,
broadcaster,
);
let s = protocol.read("key1", true);
assert_eq!(s, MesiState::Shared);
let s = protocol.write("key1").unwrap();
assert_eq!(s, MesiState::Modified);
}
#[test]
fn test_strategy_access() {
let broadcaster = Arc::new(NoopBroadcaster);
let protocol = CacheCoherenceProtocol::new(
"instance-A".to_string(),
ConsistencyStrategy::WriteBehind,
broadcaster,
);
assert_eq!(protocol.strategy(), ConsistencyStrategy::WriteBehind);
}
}