use rustc_hash::FxHashMap;
use crate::{
types::{
ArchivedBalancesDiff, BalancesDiff, BitTagVec, SET_NO_CHANGE, SET_TO_DIFF,
SET_TO_TARGET_VALUE, SET_TO_ZERO,
},
Error,
};
pub fn diff_balances(base: &[u64], target: &[u64]) -> BalancesDiff {
diff_balances_iter(base.iter().copied(), target.iter().copied())
}
pub fn diff_balances_iter<I1, I2>(mut base: I1, mut target: I2) -> BalancesDiff
where
I1: ExactSizeIterator<Item = u64>,
I2: ExactSizeIterator<Item = u64>,
{
let common_len = base.len().min(target.len());
let mut changes = Vec::with_capacity(1024);
for idx in 0..common_len {
let Some(v1) = base.next() else {
break;
};
let Some(v2) = target.next() else {
break;
};
if v1 != v2 {
let diff = v2
.checked_sub(v1)
.and_then(|value| i64::try_from(value).ok())
.or_else(|| {
v1.checked_sub(v2)
.and_then(|value| i64::try_from(value).ok())
.map(|value| -value)
});
changes.push(Change {
idx,
diff,
target: v2,
});
}
}
let mode = find_mode(&changes);
let (tags, varint_payload, target_values) = encode(common_len, &changes, mode);
BalancesDiff {
tags,
mode,
varint_payload,
target_values,
appended_balances: target.collect(),
}
}
pub fn apply_balances_iter<T: crate::ListMutTarget<u64>>(
target: &mut T,
delta: &ArchivedBalancesDiff,
) -> Result<(), Error> {
let mode = delta.mode.to_native();
let tag_len = usize::try_from(delta.tags.len.to_native())
.map_err(|_| Error::MalformedDelta("delta tag length does not fit in usize".into()))?;
if target.len() != tag_len {
return Err(Error::MalformedDelta(format!(
"target length {} does not match delta tag length {tag_len}",
target.len()
)));
}
let mut target_iter = delta.target_values.iter();
let payload = delta.varint_payload.as_slice();
let mut payload_cursor = 0usize;
let mut base_idx = 0usize;
for &tag_byte in delta.tags.data.iter() {
if base_idx >= tag_len {
break;
}
if tag_byte == 0 {
base_idx = (base_idx + 4).min(tag_len);
continue;
}
for bit in 0..4 {
if base_idx >= tag_len {
break;
}
let tag = (tag_byte >> (bit * 2)) & 0b11;
match tag {
SET_NO_CHANGE => {}
SET_TO_ZERO => {
let Some(value) = target.get_mut(base_idx) else {
return Err(Error::MalformedDelta(format!(
"target collection is missing balance at index {base_idx}"
)));
};
*value = 0;
}
SET_TO_TARGET_VALUE => {
let Some(target_value) = target_iter.next() else {
return Err(Error::MalformedDelta(
"target-value payload is shorter than the encoded target-value tags"
.into(),
));
};
let Some(value) = target.get_mut(base_idx) else {
return Err(Error::MalformedDelta(format!(
"target collection is missing balance at index {base_idx}"
)));
};
*value = target_value.to_native();
}
SET_TO_DIFF => {
let encoded = read_varint(payload, &mut payload_cursor)?;
let corrected = zigzag_decode(encoded);
let Some(diff) = corrected.checked_add(mode) else {
return Err(Error::MalformedDelta(
"decoded balance difference overflows i64".into(),
));
};
let Some(value) = target.get_mut(base_idx) else {
return Err(Error::MalformedDelta(format!(
"target collection is missing balance at index {base_idx}"
)));
};
let base_value = i64::try_from(*value).map_err(|_| {
Error::MalformedDelta(format!(
"base balance at index {base_idx} exceeds i64 range"
))
})?;
let updated = base_value
.checked_add(diff)
.and_then(|value| u64::try_from(value).ok())
.ok_or_else(|| {
Error::MalformedDelta(format!(
"decoded balance difference produces an invalid value at index {base_idx}"
))
})?;
*value = updated;
}
_ => {
return Err(Error::MalformedDelta(format!(
"invalid balance tag {tag} at index {base_idx}"
)));
}
}
base_idx += 1;
}
}
if !delta.appended_balances.is_empty() {
for value in delta.appended_balances.iter() {
target.push(value.to_native());
}
}
Ok(())
}
pub fn apply_balances(base: &mut Vec<u64>, delta: &ArchivedBalancesDiff) -> Result<(), Error> {
apply_balances_iter(base, delta)
}
struct Change {
idx: usize,
diff: Option<i64>,
target: u64,
}
fn find_mode(changes: &[Change]) -> i64 {
let mut freq_map = FxHashMap::default();
freq_map.reserve(256);
for change in changes {
let Some(diff) = change.diff else {
continue;
};
if i32::try_from(diff).is_ok() {
*freq_map.entry(diff).or_insert(0usize) += 1;
}
}
freq_map
.into_iter()
.max_by_key(|&(_, count)| count)
.map(|(value, _)| value)
.unwrap_or(0)
}
fn encode(common_len: usize, changes: &[Change], mode: i64) -> (BitTagVec, Vec<u8>, Vec<u64>) {
let mut tags = BitTagVec::new(common_len);
let mut varint_payload = Vec::with_capacity(changes.len());
let mut target_values = Vec::new();
for change in changes {
let Change { idx, diff, target } = *change;
if target == 0 {
tags.set(idx, SET_TO_ZERO);
continue;
}
let Some(diff) = diff else {
tags.set(idx, SET_TO_TARGET_VALUE);
target_values.push(target);
continue;
};
if i32::try_from(diff).is_ok() {
tags.set(idx, SET_TO_DIFF);
let corrected = diff - mode;
write_varint(zigzag_encode(corrected), &mut varint_payload);
} else {
tags.set(idx, SET_TO_TARGET_VALUE);
target_values.push(target);
}
}
(tags, varint_payload, target_values)
}
#[inline]
fn zigzag_encode(n: i64) -> u64 {
((n as u64) << 1) ^ ((n >> 63) as u64)
}
#[inline]
fn zigzag_decode(n: u64) -> i64 {
((n >> 1) as i64) ^ -((n & 1) as i64)
}
#[inline]
pub(super) fn write_varint(mut val: u64, buf: &mut Vec<u8>) {
loop {
if val < 0x80 {
buf.push(val as u8);
break;
}
buf.push((val as u8) | 0x80);
val >>= 7;
}
}
#[inline]
pub(super) fn read_varint(buf: &[u8], cursor: &mut usize) -> Result<u64, Error> {
let mut value = 0u64;
for shift in (0..=63).step_by(7) {
let Some(&byte) = buf.get(*cursor) else {
return Err(Error::MalformedDelta(format!(
"truncated varint payload at byte offset {}",
*cursor
)));
};
*cursor += 1;
let payload = (byte & 0x7f) as u64;
if shift == 63 && payload > 1 {
return Err(Error::MalformedDelta(format!(
"varint overflows u64 at byte offset {}",
*cursor - 1
)));
}
value |= payload << shift;
if byte & 0x80 == 0 {
return Ok(value);
}
}
Err(Error::MalformedDelta(
"varint exceeds the maximum length for a u64".into(),
))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::ArchivedBalancesDiff;
use crate::ListMutTarget;
struct MockTarget {
inner: Vec<u64>,
}
impl ListMutTarget<u64> for MockTarget {
fn len(&self) -> usize {
self.inner.len()
}
fn get_mut(&mut self, index: usize) -> Option<&mut u64> {
self.inner.get_mut(index)
}
fn push(&mut self, value: u64) {
self.inner.push(value);
}
}
fn assert_roundtrip(base: &[u64], target: &[u64]) {
let delta = diff_balances(base, target);
let bytes = rkyv::to_bytes::<rkyv::rancor::Error>(&delta)
.expect("test setup: failed to serialize delta");
let archived = rkyv::access::<ArchivedBalancesDiff, rkyv::rancor::Error>(&bytes)
.expect("test setup: failed to access archived delta");
let mut reconstructed = base.to_vec();
apply_balances(&mut reconstructed, archived)
.expect("test setup: failed to apply valid delta");
assert_eq!(
reconstructed, target,
"Roundtrip failed: base={base:?}, target={target:?}, delta={delta:?}"
);
}
#[test]
fn test_no_changes() {
let state = vec![1000; 10];
assert_roundtrip(&state, &state);
}
#[test]
fn test_single_balance_change() {
let base = vec![100, 200, 300];
let target = vec![100, 205, 300];
assert_roundtrip(&base, &target);
}
#[test]
fn test_all_set_to_zero() {
let base = vec![100, 200, 300];
let target = vec![0, 0, 0];
assert_roundtrip(&base, &target);
}
#[test]
fn test_appended_balances() {
let base = vec![100, 200];
let target = vec![100, 200, 300, 400];
assert_roundtrip(&base, &target);
}
#[test]
fn test_only_appended_balances() {
let base = vec![];
let target = vec![10, 20, 30];
assert_roundtrip(&base, &target);
}
#[test]
fn test_mode_selection_frequent_diff() {
let base: Vec<u64> = (0..1000).map(|i| (i * 100) as u64).collect();
let mut target = base.clone();
for i in (0..1000).step_by(2) {
target[i] += 1000;
}
let delta = diff_balances(&base, &target);
assert_eq!(
delta.mode, 1000,
"Mode should select the most frequent difference"
);
assert_roundtrip(&base, &target);
}
#[test]
fn test_i32_boundary_max_diff() {
let base = vec![100];
let target = vec![100 + i32::MAX as u64];
assert_roundtrip(&base, &target);
}
#[test]
fn test_i32_boundary_min_diff() {
let base = vec![100 + 2147483647];
let target = vec![100];
assert_roundtrip(&base, &target);
}
#[test]
fn test_i32_overflow_uses_target_value() {
let base = vec![100];
let target = vec![100 + i32::MAX as u64 + 1];
let delta = diff_balances(&base, &target);
assert_eq!(
delta.target_values.len(),
1,
"u64 diff exceeding i32::MAX must use SET_TO_TARGET_VALUE"
);
assert_eq!(
delta.varint_payload.len(),
0,
"Should not use varint payload for unrepresentable diff"
);
assert_roundtrip(&base, &target);
}
#[test]
fn test_u64_extreme_values_use_target_value() {
let base = vec![0];
let target = vec![u64::MAX];
let delta = diff_balances(&base, &target);
assert_eq!(
delta.target_values.len(),
1,
"u64::MAX diff must use SET_TO_TARGET_VALUE"
);
assert_eq!(
delta.varint_payload.len(),
0,
"Should not use varint payload for unrepresentable diff"
);
assert_roundtrip(&[0], &[u64::MAX]);
assert_roundtrip(&[u64::MAX], &[0]);
}
#[test]
fn test_apply_length_mismatch_returns_error() {
let base = vec![100, 200];
let target = vec![100, 200];
let delta = diff_balances(&base, &target);
let bytes = rkyv::to_bytes::<rkyv::rancor::Error>(&delta)
.expect("test setup: failed to serialize delta");
let archived = rkyv::access::<ArchivedBalancesDiff, rkyv::rancor::Error>(&bytes)
.expect("test setup: failed to access archived delta");
let mut wrong_base = MockTarget { inner: vec![100] }; let result = apply_balances_iter(&mut wrong_base, archived);
assert!(result.is_err(), "Should error on length mismatch");
let err_str = format!("{}", result.expect_err("test setup: expected error"));
assert!(
err_str.contains("does not match delta tag length"),
"Error message should mention length mismatch"
);
}
#[test]
fn test_apply_i64_overflow_returns_error() {
let base = vec![100, 200];
let target = vec![101, 201];
let delta = diff_balances(&base, &target);
assert_eq!(delta.mode, 1);
let mut bytes = rkyv::to_bytes::<rkyv::rancor::Error>(&delta)
.expect("test setup: failed to serialize delta");
let mode_bytes = 1i64.to_le_bytes();
let max_bytes = i64::MAX.to_le_bytes();
let pos = bytes
.windows(8)
.position(|w| w == mode_bytes)
.expect("test setup: could not find mode in serialized delta");
bytes
.get_mut(pos..pos + 8)
.expect("test setup: `pos` is valid and derived from a window of length 8")
.copy_from_slice(&max_bytes);
let archived = rkyv::access::<ArchivedBalancesDiff, rkyv::rancor::Error>(&bytes)
.expect("test setup: failed to access archived delta");
let mut state = MockTarget { inner: base };
let result = apply_balances_iter(&mut state, archived);
assert!(result.is_err(), "Should error on overflow during apply");
let err_str = format!("{}", result.expect_err("test setup"));
assert!(
err_str.contains("produces an invalid value"),
"Error message should mention the invalid resulting value"
);
}
#[test]
fn test_apply_base_exceeds_i64_range_returns_error() {
let base = vec![i64::MAX as u64 + 1];
let target = vec![i64::MAX as u64 + 2];
let delta = diff_balances(&base, &target);
let bytes = rkyv::to_bytes::<rkyv::rancor::Error>(&delta)
.expect("test setup: failed to serialize delta");
let archived = rkyv::access::<ArchivedBalancesDiff, rkyv::rancor::Error>(&bytes)
.expect("test setup: failed to access archived delta");
let mut state = MockTarget { inner: base };
let result = apply_balances_iter(&mut state, archived);
assert!(result.is_err());
let err_str = format!("{}", result.expect_err("test setup"));
assert!(
err_str.contains("exceeds i64 range"),
"Error message should mention base balance exceeding i64"
);
}
#[test]
fn test_zigzag_roundtrip() {
let values = [
0i64,
1,
-1,
2,
-2,
i32::MAX as i64,
i32::MIN as i64,
i64::MAX,
i64::MIN,
];
for &v in &values {
assert_eq!(zigzag_decode(zigzag_encode(v)), v);
}
}
#[test]
fn test_write_and_read_varint_roundtrip() {
let values = [0u64, 1, 127, 128, 255, 16383, 16384, u64::MAX];
for &v in &values {
let mut buf = Vec::new();
write_varint(v, &mut buf);
let mut cursor = 0;
let decoded = read_varint(&buf, &mut cursor).expect("valid write");
assert_eq!(decoded, v);
assert_eq!(cursor, buf.len());
}
}
#[test]
fn test_read_varint_zero() {
let buf = [0u8];
let mut cursor = 0;
assert_eq!(read_varint(&buf, &mut cursor).expect("valid zero"), 0);
assert_eq!(cursor, 1);
}
#[test]
fn test_read_varint_two_bytes() {
let buf = [0x80, 0x01];
let mut cursor = 0;
assert_eq!(read_varint(&buf, &mut cursor).expect("valid 128"), 128);
assert_eq!(cursor, 2);
}
#[test]
fn test_read_varint_max_u64() {
let buf: [u8; 10] = [0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x01];
let mut cursor = 0;
assert_eq!(read_varint(&buf, &mut cursor).expect("valid max"), u64::MAX);
assert_eq!(cursor, 10);
}
#[test]
fn test_read_varint_overflow_u64() {
let buf: [u8; 10] = [0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x02];
let mut cursor = 0;
let result = read_varint(&buf, &mut cursor);
assert!(result.is_err());
assert!(format!("{}", result.expect_err("test setup")).contains("overflows u64"));
}
#[test]
fn test_read_varint_truncated() {
let buf = [0xFF]; let mut cursor = 0;
let result = read_varint(&buf, &mut cursor);
assert!(result.is_err());
assert!(format!("{}", result.expect_err("test setup")).contains("truncated"));
}
#[test]
fn test_read_varint_exceeds_max_length() {
let buf: [u8; 10] = [0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x81];
let mut cursor = 0;
let result = read_varint(&buf, &mut cursor);
assert!(result.is_err());
assert!(
format!("{}", result.expect_err("test setup")).contains("exceeds the maximum length")
);
}
}