use sha2::{Digest, Sha256};
use thiserror::Error;
use crate::filter::Filter;
use crate::message::SubscriptionId;
use crate::util::hex;
const FINGERPRINT_BYTES: usize = 16;
const ID_BYTES: usize = 32;
const INFINITY_TIMESTAMP: u64 = u64::MAX;
const RESERVED_TIMESTAMP_INFINITY_OFFSET: u64 = 0;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct NegProtocolVersion(pub u8);
impl NegProtocolVersion {
pub const V1: Self = Self(0x61);
#[must_use]
pub const fn from_byte(byte: u8) -> Option<Self> {
if byte >= 0x60 && byte < 0x70 {
Some(Self(byte))
} else {
None
}
}
#[must_use]
pub const fn as_byte(self) -> u8 {
self.0
}
#[must_use]
pub const fn version(self) -> u8 {
self.0.wrapping_sub(0x60)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct NegItem {
pub timestamp: u64,
pub id: [u8; ID_BYTES],
}
impl NegItem {
#[must_use]
pub const fn new(timestamp: u64, id: [u8; ID_BYTES]) -> Self {
Self { timestamp, id }
}
#[must_use]
pub const fn infinity() -> Self {
Self {
timestamp: INFINITY_TIMESTAMP,
id: [0u8; ID_BYTES],
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NegBound {
pub timestamp: u64,
pub id_prefix: Vec<u8>,
}
impl NegBound {
#[must_use]
pub const fn infinity() -> Self {
Self {
timestamp: INFINITY_TIMESTAMP,
id_prefix: Vec::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum NegRangeMode {
Skip,
Fingerprint([u8; FINGERPRINT_BYTES]),
IdList(Vec<[u8; ID_BYTES]>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NegRange {
pub upper_bound: NegBound,
pub mode: NegRangeMode,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NegPayload {
pub version: NegProtocolVersion,
pub ranges: Vec<NegRange>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NegOpen {
pub subscription_id: SubscriptionId,
pub filter: Filter,
pub payload: NegPayload,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum NegMessage {
Open(Box<NegOpen>),
Msg {
subscription_id: SubscriptionId,
payload: NegPayload,
},
Close {
subscription_id: SubscriptionId,
},
Err {
subscription_id: SubscriptionId,
reason: String,
},
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum NegentropyError {
#[error("unexpected end of input while decoding varint")]
UnexpectedEof,
#[error("varint exceeds 10 bytes")]
VarintOverflow,
#[error("buffer ended {expected} bytes before payload completed")]
PayloadTruncated {
expected: usize,
},
#[error("unknown range mode {0}")]
UnknownRangeMode(u64),
#[error("unsupported Negentropy protocol version 0x{0:02x}")]
UnsupportedVersion(u8),
#[error("hex decode failure: {0}")]
Hex(String),
}
fn write_varint(value: u64, out: &mut Vec<u8>) {
let mut value = value;
let mut tmp: Vec<u8> = Vec::with_capacity(10);
loop {
let byte = u8::try_from(value & 0x7f).unwrap_or(0);
tmp.push(byte);
value >>= 7;
if value == 0 {
break;
}
}
for (i, byte) in tmp.iter().rev().enumerate() {
let with_continuation = if i + 1 < tmp.len() {
*byte | 0x80
} else {
*byte
};
out.push(with_continuation);
}
}
fn read_varint(buf: &[u8], cursor: &mut usize) -> Result<u64, NegentropyError> {
let mut value: u64 = 0;
for _ in 0..10 {
let byte = *buf.get(*cursor).ok_or(NegentropyError::UnexpectedEof)?;
*cursor += 1;
value = value
.checked_shl(7)
.ok_or(NegentropyError::VarintOverflow)?
| u64::from(byte & 0x7f);
if byte & 0x80 == 0 {
return Ok(value);
}
}
Err(NegentropyError::VarintOverflow)
}
fn encode_bound(bound: &NegBound, prev_timestamp: &mut u64, out: &mut Vec<u8>) {
let encoded_ts = if bound.timestamp == INFINITY_TIMESTAMP {
RESERVED_TIMESTAMP_INFINITY_OFFSET
} else {
bound
.timestamp
.saturating_sub(*prev_timestamp)
.saturating_add(1)
};
write_varint(encoded_ts, out);
if bound.timestamp != INFINITY_TIMESTAMP {
*prev_timestamp = bound.timestamp;
}
let len = bound.id_prefix.len().min(ID_BYTES);
write_varint(len as u64, out);
out.extend_from_slice(bound.id_prefix.get(..len).unwrap_or(&[]));
}
fn decode_bound(
buf: &[u8],
cursor: &mut usize,
prev_timestamp: &mut u64,
) -> Result<NegBound, NegentropyError> {
let ts_field = read_varint(buf, cursor)?;
let timestamp = if ts_field == RESERVED_TIMESTAMP_INFINITY_OFFSET {
INFINITY_TIMESTAMP
} else {
let value = prev_timestamp.saturating_add(ts_field.saturating_sub(1));
*prev_timestamp = value;
value
};
let len_u64 = read_varint(buf, cursor)?;
let len = usize::try_from(len_u64).map_err(|_| NegentropyError::VarintOverflow)?;
if len > ID_BYTES {
return Err(NegentropyError::PayloadTruncated { expected: len });
}
let chunk = read_chunk(buf, cursor, len)?;
Ok(NegBound {
timestamp,
id_prefix: chunk.to_vec(),
})
}
fn encode_range(range: &NegRange, prev_timestamp: &mut u64, out: &mut Vec<u8>) {
encode_bound(&range.upper_bound, prev_timestamp, out);
match &range.mode {
NegRangeMode::Skip => write_varint(0, out),
NegRangeMode::Fingerprint(fp) => {
write_varint(1, out);
out.extend_from_slice(fp);
}
NegRangeMode::IdList(ids) => {
write_varint(2, out);
write_varint(ids.len() as u64, out);
for id in ids {
out.extend_from_slice(id);
}
}
}
}
fn read_chunk<'a>(
buf: &'a [u8],
cursor: &mut usize,
len: usize,
) -> Result<&'a [u8], NegentropyError> {
let end = cursor
.checked_add(len)
.ok_or(NegentropyError::VarintOverflow)?;
let chunk = buf.get(*cursor..end);
chunk.map_or_else(
|| {
Err(NegentropyError::PayloadTruncated {
expected: end.saturating_sub(buf.len()),
})
},
|chunk| {
*cursor = end;
Ok(chunk)
},
)
}
fn decode_range(
buf: &[u8],
cursor: &mut usize,
prev_timestamp: &mut u64,
) -> Result<NegRange, NegentropyError> {
let upper_bound = decode_bound(buf, cursor, prev_timestamp)?;
let mode = read_varint(buf, cursor)?;
let mode = match mode {
0 => NegRangeMode::Skip,
1 => {
let chunk = read_chunk(buf, cursor, FINGERPRINT_BYTES)?;
let mut fp = [0u8; FINGERPRINT_BYTES];
fp.copy_from_slice(chunk);
NegRangeMode::Fingerprint(fp)
}
2 => {
let count_u64 = read_varint(buf, cursor)?;
let count = usize::try_from(count_u64).map_err(|_| NegentropyError::VarintOverflow)?;
let mut ids = Vec::with_capacity(count);
for _ in 0..count {
let chunk = read_chunk(buf, cursor, ID_BYTES)?;
let mut id = [0u8; ID_BYTES];
id.copy_from_slice(chunk);
ids.push(id);
}
NegRangeMode::IdList(ids)
}
other => return Err(NegentropyError::UnknownRangeMode(other)),
};
Ok(NegRange { upper_bound, mode })
}
#[must_use]
pub fn encode_payload(payload: &NegPayload) -> Vec<u8> {
let mut out = Vec::with_capacity(1 + payload.ranges.len() * 32);
out.push(payload.version.as_byte());
let mut prev_timestamp: u64 = 0;
for range in &payload.ranges {
encode_range(range, &mut prev_timestamp, &mut out);
}
out
}
pub fn decode_payload(buf: &[u8]) -> Result<NegPayload, NegentropyError> {
let (version_byte, rest) = buf.split_first().ok_or(NegentropyError::UnexpectedEof)?;
let version = NegProtocolVersion::from_byte(*version_byte)
.ok_or(NegentropyError::UnsupportedVersion(*version_byte))?;
let mut cursor = 0;
let mut prev_timestamp: u64 = 0;
let mut ranges = Vec::new();
while cursor < rest.len() {
ranges.push(decode_range(rest, &mut cursor, &mut prev_timestamp)?);
}
Ok(NegPayload { version, ranges })
}
#[must_use]
pub fn encode_payload_hex(payload: &NegPayload) -> String {
hex::encode(encode_payload(payload))
}
pub fn decode_payload_hex(hex_str: &str) -> Result<NegPayload, NegentropyError> {
let bytes = hex::decode(hex_str).map_err(|e| NegentropyError::Hex(e.to_string()))?;
decode_payload(&bytes)
}
#[must_use]
pub fn fingerprint(ids: &[[u8; ID_BYTES]]) -> [u8; FINGERPRINT_BYTES] {
let mut sum = [0u8; ID_BYTES];
for id in ids {
let mut carry: u16 = 0;
for (sum_byte, id_byte) in sum.iter_mut().zip(id.iter()) {
let total = u16::from(*sum_byte) + u16::from(*id_byte) + carry;
*sum_byte = u8::try_from(total & 0xff).unwrap_or(0);
carry = total >> 8;
}
}
let mut hasher = Sha256::new();
hasher.update(sum);
let mut count_buf = Vec::new();
write_varint(ids.len() as u64, &mut count_buf);
hasher.update(&count_buf);
let digest = hasher.finalize();
let mut out = [0u8; FINGERPRINT_BYTES];
for (slot, byte) in out.iter_mut().zip(digest.iter()) {
*slot = *byte;
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn varint_round_trip() {
for &value in &[0u64, 1, 127, 128, 300, 1_000_000, u64::MAX / 2] {
let mut buf = Vec::new();
write_varint(value, &mut buf);
let mut cursor = 0;
assert_eq!(read_varint(&buf, &mut cursor).unwrap(), value);
assert_eq!(cursor, buf.len());
}
}
#[test]
fn payload_round_trip() {
let payload = NegPayload {
version: NegProtocolVersion::V1,
ranges: vec![
NegRange {
upper_bound: NegBound {
timestamp: 1_700_000_000,
id_prefix: vec![0xab, 0xcd],
},
mode: NegRangeMode::Fingerprint([0xff; FINGERPRINT_BYTES]),
},
NegRange {
upper_bound: NegBound::infinity(),
mode: NegRangeMode::Skip,
},
],
};
let bytes = encode_payload(&payload);
let parsed = decode_payload(&bytes).unwrap();
assert_eq!(parsed, payload);
}
#[test]
fn id_list_payload_round_trip() {
let ids = vec![[0x11; 32], [0x22; 32]];
let payload = NegPayload {
version: NegProtocolVersion::V1,
ranges: vec![NegRange {
upper_bound: NegBound::infinity(),
mode: NegRangeMode::IdList(ids.clone()),
}],
};
let hex_str = encode_payload_hex(&payload);
let parsed = decode_payload_hex(&hex_str).unwrap();
match &parsed.ranges[0].mode {
NegRangeMode::IdList(decoded) => assert_eq!(decoded, &ids),
other => panic!("unexpected mode {other:?}"),
}
}
#[test]
fn unsupported_version_is_rejected() {
let bytes = vec![0x10];
assert!(matches!(
decode_payload(&bytes),
Err(NegentropyError::UnsupportedVersion(0x10))
));
}
#[test]
fn fingerprint_is_stable() {
let ids = vec![[1u8; 32], [2u8; 32]];
let fp = fingerprint(&ids);
let fp_again = fingerprint(&ids);
assert_eq!(fp, fp_again);
let fp_other = fingerprint(&[[3u8; 32]]);
assert_ne!(fp, fp_other);
}
}