use super::message::{MessageError, SyncMessage};
use super::tracker::SyncTracker;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum SyncError {
#[error("message error: {0}")]
Message(#[from] MessageError),
#[error("diff decode error: {0}")]
DiffDecode(String),
#[error("diff apply error: {0}")]
DiffApply(String),
#[error("version mismatch: expected base {expected}, got {actual}")]
VersionMismatch {
expected: u64,
actual: u64,
},
#[error("state not initialized")]
NotInitialized,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ProcessResult {
Updated,
AckOnly,
Duplicate,
}
pub struct SyncEngine<S, D> {
tracker: SyncTracker,
state: Option<S>,
acked_snapshot: Option<S>,
encode_diff: fn(&D) -> Vec<u8>,
decode_diff: fn(&[u8]) -> Result<D, String>,
compute_diff: fn(&S, &S) -> D,
apply_diff: fn(&mut S, &D) -> Result<(), String>,
is_diff_empty: fn(&D) -> bool,
}
impl<S: Clone, D> SyncEngine<S, D> {
pub fn new(
encode_diff: fn(&D) -> Vec<u8>,
decode_diff: fn(&[u8]) -> Result<D, String>,
compute_diff: fn(&S, &S) -> D,
apply_diff: fn(&mut S, &D) -> Result<(), String>,
is_diff_empty: fn(&D) -> bool,
) -> Self {
Self {
tracker: SyncTracker::new(),
state: None,
acked_snapshot: None,
encode_diff,
decode_diff,
compute_diff,
apply_diff,
is_diff_empty,
}
}
pub fn init(&mut self, initial_state: S) {
self.state = Some(initial_state.clone());
self.acked_snapshot = Some(initial_state);
self.tracker.reset();
}
pub fn is_initialized(&self) -> bool {
self.state.is_some()
}
pub fn state(&self) -> Option<&S> {
self.state.as_ref()
}
pub fn state_mut(&mut self) -> Option<&mut S> {
self.state.as_mut()
}
pub fn mark_changed(&mut self) -> u64 {
self.tracker.bump_version()
}
pub fn update_state(&mut self, new_state: S) -> u64 {
self.state = Some(new_state);
self.tracker.bump_version()
}
pub fn tracker(&self) -> &SyncTracker {
&self.tracker
}
pub fn has_pending_updates(&self) -> bool {
self.tracker.has_pending_updates()
}
pub fn needs_ack(&self) -> bool {
self.tracker.needs_ack()
}
pub fn generate_message(&mut self) -> Result<Option<SyncMessage>, SyncError> {
let state = self.state.as_ref().ok_or(SyncError::NotInitialized)?;
if !self.tracker.has_pending_updates() && !self.tracker.needs_ack() {
return Ok(None);
}
if !self.tracker.has_pending_updates() {
let msg = self.tracker.create_ack();
return Ok(Some(msg));
}
let base_state = self.acked_snapshot.as_ref().ok_or(SyncError::NotInitialized)?;
let diff = (self.compute_diff)(base_state, state);
let diff_bytes = if (self.is_diff_empty)(&diff) {
Vec::new()
} else {
(self.encode_diff)(&diff)
};
let base_version = self.tracker.diff_base_version();
let msg = self.tracker.create_message(diff_bytes, base_version);
self.tracker.record_sent(self.tracker.current_version());
Ok(Some(msg))
}
pub fn generate_ack(&self) -> Result<SyncMessage, SyncError> {
if !self.is_initialized() {
return Err(SyncError::NotInitialized);
}
Ok(self.tracker.create_ack())
}
pub fn process_message(&mut self, msg: &SyncMessage) -> Result<ProcessResult, SyncError> {
let state = self.state.as_mut().ok_or(SyncError::NotInitialized)?;
let is_new = self.tracker.process_incoming(msg);
if msg.is_ack_only() {
if msg.acked_state_num > 0 {
self.update_acked_snapshot();
}
return Ok(ProcessResult::AckOnly);
}
if !is_new {
return Ok(ProcessResult::Duplicate);
}
if !msg.diff.is_empty() {
let diff = (self.decode_diff)(&msg.diff)
.map_err(SyncError::DiffDecode)?;
(self.apply_diff)(state, &diff)
.map_err(SyncError::DiffApply)?;
}
if msg.acked_state_num > 0 {
self.update_acked_snapshot();
}
Ok(ProcessResult::Updated)
}
fn update_acked_snapshot(&mut self) {
if let Some(state) = &self.state {
if self.tracker.last_acked_version() > 0 {
self.acked_snapshot = Some(state.clone());
}
}
}
pub fn current_version(&self) -> u64 {
self.tracker.current_version()
}
pub fn peer_version(&self) -> u64 {
self.tracker.peer_version()
}
pub fn is_synchronized(&self) -> bool {
self.tracker.is_synchronized()
}
pub fn reset(&mut self) {
self.tracker.reset();
self.state = None;
self.acked_snapshot = None;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Debug, Clone, PartialEq)]
struct TestState {
value: i32,
}
#[derive(Debug, Clone, PartialEq)]
struct TestDiff {
delta: i32,
}
fn encode_diff(diff: &TestDiff) -> Vec<u8> {
diff.delta.to_le_bytes().to_vec()
}
fn decode_diff(data: &[u8]) -> Result<TestDiff, String> {
if data.len() != 4 {
return Err("invalid diff length".to_string());
}
let delta = i32::from_le_bytes(data.try_into().unwrap());
Ok(TestDiff { delta })
}
fn compute_diff(old: &TestState, new: &TestState) -> TestDiff {
TestDiff {
delta: new.value - old.value,
}
}
fn apply_diff(state: &mut TestState, diff: &TestDiff) -> Result<(), String> {
state.value += diff.delta;
Ok(())
}
fn is_diff_empty(diff: &TestDiff) -> bool {
diff.delta == 0
}
fn create_engine() -> SyncEngine<TestState, TestDiff> {
SyncEngine::new(encode_diff, decode_diff, compute_diff, apply_diff, is_diff_empty)
}
#[test]
fn test_init() {
let mut engine = create_engine();
assert!(!engine.is_initialized());
engine.init(TestState { value: 42 });
assert!(engine.is_initialized());
assert_eq!(engine.state().unwrap().value, 42);
}
#[test]
fn test_update_state() {
let mut engine = create_engine();
engine.init(TestState { value: 0 });
let version = engine.update_state(TestState { value: 100 });
assert_eq!(version, 1);
assert_eq!(engine.state().unwrap().value, 100);
assert!(engine.has_pending_updates());
}
#[test]
fn test_generate_message() {
let mut engine = create_engine();
engine.init(TestState { value: 0 });
let msg = engine.generate_message().unwrap();
assert!(msg.is_none());
engine.update_state(TestState { value: 10 });
let msg = engine.generate_message().unwrap().unwrap();
assert_eq!(msg.sender_state_num, 1);
assert!(!msg.is_ack_only());
let diff = decode_diff(&msg.diff).unwrap();
assert_eq!(diff.delta, 10);
}
#[test]
fn test_process_message() {
let mut engine = create_engine();
engine.init(TestState { value: 0 });
let diff = TestDiff { delta: 50 };
let msg = SyncMessage::new(1, 0, 0, encode_diff(&diff));
let result = engine.process_message(&msg).unwrap();
assert_eq!(result, ProcessResult::Updated);
assert_eq!(engine.state().unwrap().value, 50);
assert_eq!(engine.peer_version(), 1);
}
#[test]
fn test_process_ack_only() {
let mut engine = create_engine();
engine.init(TestState { value: 0 });
engine.update_state(TestState { value: 10 });
let msg = SyncMessage::ack_only(1, 1);
let result = engine.process_message(&msg).unwrap();
assert_eq!(result, ProcessResult::AckOnly);
}
#[test]
fn test_duplicate_message() {
let mut engine = create_engine();
engine.init(TestState { value: 0 });
let diff = TestDiff { delta: 10 };
let msg = SyncMessage::new(1, 0, 0, encode_diff(&diff));
engine.process_message(&msg).unwrap();
let result = engine.process_message(&msg).unwrap();
assert_eq!(result, ProcessResult::Duplicate);
}
#[test]
fn test_bidirectional_sync() {
let mut engine_a = create_engine();
let mut engine_b = create_engine();
engine_a.init(TestState { value: 0 });
engine_b.init(TestState { value: 0 });
engine_a.update_state(TestState { value: 100 });
let msg_from_a = engine_a.generate_message().unwrap().unwrap();
engine_b.process_message(&msg_from_a).unwrap();
assert_eq!(engine_b.state().unwrap().value, 100);
assert_eq!(engine_b.peer_version(), 1);
let ack_from_b = engine_b.generate_ack().unwrap();
engine_a.process_message(&ack_from_b).unwrap();
assert_eq!(engine_a.tracker().last_acked_version(), 1);
}
#[test]
fn test_not_initialized_error() {
let mut engine = create_engine();
let result = engine.generate_message();
assert!(matches!(result, Err(SyncError::NotInitialized)));
let msg = SyncMessage::ack_only(1, 0);
let result = engine.process_message(&msg);
assert!(matches!(result, Err(SyncError::NotInitialized)));
}
#[test]
fn test_empty_diff() {
let mut engine = create_engine();
engine.init(TestState { value: 42 });
engine.mark_changed();
let msg = engine.generate_message().unwrap().unwrap();
assert!(msg.diff.is_empty() || msg.is_ack_only());
}
#[test]
fn test_reset() {
let mut engine = create_engine();
engine.init(TestState { value: 100 });
engine.update_state(TestState { value: 200 });
engine.reset();
assert!(!engine.is_initialized());
assert!(engine.state().is_none());
assert_eq!(engine.current_version(), 0);
}
}