Skip to main content

polyester/codecs/
scalars.rs

1//! Strict decimal ↔ scaled-integer codecs.
2
3use crate::errors::{Error, Result};
4use rust_decimal::Decimal;
5use std::str::FromStr;
6
7pub const PRICE_TICK_SCALE: u32 = 6;
8pub const LEDGER_SCALE: u32 = 18;
9/// Maximum accepted quantity/ledger scale for public formatters and catalog hydration.
10///
11/// Values above this are rejected with [`Error::Validation`] instead of allocating
12/// pathological padding or panicking in `format!` width formatting (scale ≥ 65535).
13pub const MAX_PROTOCOL_SCALE: u32 = 36;
14pub const INT64_MAX: i128 = i64::MAX as i128;
15pub const INT64_MIN: i128 = i64::MIN as i128;
16pub const UINT64_MAX: u128 = u64::MAX as u128;
17
18/// Validate a caller/catalog scale before padding or allocation.
19pub fn validate_protocol_scale(scale: u32) -> Result<()> {
20    if scale > MAX_PROTOCOL_SCALE {
21        return Err(Error::validation(format!(
22            "scale {scale} exceeds maximum protocol scale {MAX_PROTOCOL_SCALE}"
23        )));
24    }
25    Ok(())
26}
27
28/// Strict non-negative decimal: digits with optional fractional part.
29fn decimal_string_from_input(raw: &str, field_name: &str) -> Result<String> {
30    let text = raw.trim();
31    if text.is_empty() || !is_strict_decimal(text) {
32        return Err(Error::validation(format!(
33            "{field_name} must be a valid decimal string"
34        )));
35    }
36    Ok(text.to_owned())
37}
38
39fn decimal_string_from_decimal(raw: Decimal, field_name: &str) -> Result<String> {
40    if raw.is_sign_negative() {
41        return Err(Error::validation(format!(
42            "{field_name} must be non-negative"
43        )));
44    }
45    let text = format!("{raw}");
46    // Normalize trailing zeros from Decimal display.
47    let text = if let Some((h, t)) = text.split_once('.') {
48        let t = t.trim_end_matches('0');
49        if t.is_empty() {
50            h.to_owned()
51        } else {
52            format!("{h}.{t}")
53        }
54    } else {
55        text
56    };
57    decimal_string_from_input(&text, field_name)
58}
59
60/// Strict non-negative decimal form: `digits` or `digits.digits`.
61///
62/// Matches TypeScript/Go/Python (`^\d+(?:\.\d+)?$`): a trailing bare `.`
63/// (e.g. `"65000."`) is rejected. Callers may trim surrounding whitespace
64/// before invoking this (same as TS `value.trim()`).
65fn is_strict_decimal(text: &str) -> bool {
66    let mut chars = text.chars();
67    let Some(first) = chars.next() else {
68        return false;
69    };
70    if !first.is_ascii_digit() {
71        return false;
72    }
73    let mut saw_dot = false;
74    let mut frac_digits = 0usize;
75    for c in chars {
76        if c == '.' {
77            if saw_dot {
78                return false;
79            }
80            saw_dot = true;
81            continue;
82        }
83        if !c.is_ascii_digit() {
84            return false;
85        }
86        if saw_dot {
87            frac_digits += 1;
88        }
89    }
90    !saw_dot || frac_digits > 0
91}
92
93/// Strict decimal→scaled. Never rounds; excess fractional digits fail.
94pub fn try_decimal_to_scaled(decimal: &str, scale: u32) -> std::result::Result<i128, &'static str> {
95    if scale > MAX_PROTOCOL_SCALE {
96        return Err("scale");
97    }
98    let raw = decimal.trim();
99    if !is_strict_decimal(raw) {
100        return Err("invalid");
101    }
102    let (int_part, frac_part) = match raw.split_once('.') {
103        Some((i, f)) => (i, f),
104        None => (raw, ""),
105    };
106    if frac_part.len() as u32 > scale {
107        return Err("precision");
108    }
109    let mut digits = String::with_capacity(int_part.len() + scale as usize);
110    digits.push_str(int_part);
111    digits.push_str(frac_part);
112    let pad = scale as usize - frac_part.len();
113    digits.extend(std::iter::repeat_n('0', pad));
114    if digits.is_empty() {
115        digits.push('0');
116    }
117    digits.parse::<i128>().map_err(|_| "invalid")
118}
119
120pub fn decimal_to_scaled_str(raw: &str, scale: u32, field_name: &str) -> Result<i128> {
121    validate_protocol_scale(scale)?;
122    let text = decimal_string_from_input(raw, field_name)?;
123    match try_decimal_to_scaled(&text, scale) {
124        Ok(v) => Ok(v),
125        Err("precision") => Err(Error::validation(format!(
126            "{field_name} supports at most {scale} decimal places: {text}"
127        ))),
128        Err("scale") => Err(Error::validation(format!(
129            "{field_name} scale {scale} exceeds maximum protocol scale {MAX_PROTOCOL_SCALE}"
130        ))),
131        Err(_) => Err(Error::validation(format!(
132            "{field_name} must be a valid decimal string"
133        ))),
134    }
135}
136
137pub fn decimal_to_scaled(raw: Decimal, scale: u32, field_name: &str) -> Result<i128> {
138    let text = decimal_string_from_decimal(raw, field_name)?;
139    decimal_to_scaled_str(&text, scale, field_name)
140}
141
142pub fn parse_price_ticks_str(raw: &str, field_name: &str) -> Result<i64> {
143    let scaled = decimal_to_scaled_str(raw, PRICE_TICK_SCALE, field_name)?;
144    if scaled < 0 {
145        return Err(Error::validation(format!(
146            "{field_name} must be non-negative"
147        )));
148    }
149    if scaled > INT64_MAX {
150        return Err(Error::validation(format!(
151            "{field_name} exceeds int64 range"
152        )));
153    }
154    Ok(scaled as i64)
155}
156
157pub fn parse_price_ticks(raw: Decimal, field_name: &str) -> Result<i64> {
158    let scaled = decimal_to_scaled(raw, PRICE_TICK_SCALE, field_name)?;
159    if scaled > INT64_MAX {
160        return Err(Error::validation(format!(
161            "{field_name} exceeds int64 range"
162        )));
163    }
164    Ok(scaled as i64)
165}
166
167pub fn format_price_ticks(ticks: i64) -> String {
168    format_scaled(ticks as i128, PRICE_TICK_SCALE)
169        .expect("PRICE_TICK_SCALE is within MAX_PROTOCOL_SCALE")
170}
171
172pub fn parse_qty_scaled_str(raw: &str, scale: u32, field_name: &str) -> Result<i64> {
173    let scaled = decimal_to_scaled_str(raw, scale, field_name)?;
174    if scaled <= 0 {
175        return Err(Error::validation(format!("{field_name} must be positive")));
176    }
177    if scale != 18 && scaled > INT64_MAX {
178        return Err(Error::validation(format!(
179            "{field_name} exceeds int64 range"
180        )));
181    }
182    if scaled > INT64_MAX {
183        return Err(Error::validation(format!(
184            "{field_name} exceeds int64 range"
185        )));
186    }
187    Ok(scaled as i64)
188}
189
190pub fn parse_qty_scaled(raw: Decimal, scale: u32, field_name: &str) -> Result<i64> {
191    let text = decimal_string_from_decimal(raw, field_name)?;
192    parse_qty_scaled_str(&text, scale, field_name)
193}
194
195pub fn format_qty_scaled(qty_scaled: i64, scale: u32) -> Result<String> {
196    format_scaled(qty_scaled as i128, scale)
197}
198
199/// Format a smaller ledger quantity integer by `10^scale`.
200pub fn format_ledger_u64(value: u64, scale: u32) -> Result<String> {
201    let scale = if scale == 0 { LEDGER_SCALE } else { scale };
202    format_scaled(value as i128, scale)
203}
204
205/// Format a full-width unsigned ledger integer string by `10^scale`.
206///
207/// Balance models expose protobuf `u128` values as decimal strings so no
208/// precision is lost. This helper formats those strings without narrowing to
209/// `u64` or a floating-point type.
210pub fn format_ledger_u128(value: &str, scale: u32) -> Result<String> {
211    let digits = value.trim();
212    if digits.is_empty() || !digits.bytes().all(|byte| byte.is_ascii_digit()) {
213        return Err(Error::validation(
214            "ledger value must be an unsigned decimal integer string",
215        ));
216    }
217    let digits = digits.trim_start_matches('0');
218    let digits = if digits.is_empty() { "0" } else { digits };
219    const U128_MAX_DECIMAL: &str = "340282366920938463463374607431768211455";
220    if digits.len() > U128_MAX_DECIMAL.len()
221        || (digits.len() == U128_MAX_DECIMAL.len() && digits > U128_MAX_DECIMAL)
222    {
223        return Err(Error::validation("ledger value exceeds u128 range"));
224    }
225    let scale = if scale == 0 { LEDGER_SCALE } else { scale };
226    validate_protocol_scale(scale)?;
227    if scale == 0 {
228        return Ok(digits.to_owned());
229    }
230    let width = (scale as usize)
231        .checked_add(1)
232        .ok_or_else(|| Error::validation("scale width overflow"))?;
233    let padded = format!("{digits:0>width$}");
234    let (head, tail) = padded.split_at(padded.len() - scale as usize);
235    let head = head.trim_start_matches('0');
236    let head = if head.is_empty() { "0" } else { head };
237    let tail = tail.trim_end_matches('0');
238    Ok(if tail.is_empty() {
239        head.to_owned()
240    } else {
241        format!("{head}.{tail}")
242    })
243}
244
245fn format_scaled(value: i128, scale: u32) -> Result<String> {
246    validate_protocol_scale(scale)?;
247    if scale == 0 {
248        return Ok(value.to_string());
249    }
250    let neg = value < 0;
251    let digits = value.abs().to_string();
252    let width = (scale as usize)
253        .checked_add(1)
254        .ok_or_else(|| Error::validation("scale width overflow"))?;
255    let padded = format!("{digits:0>width$}");
256    let (head, tail) = padded.split_at(padded.len() - scale as usize);
257    let head = head.trim_start_matches('0');
258    let head = if head.is_empty() { "0" } else { head };
259    let tail = tail.trim_end_matches('0');
260    let raw = if tail.is_empty() {
261        head.to_owned()
262    } else {
263        format!("{head}.{tail}")
264    };
265    Ok(if neg { format!("-{raw}") } else { raw })
266}
267
268fn base58_to_u64(value: &str, label: &str) -> Result<u64> {
269    let bytes = bs58::decode(value)
270        .into_vec()
271        .map_err(|_| Error::validation(format!("{label} must be base58 or decimal uint64")))?;
272    if bytes.len() > 8 {
273        return Err(Error::validation(format!("{label} exceeds uint64 range")));
274    }
275    let mut buf = [0u8; 8];
276    buf[8 - bytes.len()..].copy_from_slice(&bytes);
277    Ok(u64::from_be_bytes(buf))
278}
279
280/// Parse a public id that may be base58 or decimal.
281///
282/// All-digit strings are ambiguous: `format_id(4)` is `"5"`, which is also a
283/// valid decimal. Prefer the canonical base58 decode when `format_id(b) == input`;
284/// otherwise treat the value as decimal.
285pub fn id_to_u64(value: &str, label: &str) -> Result<u64> {
286    let value = value.trim();
287    if value.is_empty() {
288        return Err(Error::validation(format!(
289            "{label} must be base58 or decimal uint64"
290        )));
291    }
292    if value.chars().all(|c| c.is_ascii_digit()) {
293        let decimal = value
294            .parse::<u64>()
295            .map_err(|_| Error::validation(format!("{label} exceeds uint64 range")))?;
296        if let Ok(canonical) = base58_to_u64(value, label)
297            && format_id(canonical) == value
298        {
299            return Ok(canonical);
300        }
301        return Ok(decimal);
302    }
303    base58_to_u64(value, label)
304}
305
306pub fn format_id(id: u64) -> String {
307    if id == 0 {
308        return bs58::encode([0u8]).into_string();
309    }
310    let bytes = id.to_be_bytes();
311    let start = bytes.iter().position(|&b| b != 0).unwrap_or(7);
312    bs58::encode(&bytes[start..]).into_string()
313}
314
315/// Format a uint64 id as base58, or `"0"` when zero (order/trade ids).
316pub fn format_uint64_id(id: u64) -> String {
317    if id == 0 {
318        "0".to_owned()
319    } else {
320        format_id(id)
321    }
322}
323
324/// Format protobuf U128 hi/lo as a decimal string.
325pub fn u128_to_str(hi: u64, lo: u64) -> String {
326    let value = (u128::from(hi) << 64) | u128::from(lo);
327    value.to_string()
328}
329
330/// Encode a non-negative scaled integer as protobuf `U128` (hi/lo).
331pub fn i128_to_u128(n: i128) -> Result<crate::proto::polyester::r#type::v1::U128> {
332    if n < 0 {
333        return Err(Error::validation("u128 value must be non-negative"));
334    }
335    let value = n as u128;
336    Ok(crate::proto::polyester::r#type::v1::U128 {
337        hi: (value >> 64) as u64,
338        lo: value as u64,
339        ..Default::default()
340    })
341}
342
343/// Encode a `u128` as protobuf `U128` (hi/lo).
344pub fn u128_to_proto(value: u128) -> crate::proto::polyester::r#type::v1::U128 {
345    crate::proto::polyester::r#type::v1::U128 {
346        hi: (value >> 64) as u64,
347        lo: value as u64,
348        ..Default::default()
349    }
350}
351
352pub fn parse_decimal_input(raw: &str) -> Result<Decimal> {
353    Decimal::from_str(raw.trim()).map_err(|_| Error::validation("invalid decimal".to_owned()))
354}
355
356#[cfg(test)]
357mod tests {
358    use super::*;
359
360    #[test]
361    fn price_ticks_round_trip() {
362        let ticks = parse_price_ticks_str("1.5", "price").unwrap();
363        assert_eq!(ticks, 1_500_000);
364        assert_eq!(format_price_ticks(ticks), "1.5");
365    }
366
367    #[test]
368    fn reject_excess_precision() {
369        let err = parse_price_ticks_str("1.1234567", "price").unwrap_err();
370        assert!(err.to_string().contains("at most 6"));
371    }
372
373    #[test]
374    fn qty_positive() {
375        assert!(parse_qty_scaled_str("0", 8, "qty").is_err());
376        assert_eq!(parse_qty_scaled_str("0.00000001", 8, "qty").unwrap(), 1);
377    }
378
379    #[test]
380    fn qty_rejects_excess_precision() {
381        let err = parse_qty_scaled_str("1.123456789", 8, "qty").unwrap_err();
382        assert!(err.to_string().contains("at most") || err.to_string().contains("precision"));
383    }
384
385    #[test]
386    fn price_rejects_negative_string() {
387        assert!(parse_price_ticks_str("-1", "price").is_err());
388    }
389
390    #[test]
391    fn price_rejects_trailing_dot_and_accepts_trimmed_whitespace() {
392        // TS/Go/Python reject bare trailing dots via `^\d+(?:\.\d+)?$`.
393        assert!(parse_price_ticks_str("65000.", "price").is_err());
394        assert!(parse_price_ticks_str("65.", "price").is_err());
395        // Leading/trailing whitespace is trimmed (TS `value.trim()` parity).
396        assert_eq!(
397            parse_price_ticks_str(" 65000", "price").unwrap(),
398            65_000_000_000
399        );
400        assert_eq!(
401            parse_price_ticks_str("65000 ", "price").unwrap(),
402            65_000_000_000
403        );
404        assert_eq!(
405            parse_price_ticks_str("65000.0", "price").unwrap(),
406            65_000_000_000
407        );
408    }
409
410    #[test]
411    fn format_qty_scaled_round_trip() {
412        assert_eq!(format_qty_scaled(1_000_000, 8).unwrap(), "0.01");
413    }
414
415    #[test]
416    fn format_rejects_scale_above_max_protocol_scale() {
417        assert!(format_qty_scaled(1, MAX_PROTOCOL_SCALE).is_ok());
418        assert!(format_qty_scaled(1, MAX_PROTOCOL_SCALE + 1).is_err());
419        assert!(format_qty_scaled(1, 65535).is_err());
420        assert!(format_ledger_u64(1, 65535).is_err());
421        assert!(format_ledger_u128("1", 65535).is_err());
422    }
423
424    #[test]
425    fn format_full_width_ledger_integer_string() {
426        assert_eq!(
427            format_ledger_u128("1000000000000000001", 18).unwrap(),
428            "1.000000000000000001"
429        );
430        assert_eq!(format_ledger_u128("000000", 18).unwrap(), "0");
431        assert!(format_ledger_u128("-1", 18).is_err());
432        assert!(format_ledger_u128("1.5", 18).is_err());
433        assert!(format_ledger_u128("340282366920938463463374607431768211456", 18).is_err());
434    }
435
436    #[test]
437    fn id_round_trip_prefers_canonical_base58_for_all_digit_encodings() {
438        // format_id(4) == "5"; decimal parse would wrongly yield 5.
439        assert_eq!(format_id(4), "5");
440        assert_eq!(id_to_u64("5", "order_id").unwrap(), 4);
441        // Zero has its own canonical base58 encoding and must not alias id 1.
442        assert_eq!(format_id(0), "1");
443        assert_eq!(format_id(1), "2");
444        assert_ne!(format_id(0), format_id(1));
445        for id in 0u64..200 {
446            let encoded = format_id(id);
447            assert_eq!(
448                id_to_u64(&encoded, "id").unwrap(),
449                id,
450                "round-trip failed for id={id} encoded={encoded}"
451            );
452        }
453    }
454
455    #[test]
456    fn format_uint64_id_preserves_wire_zero_as_decimal_zero() {
457        assert_eq!(format_uint64_id(0), "0");
458        assert_eq!(format_uint64_id(1), "2");
459    }
460
461    #[test]
462    fn id_to_u64_still_accepts_non_canonical_decimal() {
463        // "10" is not the canonical encoding of any small id via format_id.
464        assert_ne!(format_id(10), "10");
465        assert_eq!(id_to_u64("10", "order_id").unwrap(), 10);
466        assert_eq!(id_to_u64("100", "order_id").unwrap(), 100);
467    }
468
469    #[test]
470    fn id_to_u64_rejects_invalid() {
471        assert!(id_to_u64("", "id").is_err());
472        assert!(id_to_u64("not a trigger id", "id").is_err());
473    }
474}