use borsh::{BorshDeserialize, BorshSerialize};
use super::handshake::{SyncCapabilities, SyncHandshake};
use super::levelwise::should_use_levelwise;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, BorshSerialize, BorshDeserialize)]
pub enum SyncProtocolKind {
None,
DeltaSync,
HashComparison,
Snapshot,
BloomFilter,
SubtreePrefetch,
LevelWise,
}
#[derive(Clone, Debug, PartialEq, BorshSerialize, BorshDeserialize)]
pub enum SyncProtocol {
None,
DeltaSync {
missing_delta_ids: Vec<[u8; 32]>,
},
HashComparison {
root_hash: [u8; 32],
},
Snapshot {
compressed: bool,
verified: bool,
},
BloomFilter {
filter_size: u64,
false_positive_rate: f64,
},
SubtreePrefetch {
subtree_roots: Vec<[u8; 32]>,
},
LevelWise {
max_depth: u32,
},
}
impl Default for SyncProtocol {
fn default() -> Self {
Self::None
}
}
impl SyncProtocol {
#[must_use]
pub fn kind(&self) -> SyncProtocolKind {
SyncProtocolKind::from(self)
}
}
impl From<&SyncProtocol> for SyncProtocolKind {
fn from(protocol: &SyncProtocol) -> Self {
match protocol {
SyncProtocol::None => Self::None,
SyncProtocol::DeltaSync { .. } => Self::DeltaSync,
SyncProtocol::HashComparison { .. } => Self::HashComparison,
SyncProtocol::Snapshot { .. } => Self::Snapshot,
SyncProtocol::BloomFilter { .. } => Self::BloomFilter,
SyncProtocol::SubtreePrefetch { .. } => Self::SubtreePrefetch,
SyncProtocol::LevelWise { .. } => Self::LevelWise,
}
}
}
#[derive(Clone, Debug)]
pub struct ProtocolSelection {
pub protocol: SyncProtocol,
pub reason: &'static str,
}
#[must_use]
pub fn calculate_divergence(local: &SyncHandshake, remote: &SyncHandshake) -> f64 {
let diff = local.entity_count.abs_diff(remote.entity_count);
let denominator = remote.entity_count.max(1);
diff as f64 / denominator as f64
}
#[must_use]
pub fn select_protocol(local: &SyncHandshake, remote: &SyncHandshake) -> ProtocolSelection {
if local.root_hash == remote.root_hash {
return ProtocolSelection {
protocol: SyncProtocol::None,
reason: "root hashes match, already in sync",
};
}
if !local.is_version_compatible(remote) {
return ProtocolSelection {
protocol: SyncProtocol::HashComparison {
root_hash: remote.root_hash,
},
reason: "version mismatch, using safe fallback",
};
}
if !local.has_state {
return ProtocolSelection {
protocol: SyncProtocol::Snapshot {
compressed: remote.entity_count > 100,
verified: true,
},
reason: "fresh node bootstrap via snapshot",
};
}
if !remote.has_state {
return ProtocolSelection {
protocol: SyncProtocol::None,
reason: "remote is fresh, skipping (remote must initiate its own sync)",
};
}
let divergence = calculate_divergence(local, remote);
if divergence > 0.5 {
return ProtocolSelection {
protocol: SyncProtocol::HashComparison {
root_hash: remote.root_hash,
},
reason: "high divergence (>50%), using hash comparison with CRDT merge",
};
}
if remote.max_depth > 3 && divergence < 0.2 {
return ProtocolSelection {
protocol: SyncProtocol::SubtreePrefetch {
subtree_roots: vec![], },
reason: "deep tree with low divergence, using subtree prefetch",
};
}
if remote.entity_count > 50 && divergence < 0.1 {
return ProtocolSelection {
protocol: SyncProtocol::BloomFilter {
filter_size: remote.entity_count.saturating_mul(10).min(10_000),
false_positive_rate: 0.01,
},
reason: "large tree with small divergence, using bloom filter",
};
}
let max_depth_usize = remote.max_depth as usize;
let avg_children_per_level = if remote.max_depth > 0 {
(remote.entity_count / u64::from(remote.max_depth)) as usize
} else {
0
};
if should_use_levelwise(max_depth_usize, avg_children_per_level) {
return ProtocolSelection {
protocol: SyncProtocol::LevelWise {
max_depth: remote.max_depth,
},
reason: "wide shallow tree, using level-wise sync",
};
}
ProtocolSelection {
protocol: SyncProtocol::HashComparison {
root_hash: remote.root_hash,
},
reason: "default: using hash comparison",
}
}
#[must_use]
pub fn is_protocol_supported(protocol: &SyncProtocol, capabilities: &SyncCapabilities) -> bool {
capabilities.supported_protocols.contains(&protocol.kind())
}
#[must_use]
pub fn select_protocol_with_fallback(
local: &SyncHandshake,
remote: &SyncHandshake,
remote_capabilities: &SyncCapabilities,
) -> ProtocolSelection {
let preferred = select_protocol(local, remote);
if is_protocol_supported(&preferred.protocol, remote_capabilities) {
return preferred;
}
if local.has_state {
let fallback = SyncProtocol::HashComparison {
root_hash: remote.root_hash,
};
if is_protocol_supported(&fallback, remote_capabilities) {
return ProtocolSelection {
protocol: fallback,
reason: "fallback to hash comparison (preferred not supported)",
};
}
}
ProtocolSelection {
protocol: SyncProtocol::None,
reason: "no mutually supported protocol found",
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sync::handshake::SYNC_PROTOCOL_VERSION;
#[test]
fn test_sync_protocol_roundtrip() {
let protocols = vec![
SyncProtocol::None,
SyncProtocol::DeltaSync {
missing_delta_ids: vec![[1; 32], [2; 32]],
},
SyncProtocol::HashComparison { root_hash: [3; 32] },
SyncProtocol::Snapshot {
compressed: true,
verified: false,
},
SyncProtocol::BloomFilter {
filter_size: 1024,
false_positive_rate: 0.01,
},
SyncProtocol::SubtreePrefetch {
subtree_roots: vec![[5; 32], [6; 32]],
},
SyncProtocol::LevelWise { max_depth: 3 },
];
for protocol in protocols {
let encoded = borsh::to_vec(&protocol).expect("serialize");
let decoded: SyncProtocol = borsh::from_slice(&encoded).expect("deserialize");
assert_eq!(protocol, decoded);
}
}
#[test]
fn test_sync_protocol_kind_roundtrip() {
let kinds = vec![
SyncProtocolKind::None,
SyncProtocolKind::DeltaSync,
SyncProtocolKind::HashComparison,
SyncProtocolKind::Snapshot,
SyncProtocolKind::BloomFilter,
SyncProtocolKind::SubtreePrefetch,
SyncProtocolKind::LevelWise,
];
for kind in kinds {
let encoded = borsh::to_vec(&kind).expect("serialize");
let decoded: SyncProtocolKind = borsh::from_slice(&encoded).expect("deserialize");
assert_eq!(kind, decoded);
}
}
#[test]
fn test_sync_protocol_kind_conversion() {
assert_eq!(SyncProtocol::None.kind(), SyncProtocolKind::None);
assert_eq!(
SyncProtocol::DeltaSync {
missing_delta_ids: vec![[1; 32]]
}
.kind(),
SyncProtocolKind::DeltaSync
);
assert_eq!(
SyncProtocol::HashComparison { root_hash: [2; 32] }.kind(),
SyncProtocolKind::HashComparison
);
assert_eq!(
SyncProtocol::Snapshot {
compressed: true,
verified: true
}
.kind(),
SyncProtocolKind::Snapshot
);
assert_eq!(
SyncProtocol::BloomFilter {
filter_size: 1024,
false_positive_rate: 0.01
}
.kind(),
SyncProtocolKind::BloomFilter
);
assert_eq!(
SyncProtocol::SubtreePrefetch {
subtree_roots: vec![]
}
.kind(),
SyncProtocolKind::SubtreePrefetch
);
assert_eq!(
SyncProtocol::LevelWise { max_depth: 5 }.kind(),
SyncProtocolKind::LevelWise
);
let protocol = SyncProtocol::HashComparison { root_hash: [3; 32] };
let kind: SyncProtocolKind = (&protocol).into();
assert_eq!(kind, SyncProtocolKind::HashComparison);
}
#[test]
fn test_calculate_divergence() {
let local = SyncHandshake::new([1; 32], 100, 5, vec![]);
let remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
assert!((calculate_divergence(&local, &remote) - 0.0).abs() < f64::EPSILON);
let local = SyncHandshake::new([1; 32], 50, 5, vec![]);
let remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
assert!((calculate_divergence(&local, &remote) - 0.5).abs() < f64::EPSILON);
let local = SyncHandshake::new([1; 32], 0, 0, vec![]);
let remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
assert!((calculate_divergence(&local, &remote) - 1.0).abs() < f64::EPSILON);
let local = SyncHandshake::new([1; 32], 100, 5, vec![]);
let remote = SyncHandshake::new([2; 32], 0, 0, vec![]);
assert!((calculate_divergence(&local, &remote) - 100.0).abs() < f64::EPSILON);
}
#[test]
fn test_select_protocol_rule1_already_synced() {
let local = SyncHandshake::new([42; 32], 100, 5, vec![]);
let remote = SyncHandshake::new([42; 32], 200, 3, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(selection.protocol, SyncProtocol::None));
assert!(selection.reason.contains("already in sync"));
}
#[test]
fn test_select_protocol_rule2_fresh_node_gets_snapshot() {
let local = SyncHandshake::new([0; 32], 0, 0, vec![]); let remote = SyncHandshake::new([42; 32], 200, 5, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(selection.protocol, SyncProtocol::Snapshot { .. }));
assert!(selection.reason.contains("fresh node"));
}
#[test]
fn test_select_protocol_rule2b_remote_fresh_returns_none() {
let local = SyncHandshake::new([42; 32], 100, 5, vec![[1; 32]]);
let remote = SyncHandshake::new([0; 32], 0, 0, vec![]);
let selection = select_protocol(&local, &remote);
assert!(
matches!(selection.protocol, SyncProtocol::None),
"should be None when remote is fresh, got {:?}",
selection.protocol
);
assert!(selection.reason.contains("remote is fresh"));
}
#[test]
fn test_select_protocol_both_fresh_rule2a_takes_precedence() {
let local = SyncHandshake::new([0; 32], 0, 0, vec![]); let remote = SyncHandshake::new([0; 32], 0, 0, vec![]);
let selection = select_protocol(&local, &remote);
assert!(
matches!(selection.protocol, SyncProtocol::None),
"two fresh nodes with same root hash should be None (Rule 1), got {:?}",
selection.protocol
);
assert!(selection.reason.contains("already in sync"));
let local = SyncHandshake::new([0; 32], 0, 0, vec![]); let remote = SyncHandshake::new([42; 32], 0, 0, vec![]);
let selection = select_protocol(&local, &remote);
assert!(
matches!(selection.protocol, SyncProtocol::Snapshot { .. }),
"fresh local should get Snapshot (Rule 2a) not None (Rule 2b), got {:?}",
selection.protocol
);
assert!(selection.reason.contains("fresh node"));
}
#[test]
fn test_select_protocol_rule3_initialized_node_never_gets_snapshot() {
let local = SyncHandshake::new([1; 32], 1, 1, vec![]); let remote = SyncHandshake::new([42; 32], 200, 5, vec![]);
let selection = select_protocol(&local, &remote);
assert!(!matches!(selection.protocol, SyncProtocol::Snapshot { .. }));
}
#[test]
fn test_select_protocol_rule3_high_divergence_uses_hash_comparison() {
let local = SyncHandshake::new([1; 32], 10, 2, vec![]); let remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(
selection.protocol,
SyncProtocol::HashComparison { .. }
));
assert!(selection.reason.contains("divergence"));
}
#[test]
fn test_select_protocol_rule4_deep_tree_uses_subtree_prefetch() {
let local = SyncHandshake::new([1; 32], 90, 5, vec![]); let remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(
selection.protocol,
SyncProtocol::SubtreePrefetch { .. }
));
assert!(selection.reason.contains("subtree"));
}
#[test]
fn test_select_protocol_rule5_large_tree_small_diff_uses_bloom() {
let local = SyncHandshake::new([1; 32], 95, 2, vec![]); let remote = SyncHandshake::new([2; 32], 100, 2, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(
selection.protocol,
SyncProtocol::BloomFilter { .. }
));
assert!(selection.reason.contains("bloom"));
}
#[test]
fn test_select_protocol_rule6_wide_shallow_uses_levelwise() {
let local = SyncHandshake::new([1; 32], 40, 2, vec![]);
let remote = SyncHandshake::new([2; 32], 40, 2, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(selection.protocol, SyncProtocol::LevelWise { .. }));
assert!(selection.reason.contains("level"));
}
#[test]
fn test_select_protocol_rule7_default_uses_hash_comparison() {
let local = SyncHandshake::new([1; 32], 30, 2, vec![]); let remote = SyncHandshake::new([2; 32], 40, 3, vec![]);
let selection = select_protocol(&local, &remote);
assert!(matches!(
selection.protocol,
SyncProtocol::HashComparison { .. }
));
assert!(selection.reason.contains("default"));
}
#[test]
fn test_select_protocol_version_mismatch_uses_safe_fallback() {
let local = SyncHandshake::new([1; 32], 100, 5, vec![]);
let mut remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
remote.version = SYNC_PROTOCOL_VERSION + 1;
let selection = select_protocol(&local, &remote);
assert!(matches!(
selection.protocol,
SyncProtocol::HashComparison { .. }
));
assert!(selection.reason.contains("version mismatch"));
}
#[test]
fn test_is_protocol_supported() {
let caps = SyncCapabilities::default();
assert!(is_protocol_supported(&SyncProtocol::None, &caps));
assert!(is_protocol_supported(
&SyncProtocol::HashComparison { root_hash: [0; 32] },
&caps
));
assert!(!is_protocol_supported(
&SyncProtocol::SubtreePrefetch {
subtree_roots: vec![]
},
&caps
));
assert!(is_protocol_supported(
&SyncProtocol::LevelWise { max_depth: 2 },
&caps
));
}
#[test]
fn test_select_protocol_with_fallback() {
let local = SyncHandshake::new([1; 32], 90, 5, vec![]); let remote = SyncHandshake::new([2; 32], 100, 5, vec![]);
let caps = SyncCapabilities::default();
let selection = select_protocol_with_fallback(&local, &remote, &caps);
assert!(matches!(
selection.protocol,
SyncProtocol::HashComparison { .. }
));
assert!(selection.reason.contains("fallback"));
}
}