use crate::sync::version_vector::{ConflictInfo, VectorComparison, VersionVector};
use serde_json::Value;
use std::sync::Arc;
#[async_trait::async_trait]
pub trait ConflictResolver: Send + Sync {
async fn resolve(&self, conflict: &ConflictInfo) -> ConflictResolution;
fn name(&self) -> &'static str;
}
#[derive(Debug, Clone)]
pub enum ConflictResolution {
LocalWins,
RemoteWins,
Merged(Value),
KeepBoth { local: Value, remote: Value },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConflictResolutionStrategy {
LastWriteWins,
Deterministic,
AutomaticMerge,
Manual,
CustomScript,
}
impl ConflictResolutionStrategy {
pub fn create_resolver(&self) -> Arc<dyn ConflictResolver> {
match self {
ConflictResolutionStrategy::LastWriteWins => Arc::new(LastWriteWinsResolver),
ConflictResolutionStrategy::Deterministic => Arc::new(DeterministicResolver),
ConflictResolutionStrategy::AutomaticMerge => Arc::new(AutomaticMergeResolver),
ConflictResolutionStrategy::Manual => Arc::new(ManualResolver),
ConflictResolutionStrategy::CustomScript => Arc::new(ManualResolver),
}
}
}
pub struct LastWriteWinsResolver;
#[async_trait::async_trait]
impl ConflictResolver for LastWriteWinsResolver {
async fn resolve(&self, conflict: &ConflictInfo) -> ConflictResolution {
let local_ts = conflict.local_vector.hlc_timestamp();
let remote_ts = conflict.remote_vector.hlc_timestamp();
if local_ts > remote_ts {
ConflictResolution::LocalWins
} else if remote_ts > local_ts {
ConflictResolution::RemoteWins
} else {
let local_cnt = conflict.local_vector.hlc_counter();
let remote_cnt = conflict.remote_vector.hlc_counter();
if local_cnt >= remote_cnt {
ConflictResolution::LocalWins
} else {
ConflictResolution::RemoteWins
}
}
}
fn name(&self) -> &'static str {
"last_write_wins"
}
}
pub struct DeterministicResolver;
#[async_trait::async_trait]
impl ConflictResolver for DeterministicResolver {
async fn resolve(&self, conflict: &ConflictInfo) -> ConflictResolution {
let local_node = conflict
.local_vector
.nodes()
.max()
.map(|s| s.as_str())
.unwrap_or("");
let remote_node = conflict
.remote_vector
.nodes()
.max()
.map(|s| s.as_str())
.unwrap_or("");
if local_node >= remote_node {
ConflictResolution::LocalWins
} else {
ConflictResolution::RemoteWins
}
}
fn name(&self) -> &'static str {
"deterministic"
}
}
pub struct AutomaticMergeResolver;
#[async_trait::async_trait]
impl ConflictResolver for AutomaticMergeResolver {
async fn resolve(&self, conflict: &ConflictInfo) -> ConflictResolution {
if let (Some(local), Some(remote)) = (&conflict.local_data, &conflict.remote_data) {
if let Some(merged) = attempt_crdt_merge(local, remote) {
return ConflictResolution::Merged(merged);
}
}
let resolver = LastWriteWinsResolver;
resolver.resolve(conflict).await
}
fn name(&self) -> &'static str {
"automatic_merge"
}
}
pub struct ManualResolver;
#[async_trait::async_trait]
impl ConflictResolver for ManualResolver {
async fn resolve(&self, conflict: &ConflictInfo) -> ConflictResolution {
ConflictResolution::KeepBoth {
local: conflict.local_data.clone().unwrap_or(Value::Null),
remote: conflict.remote_data.clone().unwrap_or(Value::Null),
}
}
fn name(&self) -> &'static str {
"manual"
}
}
fn attempt_crdt_merge(local: &Value, remote: &Value) -> Option<Value> {
match (local, remote) {
(Value::Object(local_map), Value::Object(remote_map)) => {
let mut merged = serde_json::Map::new();
for (key, value) in local_map {
merged.insert(key.clone(), value.clone());
}
for (key, remote_value) in remote_map {
if let Some(local_value) = local_map.get(key) {
if let Some(merged_value) = attempt_crdt_merge(local_value, remote_value) {
merged.insert(key.clone(), merged_value);
} else {
merged.insert(key.clone(), remote_value.clone());
}
} else {
merged.insert(key.clone(), remote_value.clone());
}
}
Some(Value::Object(merged))
}
_ => None,
}
}
pub fn detect_conflict(
local_vector: &VersionVector,
remote_vector: &VersionVector,
) -> Option<VectorComparison> {
let comparison = local_vector.compare(remote_vector);
match comparison {
VectorComparison::Concurrent => Some(comparison),
_ => None,
}
}
pub async fn resolve_conflict(
strategy: ConflictResolutionStrategy,
conflict: &ConflictInfo,
) -> ConflictResolution {
let resolver = strategy.create_resolver();
resolver.resolve(conflict).await
}
pub fn apply_resolution(
resolution: &ConflictResolution,
local: Option<&Value>,
remote: Option<&Value>,
) -> Option<Value> {
match resolution {
ConflictResolution::LocalWins => local.cloned(),
ConflictResolution::RemoteWins => remote.cloned(),
ConflictResolution::Merged(merged) => Some(merged.clone()),
ConflictResolution::KeepBoth { local, remote } => {
let mut conflict_doc = serde_json::Map::new();
conflict_doc.insert("_conflict".to_string(), Value::Bool(true));
conflict_doc.insert("_local".to_string(), local.clone());
conflict_doc.insert("_remote".to_string(), remote.clone());
Some(Value::Object(conflict_doc))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sync::version_vector::VersionVector;
use serde_json::json;
fn create_conflict_info(local_ts: u64, remote_ts: u64) -> ConflictInfo {
let mut local_vector = VersionVector::new();
local_vector.set_hlc(local_ts, 0);
local_vector.increment("node-1");
let mut remote_vector = VersionVector::new();
remote_vector.set_hlc(remote_ts, 0);
remote_vector.increment("node-2");
ConflictInfo {
document_key: "test-1".to_string(),
collection: "test".to_string(),
local_vector,
remote_vector,
local_data: Some(json!({"field": "local"})),
remote_data: Some(json!({"field": "remote"})),
detected_at: 0,
}
}
#[tokio::test]
async fn test_last_write_wins_local() {
let resolver = LastWriteWinsResolver;
let conflict = create_conflict_info(100, 50);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::LocalWins));
}
#[tokio::test]
async fn test_last_write_wins_remote() {
let resolver = LastWriteWinsResolver;
let conflict = create_conflict_info(50, 100);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::RemoteWins));
}
#[tokio::test]
async fn test_deterministic_resolution() {
let resolver = DeterministicResolver;
let conflict = create_conflict_info(50, 50);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::RemoteWins));
}
#[tokio::test]
async fn test_manual_resolution() {
let resolver = ManualResolver;
let conflict = create_conflict_info(50, 50);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::KeepBoth { .. }));
}
#[test]
fn test_crdt_merge_objects() {
let local = json!({
"name": "Alice",
"age": 30,
"crdt_counter": { "_type": "GCounter", "value": 5 }
});
let remote = json!({
"name": "Alice",
"age": 31,
"email": "alice@example.com"
});
let merged = attempt_crdt_merge(&local, &remote);
assert!(merged.is_some());
let merged = merged.unwrap();
assert_eq!(merged["name"], "Alice");
assert_eq!(merged["age"], 31); assert_eq!(merged["email"], "alice@example.com");
}
}