use serde::{Deserialize, Deserializer, Serialize, Serializer};
#[derive(Clone, Debug, PartialEq, Eq, Default)]
pub struct BigUintDec(pub Vec<u8>);
impl BigUintDec {
pub fn from_be_bytes(be: &[u8]) -> Self {
BigUintDec(strip_leading_zeros(be).to_vec())
}
pub fn as_be_bytes(&self) -> &[u8] {
&self.0
}
pub fn to_be_bytes_padded(&self, n: usize) -> Vec<u8> {
let b = strip_leading_zeros(&self.0);
if b.len() >= n {
b[b.len() - n..].to_vec()
} else {
let mut out = vec![0u8; n];
out[n - b.len()..].copy_from_slice(b);
out
}
}
}
pub fn be_to_decimal(be: &[u8]) -> String {
let be = strip_leading_zeros(be);
if be.is_empty() {
return "0".to_string();
}
let mut digits: Vec<u8> = vec![0];
for &byte in be {
let mut carry = byte as u32;
for d in digits.iter_mut() {
let v = (*d as u32) * 256 + carry;
*d = (v % 10) as u8;
carry = v / 10;
}
while carry > 0 {
digits.push((carry % 10) as u8);
carry /= 10;
}
}
digits.iter().rev().map(|d| (b'0' + d) as char).collect()
}
pub const MAX_DECIMAL_DIGITS: usize = 8192;
pub fn decimal_to_be(s: &str) -> Result<Vec<u8>, DecimalError> {
let s = s.trim();
if s.is_empty() {
return Err(DecimalError("empty decimal string"));
}
if s.starts_with('-') {
return Err(DecimalError("negative values are not supported"));
}
if s.len() > MAX_DECIMAL_DIGITS {
return Err(DecimalError("decimal string too long"));
}
let mut bytes: Vec<u8> = vec![0];
for ch in s.chars() {
let digit = ch.to_digit(10).ok_or(DecimalError("non-digit character"))?;
let mut carry = digit;
for b in bytes.iter_mut().rev() {
let v = (*b as u32) * 10 + carry;
*b = (v & 0xff) as u8;
carry = v >> 8;
}
while carry > 0 {
bytes.insert(0, (carry & 0xff) as u8);
carry >>= 8;
}
}
Ok(strip_leading_zeros(&bytes).to_vec())
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DecimalError(&'static str);
impl std::fmt::Display for DecimalError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "invalid decimal integer: {}", self.0)
}
}
impl std::error::Error for DecimalError {}
fn strip_leading_zeros(b: &[u8]) -> &[u8] {
let start = b.iter().position(|&x| x != 0).unwrap_or(b.len());
&b[start..]
}
impl Serialize for BigUintDec {
fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
let num = serde_json::Number::from_string_unchecked(be_to_decimal(&self.0));
num.serialize(s)
}
}
impl<'de> Deserialize<'de> for BigUintDec {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
use serde::de::Error as _;
let num = serde_json::Number::deserialize(d)?;
let be = decimal_to_be(num.as_str()).map_err(D::Error::custom)?;
Ok(BigUintDec(be))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decimal_roundtrip_small() {
for n in [0u64, 1, 9, 10, 255, 256, 65535, 1_000_000, u64::MAX] {
let be = n.to_be_bytes();
let dec = be_to_decimal(&be);
assert_eq!(dec, n.to_string());
let back = decimal_to_be(&dec).unwrap();
assert_eq!(BigUintDec(back), BigUintDec::from_be_bytes(&be));
}
}
#[test]
fn decimal_large_beyond_u64() {
let l = "7237005577332262213973186563042994240857116359379907606001950938285454250989";
let be = decimal_to_be(l).unwrap();
assert_eq!(be_to_decimal(&be), l);
}
#[test]
fn zero_is_canonical() {
assert_eq!(be_to_decimal(&[]), "0");
assert_eq!(be_to_decimal(&[0, 0, 0]), "0");
assert_eq!(decimal_to_be("0").unwrap(), Vec::<u8>::new());
assert_eq!(BigUintDec::from_be_bytes(&[0, 0, 5]).0, vec![5]);
}
#[test]
fn json_is_a_bare_number() {
let v = BigUintDec::from_be_bytes(&123456789u64.to_be_bytes());
let s = serde_json::to_string(&v).unwrap();
assert_eq!(s, "123456789");
let back: BigUintDec = serde_json::from_str(&s).unwrap();
assert_eq!(back, v);
}
#[test]
fn json_large_number_lossless() {
let l = "7237005577332262213973186563042994240857116359379907606001950938285454250989";
let v = BigUintDec(decimal_to_be(l).unwrap());
let s = serde_json::to_string(&v).unwrap();
assert_eq!(s, l);
let back: BigUintDec = serde_json::from_str(&s).unwrap();
assert_eq!(back, v);
}
#[test]
fn padded_width() {
let v = BigUintDec::from_be_bytes(&[0xab, 0xcd]);
assert_eq!(v.to_be_bytes_padded(4), vec![0, 0, 0xab, 0xcd]);
assert_eq!(v.to_be_bytes_padded(2), vec![0xab, 0xcd]);
assert_eq!(v.to_be_bytes_padded(1), vec![0xcd]);
}
#[test]
fn rejects_bad_input() {
assert!(decimal_to_be("").is_err());
assert!(decimal_to_be("-5").is_err());
assert!(decimal_to_be("12a3").is_err());
}
#[test]
fn rejects_oversized_decimal_but_accepts_paillier_sized() {
let too_long = "9".repeat(MAX_DECIMAL_DIGITS + 1);
assert_eq!(
decimal_to_be(&too_long),
Err(DecimalError("decimal string too long"))
);
let json = format!("{{\"v\": {too_long}}}");
#[derive(Deserialize)]
struct Wrap {
#[allow(dead_code)]
v: BigUintDec,
}
assert!(serde_json::from_str::<Wrap>(&json).is_err());
let mut paillier_sized = String::from("9");
paillier_sized.push_str(&"7".repeat(616));
assert_eq!(paillier_sized.len(), 617);
let be = decimal_to_be(&paillier_sized).unwrap();
assert_eq!(be_to_decimal(&be), paillier_sized);
let at_cap = "1".repeat(MAX_DECIMAL_DIGITS);
assert!(decimal_to_be(&at_cap).is_ok());
}
}