use crate::sync::version_vector::{ConflictInfo, VectorComparison, VersionVector};
use mlua::{Lua, Result as LuaResult, Value as LuaValue};
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"
}
}
pub struct CustomScriptResolver {
script: String,
}
impl CustomScriptResolver {
pub fn new(script: impl Into<String>) -> Self {
Self {
script: script.into(),
}
}
fn json_to_lua(lua: &Lua, value: &Value) -> LuaResult<LuaValue> {
match value {
Value::Null => Ok(LuaValue::Nil),
Value::Bool(b) => Ok(LuaValue::Boolean(*b)),
Value::Number(n) => {
if let Some(i) = n.as_i64() {
Ok(LuaValue::Integer(i))
} else if let Some(f) = n.as_f64() {
Ok(LuaValue::Number(f))
} else {
Ok(LuaValue::Nil)
}
}
Value::String(s) => lua.create_string(s).map(LuaValue::String),
Value::Array(arr) => {
let table = lua.create_table()?;
for (i, v) in arr.iter().enumerate() {
table.set(i + 1, Self::json_to_lua(lua, v)?)?;
}
Ok(LuaValue::Table(table))
}
Value::Object(obj) => {
let table = lua.create_table()?;
for (k, v) in obj {
table.set(k.as_str(), Self::json_to_lua(lua, v)?)?;
}
Ok(LuaValue::Table(table))
}
}
}
fn lua_to_json(value: LuaValue) -> Option<Value> {
match value {
LuaValue::Nil => Some(Value::Null),
LuaValue::Boolean(b) => Some(Value::Bool(b)),
LuaValue::Integer(i) => Some(Value::Number(i.into())),
LuaValue::Number(f) => serde_json::Number::from_f64(f).map(Value::Number),
LuaValue::String(s) => Some(Value::String(s.to_string_lossy().to_string())),
LuaValue::Table(table) => {
let mut has_string_keys = false;
let mut max_int_key = 0i64;
let mut int_key_count = 0usize;
for (k, _) in table.clone().pairs::<LuaValue, LuaValue>().flatten() {
match k {
LuaValue::Integer(i) => {
int_key_count += 1;
if i > max_int_key {
max_int_key = i;
}
}
LuaValue::String(_) => {
has_string_keys = true;
}
_ => {}
}
}
let is_array =
!has_string_keys && int_key_count > 0 && max_int_key == int_key_count as i64;
if is_array {
let mut arr: Vec<(i64, Value)> = Vec::new();
for (k, v) in table.pairs::<i64, LuaValue>().flatten() {
if let Some(val) = Self::lua_to_json(v) {
arr.push((k, val));
}
}
arr.sort_by_key(|(k, _)| *k);
Some(Value::Array(arr.into_iter().map(|(_, v)| v).collect()))
} else {
let mut obj = serde_json::Map::new();
for (k, v) in table.pairs::<LuaValue, LuaValue>().flatten() {
let key = match k {
LuaValue::String(s) => s.to_string_lossy().to_string(),
LuaValue::Integer(i) => i.to_string(),
LuaValue::Number(n) => n.to_string(),
_ => continue,
};
if let Some(val) = Self::lua_to_json(v) {
obj.insert(key, val);
}
}
Some(Value::Object(obj))
}
}
_ => None,
}
}
}
#[async_trait::async_trait]
impl ConflictResolver for CustomScriptResolver {
async fn resolve(&self, conflict: &ConflictInfo) -> ConflictResolution {
let lua = Lua::new();
let globals = lua.globals();
if let Ok(local_val) =
Self::json_to_lua(&lua, &conflict.local_data.clone().unwrap_or(Value::Null))
{
let _ = globals.set("local_doc", local_val);
}
if let Ok(remote_val) =
Self::json_to_lua(&lua, &conflict.remote_data.clone().unwrap_or(Value::Null))
{
let _ = globals.set("remote_doc", remote_val);
}
let _ = globals.set("key", conflict.document_key.clone());
let _ = globals.set("collection", conflict.collection.clone());
match lua.load(&self.script).eval::<LuaValue>() {
Ok(result) => match &result {
LuaValue::String(s) => {
let s_str = s.to_string_lossy();
match s_str.as_ref() {
"local" => ConflictResolution::LocalWins,
"remote" => ConflictResolution::RemoteWins,
_ => ConflictResolution::LocalWins, }
}
LuaValue::Table(_) => {
if let Some(merged) = Self::lua_to_json(result) {
ConflictResolution::Merged(merged)
} else {
ConflictResolution::LocalWins
}
}
_ => ConflictResolution::LocalWins,
},
Err(e) => {
tracing::error!("Custom conflict resolver script error: {}", e);
ConflictResolution::LocalWins
}
}
}
fn name(&self) -> &'static str {
"custom_script"
}
}
pub fn create_custom_resolver(script: impl Into<String>) -> Arc<dyn ConflictResolver> {
Arc::new(CustomScriptResolver::new(script))
}
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");
}
#[tokio::test]
async fn test_custom_script_returns_local() {
let script = r#"
return "local"
"#;
let resolver = CustomScriptResolver::new(script);
let conflict = create_conflict_info(50, 100);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::LocalWins));
}
#[tokio::test]
async fn test_custom_script_returns_remote() {
let script = r#"
return "remote"
"#;
let resolver = CustomScriptResolver::new(script);
let conflict = create_conflict_info(50, 100);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::RemoteWins));
}
#[tokio::test]
async fn test_custom_script_returns_merged() {
let script = r#"
-- Merge local and remote, taking name from local and field from remote
local result = {}
if local_doc and local_doc.name then
result.name = local_doc.name
else
result.name = "unknown"
end
if remote_doc and remote_doc.field then
result.field = remote_doc.field
end
return result
"#;
let mut local_vector = VersionVector::new();
local_vector.set_hlc(50, 0);
local_vector.increment("node-1");
let mut remote_vector = VersionVector::new();
remote_vector.set_hlc(100, 0);
remote_vector.increment("node-2");
let conflict = ConflictInfo {
document_key: "test-1".to_string(),
collection: "test".to_string(),
local_vector,
remote_vector,
local_data: Some(json!({"name": "Alice", "field": "local_value"})),
remote_data: Some(json!({"name": "Bob", "field": "remote_value"})),
detected_at: 0,
};
let resolver = CustomScriptResolver::new(script);
let result = resolver.resolve(&conflict).await;
match result {
ConflictResolution::Merged(merged) => {
assert_eq!(merged["name"], "Alice");
assert_eq!(merged["field"], "remote_value");
}
_ => panic!("Expected Merged resolution"),
}
}
#[tokio::test]
async fn test_custom_script_has_access_to_key() {
let script = r#"
-- Return remote if key starts with "important"
if string.sub(key, 1, 9) == "important" then
return "remote"
else
return "local"
end
"#;
let mut local_vector = VersionVector::new();
local_vector.set_hlc(50, 0);
local_vector.increment("node-1");
let mut remote_vector = VersionVector::new();
remote_vector.set_hlc(100, 0);
remote_vector.increment("node-2");
let conflict = ConflictInfo {
document_key: "important-doc-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,
};
let resolver = CustomScriptResolver::new(script);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::RemoteWins));
}
#[tokio::test]
async fn test_custom_script_error_defaults_to_local() {
let script = r#"
-- Invalid Lua that will error
this_function_does_not_exist()
"#;
let resolver = CustomScriptResolver::new(script);
let conflict = create_conflict_info(50, 100);
let result = resolver.resolve(&conflict).await;
assert!(matches!(result, ConflictResolution::LocalWins));
}
}