1use 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;
9pub 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
18pub 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
28fn 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 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
60fn 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
93pub 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
199pub 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
205pub 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
280pub 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
315pub fn format_uint64_id(id: u64) -> String {
317 if id == 0 {
318 "0".to_owned()
319 } else {
320 format_id(id)
321 }
322}
323
324pub 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
330pub 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
343pub 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 assert!(parse_price_ticks_str("65000.", "price").is_err());
394 assert!(parse_price_ticks_str("65.", "price").is_err());
395 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 assert_eq!(format_id(4), "5");
440 assert_eq!(id_to_u64("5", "order_id").unwrap(), 4);
441 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 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}