Skip to main content

ternaria_arith/
balanced.rs

1//! The balanced ternary numeric types: [`Tryte`], [`Word`], [`DoubleWord`].
2//!
3//! # Representation
4//!
5//! Each type stores its numeric value in a native host integer rather than a
6//! packed trit string (D-03). Balanced ternary addition needs base-3 carry
7//! propagation, which a host add does not perform, so the fast path is to
8//! unpack, operate natively, and repack.
9//!
10//! BCT, two bits per trit, is the storage and interchange format, reached via
11//! [`Word::to_bct`] and [`Word::from_bct`]. Guest memory holds BCT; registers
12//! hold values.
13//!
14//! Widths and their host types (D-02):
15//!
16//! | Type | Trits | ~ bits | Range | Host |
17//! |------|-------|--------|-------|------|
18//! | [`Tryte`]      |  9 |  14.3 | +/-9,841 | `i32` |
19//! | [`Word`]       | 27 |  42.8 | +/-3.81*10^12 | `i64` |
20//! | [`DoubleWord`] | 54 |  85.6 | +/-2.91*10^25 | `i128` |
21
22use crate::trit::Trit;
23use core::fmt;
24use core::str::FromStr;
25
26/// 3^n as an `i128`. Exact for n <= 80.
27///
28/// Most calls are in const context ([`Word::MAX`] and friends) and cost
29/// nothing at runtime. The exceptions are the trit shifts, which pass a runtime
30/// `k` - hence the inline hint. If profiling ever shows this loop mattering,
31/// the fix is a 81-entry lookup table, not a cleverer loop.
32#[inline]
33pub(crate) const fn pow3(n: u32) -> i128 {
34    let mut r: i128 = 1;
35    let mut i = 0;
36    while i < n {
37        r *= 3;
38        i += 1;
39    }
40    r
41}
42
43/// Error from parsing a balanced ternary literal.
44#[derive(Clone, Copy, PartialEq, Eq, Debug)]
45pub enum ParseTritsError {
46    /// A character outside the accepted digit set.
47    ///
48    /// Accepted are `T`, `t` and `-` for -1; `0` for zero; `1` and `+` for +1.
49    /// Underscores are stripped before parsing and never reach here.
50    BadDigit(char),
51    /// No digits at all.
52    Empty,
53    /// More digits than the type has trits.
54    TooLong {
55        /// How many digits the literal actually had.
56        got: usize,
57        /// How many trits the target type provides.
58        max: usize,
59    },
60}
61
62impl fmt::Display for ParseTritsError {
63    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64        match self {
65            ParseTritsError::BadDigit(c) => write!(f, "invalid balanced ternary digit {c:?}"),
66            ParseTritsError::Empty => f.write_str("empty balanced ternary literal"),
67            ParseTritsError::TooLong { got, max } => {
68                write!(f, "{got} digits exceeds the {max} trits available")
69            }
70        }
71    }
72}
73
74/// Error from decoding a BCT-packed value.
75#[derive(Clone, Copy, PartialEq, Eq, Debug)]
76pub struct NatError {
77    /// Index of the first trit position holding the `10` NaT pattern.
78    pub position: u32,
79}
80
81impl fmt::Display for NatError {
82    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
83        write!(f, "NaT (not-a-trit) at trit position {}", self.position)
84    }
85}
86
87macro_rules! balanced_type {
88    ($name:ident, $trits:literal, $host:ty, $bct:ty, $what:literal) => {
89        #[doc = concat!("A ", $what, ", ", stringify!($trits), " balanced trits.")]
90        ///
91        /// Stores the numeric value in the host integer; see the module docs.
92        /// Every value of this type is in range by construction, so `value()`
93        /// never returns something out of bounds.
94        #[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Default)]
95        pub struct $name($host);
96
97        impl $name {
98            /// Number of trits in this type.
99            pub const TRITS: u32 = $trits;
100
101            /// 3^TRITS, the number of distinct values.
102            pub const MODULUS: i128 = pow3($trits);
103
104            /// The largest representable value, (3^n-1)/2.
105            pub const MAX: Self = Self(((pow3($trits) - 1) / 2) as $host);
106
107            /// The smallest representable value, -(3^n-1)/2.
108            ///
109            /// The range is symmetric, so there is no value whose negation
110            /// overflows (D-01).
111            pub const MIN: Self = Self(-(((pow3($trits) - 1) / 2) as $host));
112
113            /// Zero.
114            pub const ZERO: Self = Self(0);
115            /// One.
116            pub const ONE: Self = Self(1);
117            /// Negative one.
118            pub const NEG_ONE: Self = Self(-1);
119
120            /// True when `v` is representable in this many trits.
121            #[inline]
122            pub const fn in_range(v: $host) -> bool {
123                v >= Self::MIN.0 && v <= Self::MAX.0
124            }
125
126            /// Wraps a host value, or `None` if out of range.
127            #[inline]
128            pub const fn new(v: $host) -> Option<Self> {
129                if Self::in_range(v) {
130                    Some(Self(v))
131                } else {
132                    None
133                }
134            }
135
136            /// Wraps a host value, panicking if out of range.
137            ///
138            /// For constants and tests where the range is known statically.
139            #[inline]
140            #[track_caller]
141            pub const fn from_value(v: $host) -> Self {
142                Self::new(v).expect(concat!(stringify!($name), ": value out of range"))
143            }
144
145            /// Reduces an `i128` into range modulo 3^n. Always succeeds.
146            #[inline]
147            pub const fn from_i128_wrapping(v: i128) -> Self {
148                let m = Self::MODULUS;
149                let mut r = v % m;
150                let max = (m - 1) / 2;
151                if r > max {
152                    r -= m;
153                } else if r < -max {
154                    r += m;
155                }
156                Self(r as $host)
157            }
158
159            /// Reduces an `i128` into range, or `None` if it does not fit.
160            #[inline]
161            pub const fn from_i128(v: i128) -> Option<Self> {
162                let max = (Self::MODULUS - 1) / 2;
163                if v >= -max && v <= max {
164                    Some(Self(v as $host))
165                } else {
166                    None
167                }
168            }
169
170            /// The numeric value.
171            #[inline]
172            pub const fn value(self) -> $host {
173                self.0
174            }
175
176            /// The numeric value widened to `i128`.
177            #[inline]
178            pub const fn to_i128(self) -> i128 {
179                self.0 as i128
180            }
181
182            /// The trits, least significant first.
183            pub const fn trits(self) -> [Trit; $trits] {
184                let mut out = [Trit::Zero; $trits];
185                let mut n = self.0;
186                let mut i = 0;
187                while i < $trits {
188                    // Rust's `%` truncates toward zero, so for n < 0 the
189                    // remainder lands in {-2,-1,0}. Both 2 and -2 are outside
190                    // the digit set and borrow from the next position.
191                    let mut r = n % 3;
192                    n /= 3;
193                    if r == 2 {
194                        r = -1;
195                        n += 1;
196                    } else if r == -2 {
197                        r = 1;
198                        n -= 1;
199                    }
200                    out[i] = match r {
201                        -1 => Trit::Neg,
202                        1 => Trit::Pos,
203                        _ => Trit::Zero,
204                    };
205                    i += 1;
206                }
207                out
208            }
209
210            /// Rebuilds a value from trits, least significant first.
211            pub const fn from_trits(trits: [Trit; $trits]) -> Self {
212                let mut acc: i128 = 0;
213                let mut p: i128 = 1;
214                let mut i = 0;
215                while i < $trits {
216                    acc += (trits[i] as i8 as i128) * p;
217                    p *= 3;
218                    i += 1;
219                }
220                Self(acc as $host)
221            }
222
223            /// The trit at `index`, counting from the least significant.
224            ///
225            /// # Panics
226            /// If `index >= TRITS`.
227            #[inline]
228            #[track_caller]
229            pub const fn trit(self, index: u32) -> Trit {
230                assert!(index < $trits, "trit index out of range");
231                self.trits()[index as usize]
232            }
233
234            // ---- BCT storage format (D-03) ----
235
236            /// Packs into BCT: two bits per trit, least significant trit in the
237            /// low bits.
238            pub const fn to_bct(self) -> $bct {
239                let trits = self.trits();
240                let mut out: $bct = 0;
241                let mut i = 0;
242                while i < $trits {
243                    out |= (trits[i].bct_code() as $bct) << (2 * i);
244                    i += 1;
245                }
246                out
247            }
248
249            /// Unpacks from BCT, rejecting the `10` NaT pattern.
250            ///
251            /// NaT marks poisoned host memory and must never reach a guest, so
252            /// decoding it is an error rather than a silent zero (D-03).
253            pub const fn from_bct(bits: $bct) -> Result<Self, NatError> {
254                let mut trits = [Trit::Zero; $trits];
255                let mut i = 0;
256                while i < $trits {
257                    let code = ((bits >> (2 * i)) & 0b11) as u8;
258                    match Trit::from_bct_code(code) {
259                        Some(t) => trits[i] = t,
260                        None => return Err(NatError { position: i as u32 }),
261                    }
262                    i += 1;
263                }
264                Ok(Self::from_trits(trits))
265            }
266
267            // ---- arithmetic ----
268            //
269            // Checked forms return None on overflow. Intermediates go through
270            // i128. If an i128 operation overflows, the true result is outside
271            // this type's range, so the i128 check is a sound overflow test.
272
273            /// Addition, `None` on overflow.
274            #[inline]
275            pub const fn checked_add(self, rhs: Self) -> Option<Self> {
276                match (self.0 as i128).checked_add(rhs.0 as i128) {
277                    Some(v) => Self::from_i128(v),
278                    None => None,
279                }
280            }
281
282            /// Subtraction, `None` on overflow.
283            #[inline]
284            pub const fn checked_sub(self, rhs: Self) -> Option<Self> {
285                match (self.0 as i128).checked_sub(rhs.0 as i128) {
286                    Some(v) => Self::from_i128(v),
287                    None => None,
288                }
289            }
290
291            /// Multiplication, `None` on overflow.
292            #[inline]
293            pub const fn checked_mul(self, rhs: Self) -> Option<Self> {
294                match (self.0 as i128).checked_mul(rhs.0 as i128) {
295                    Some(v) => Self::from_i128(v),
296                    None => None,
297                }
298            }
299
300            /// Truncating division, `None` only when `rhs` is zero.
301            ///
302            /// Division cannot overflow: the range is symmetric, so there is no
303            /// `MIN / -1` case (D-01).
304            #[inline]
305            pub const fn checked_div(self, rhs: Self) -> Option<Self> {
306                if rhs.0 == 0 {
307                    return None;
308                }
309                Self::from_i128((self.0 as i128) / (rhs.0 as i128))
310            }
311
312            /// Remainder, `None` only when `rhs` is zero.
313            #[inline]
314            pub const fn checked_rem(self, rhs: Self) -> Option<Self> {
315                if rhs.0 == 0 {
316                    return None;
317                }
318                Self::from_i128((self.0 as i128) % (rhs.0 as i128))
319            }
320
321            /// Negation. Cannot overflow; the range is symmetric (D-01).
322            #[inline]
323            pub const fn neg(self) -> Self {
324                Self(-self.0)
325            }
326
327            /// Absolute value. Total; the range is symmetric (D-01).
328            #[inline]
329            pub const fn abs(self) -> Self {
330                Self(if self.0 < 0 { -self.0 } else { self.0 })
331            }
332
333            /// Addition modulo 3^n.
334            #[inline]
335            pub const fn wrapping_add(self, rhs: Self) -> Self {
336                Self::from_i128_wrapping(self.0 as i128 + rhs.0 as i128)
337            }
338
339            /// Subtraction modulo 3^n.
340            #[inline]
341            pub const fn wrapping_sub(self, rhs: Self) -> Self {
342                Self::from_i128_wrapping(self.0 as i128 - rhs.0 as i128)
343            }
344
345            /// Shift left by `k` trits, multiplying by 3^k. `None` on
346            /// overflow. Exact by construction.
347            #[inline]
348            pub const fn checked_shl_trits(self, k: u32) -> Option<Self> {
349                if k >= 128 {
350                    return if self.0 == 0 { Some(Self::ZERO) } else { None };
351                }
352                match (self.0 as i128).checked_mul(pow3(k)) {
353                    Some(v) => Self::from_i128(v),
354                    None => None,
355                }
356            }
357
358            /// Shift right by `k` trits, dropping the low `k` digits. This is
359            /// division by 3^k rounded to nearest: the dropped digits are worth
360            /// at most (3^k-1)/2, half a unit in the last place.
361            ///
362            /// This is not the host's `/`, which truncates toward zero. For
363            /// 1,598,046,941,971 / 27 truncation leaves a remainder of 19
364            /// against a ULP of 27. Implementing this as `self.0 / pow3(k)` was
365            /// a bug, caught by `truncation_is_round_to_nearest`.
366            ///
367            /// Ties cannot occur: 3^k is odd, so twice the remainder never
368            /// equals it. No round-half-to-even rule is needed.
369            #[inline]
370            pub const fn shr_trits(self, k: u32) -> Self {
371                if k >= 128 {
372                    return Self::ZERO;
373                }
374                let d = pow3(k);
375                let v = self.0 as i128;
376                let q = v / d;
377                let r = v - q * d;
378                // Nudge toward the nearer multiple when the truncated-toward-zero
379                // remainder exceeds half a ULP. Symmetric in sign.
380                let adj = if 2 * r > d {
381                    1
382                } else if 2 * r < -d {
383                    -1
384                } else {
385                    0
386                };
387                Self((q + adj) as $host)
388            }
389
390            /// Three-way comparison returning one trit (D-06).
391            ///
392            /// `Neg` for less, `Zero` for equal, `Pos` for greater.
393            #[inline]
394            pub const fn cmp3(self, rhs: Self) -> Trit {
395                if self.0 < rhs.0 {
396                    Trit::Neg
397                } else if self.0 > rhs.0 {
398                    Trit::Pos
399                } else {
400                    Trit::Zero
401                }
402            }
403
404            /// The sign, as a trit.
405            #[inline]
406            pub const fn signum(self) -> Trit {
407                self.cmp3(Self::ZERO)
408            }
409
410            /// Applies a per-trit unary operation.
411            pub fn map_trits(self, f: impl Fn(Trit) -> Trit) -> Self {
412                let mut t = self.trits();
413                for slot in t.iter_mut() {
414                    *slot = f(*slot);
415                }
416                Self::from_trits(t)
417            }
418
419            /// Applies a per-trit binary operation, used by `MIN` and `MAX`.
420            pub fn zip_trits(self, rhs: Self, f: impl Fn(Trit, Trit) -> Trit) -> Self {
421                let a = self.trits();
422                let b = rhs.trits();
423                let mut out = [Trit::Zero; $trits];
424                for i in 0..$trits {
425                    out[i] = f(a[i], b[i]);
426                }
427                Self::from_trits(out)
428            }
429
430            /// Per-trit logical AND (minimum).
431            pub fn trit_and(self, rhs: Self) -> Self {
432                self.zip_trits(rhs, Trit::and)
433            }
434
435            /// Per-trit logical OR (maximum).
436            pub fn trit_or(self, rhs: Self) -> Self {
437                self.zip_trits(rhs, Trit::or)
438            }
439
440            /// Per-trit logical NOT. Identical to arithmetic negation, since
441            /// negating a balanced ternary number flips every digit (D-01).
442            pub fn trit_not(self) -> Self {
443                self.neg()
444            }
445        }
446
447        impl fmt::Display for $name {
448            /// Balanced ternary digits, most significant first, `T` for -1.
449            fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
450                let trits = self.trits();
451                let mut hi = None;
452                for i in (0..$trits).rev() {
453                    if trits[i] != Trit::Zero {
454                        hi = Some(i);
455                        break;
456                    }
457                }
458                match hi {
459                    None => f.write_str("0"),
460                    Some(h) => {
461                        for i in (0..=h).rev() {
462                            fmt::Display::fmt(&trits[i], f)?;
463                        }
464                        Ok(())
465                    }
466                }
467            }
468        }
469
470        impl fmt::Debug for $name {
471            fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
472                write!(f, "{}({} = {})", stringify!($name), self, self.0)
473            }
474        }
475
476        impl FromStr for $name {
477            type Err = ParseTritsError;
478
479            /// Parses balanced ternary digits, most significant first.
480            /// Underscores are ignored as separators.
481            fn from_str(s: &str) -> Result<Self, Self::Err> {
482                let digits: Vec<char> = s.chars().filter(|c| *c != '_').collect();
483                if digits.is_empty() {
484                    return Err(ParseTritsError::Empty);
485                }
486                if digits.len() > $trits {
487                    return Err(ParseTritsError::TooLong {
488                        got: digits.len(),
489                        max: $trits,
490                    });
491                }
492                let mut trits = [Trit::Zero; $trits];
493                // Input is most-significant-first; storage is the reverse.
494                for (i, c) in digits.iter().rev().enumerate() {
495                    trits[i] = Trit::from_char(*c).ok_or(ParseTritsError::BadDigit(*c))?;
496                }
497                Ok(Self::from_trits(trits))
498            }
499        }
500
501        impl core::ops::Neg for $name {
502            type Output = Self;
503            #[inline]
504            fn neg(self) -> Self {
505                Self::neg(self)
506            }
507        }
508    };
509}
510
511balanced_type!(Tryte, 9, i32, u32, "tryte, the addressable unit");
512balanced_type!(Word, 27, i64, u64, "word");
513balanced_type!(DoubleWord, 54, i128, u128, "double word");
514
515/// Splits a value into trytes and rebuilds it from them.
516///
517/// A tryte is nine trits, so the tryte at index `k` carries weight 3^(9k) and
518/// the decomposition is exact: `$trytes` trytes tile the range with nothing
519/// left over. Memory and memory-mapped devices both store multi-tryte values
520/// this way, least significant tryte first.
521macro_rules! tryte_split {
522    ($name:ident, $trytes:expr, $host:ty) => {
523        impl $name {
524            /// Number of trytes in the representation.
525            pub const TRYTES: usize = $trytes;
526
527            /// Tryte `index`, counting from the least significant.
528            ///
529            /// # Panics
530            /// If `index >= TRYTES`.
531            #[track_caller]
532            pub const fn tryte(self, index: usize) -> Tryte {
533                assert!(index < $trytes, "tryte index out of range");
534                let trits = self.trits();
535                let mut chunk = [Trit::Zero; 9];
536                let mut i = 0;
537                while i < 9 {
538                    chunk[i] = trits[index * 9 + i];
539                    i += 1;
540                }
541                Tryte::from_trits(chunk)
542            }
543
544            /// Rebuilds a value from its trytes, least significant first.
545            pub const fn from_trytes(trytes: [Tryte; $trytes]) -> Self {
546                let mut acc: i128 = 0;
547                let mut weight: i128 = 1;
548                let mut i = 0;
549                while i < $trytes {
550                    acc += (trytes[i].value() as i128) * weight;
551                    weight *= 19683;
552                    i += 1;
553                }
554                Self(acc as $host)
555            }
556        }
557    };
558}
559
560tryte_split!(Word, 3, i64);
561tryte_split!(DoubleWord, 6, i128);
562
563/// Adds `wrapping_mul` where the exact product fits an `i128` intermediate.
564///
565/// Not provided for [`DoubleWord`]. Two 54-trit values multiply to at most
566/// 3^108, about 4*10^51, which exceeds `i128`. `i128::wrapping_mul` would
567/// reduce modulo 2^128 rather than 3^54 and give a wrong answer.
568/// [`DoubleWord::checked_mul`] is exact and unaffected.
569macro_rules! wrapping_mul_via_i128 {
570    ($name:ident) => {
571        impl $name {
572            /// Multiplication modulo 3^n.
573            #[inline]
574            pub const fn wrapping_mul(self, rhs: Self) -> Self {
575                Self::from_i128_wrapping(self.value() as i128 * rhs.value() as i128)
576            }
577        }
578    };
579}
580
581wrapping_mul_via_i128!(Tryte);
582wrapping_mul_via_i128!(Word);
583
584impl Word {
585    /// Widens to a [`DoubleWord`]. Always exact.
586    #[inline]
587    pub const fn widen(self) -> DoubleWord {
588        DoubleWord(self.value() as i128)
589    }
590
591    /// The exact product of two words, as a double word.
592    ///
593    /// Cannot overflow: the product of two p-trit values is exactly 2p trits,
594    /// and 27 + 27 is the double-word width. A `DoubleWord` accumulator
595    /// therefore holds each product exactly and only the summation rounds
596    /// (D-02).
597    #[inline]
598    pub const fn widening_mul(self, rhs: Self) -> DoubleWord {
599        DoubleWord(self.value() as i128 * rhs.value() as i128)
600    }
601}
602
603impl Tryte {
604    /// Widens to a [`Word`]. Always exact.
605    #[inline]
606    pub const fn widen(self) -> Word {
607        Word(self.value() as i64)
608    }
609}
610
611#[cfg(test)]
612mod tryte_split_tests {
613    use super::*;
614
615    #[test]
616    fn word_trytes_carry_powers_of_three_to_the_ninth() {
617        // Tryte k has weight 3^(9k), so a value of exactly 3^(9k) puts a one
618        // in tryte k and zero everywhere else.
619        for k in 0..Word::TRYTES {
620            let v = Word::from_value(19683i64.pow(k as u32));
621            for j in 0..Word::TRYTES {
622                let want = if j == k { 1 } else { 0 };
623                assert_eq!(v.tryte(j).value(), want, "3^(9*{k}) tryte {j}");
624            }
625        }
626    }
627
628    #[test]
629    fn word_round_trips_through_its_trytes() {
630        let cases = [
631            Word::MIN,
632            Word::from_value(-1),
633            Word::ZERO,
634            Word::from_value(1),
635            Word::from_value(-4_294_967_296),
636            Word::from_value(1_598_046_941_971),
637            Word::MAX,
638        ];
639        for v in cases {
640            let trytes = [v.tryte(0), v.tryte(1), v.tryte(2)];
641            assert_eq!(Word::from_trytes(trytes), v, "value {}", v.value());
642        }
643    }
644
645    #[test]
646    fn double_word_round_trips_through_its_trytes() {
647        let cases = [
648            DoubleWord::MIN,
649            DoubleWord::ZERO,
650            DoubleWord::from_value(-987_654_321_098_765_432),
651            DoubleWord::MAX,
652        ];
653        for v in cases {
654            let mut trytes = [Tryte::ZERO; DoubleWord::TRYTES];
655            for (i, slot) in trytes.iter_mut().enumerate() {
656                *slot = v.tryte(i);
657            }
658            assert_eq!(DoubleWord::from_trytes(trytes), v, "value {}", v.value());
659        }
660    }
661
662    #[test]
663    fn every_tryte_value_round_trips_in_the_low_position() {
664        for n in Tryte::MIN.value()..=Tryte::MAX.value() {
665            let t = Tryte::from_value(n);
666            let w = Word::from_trytes([t, Tryte::ZERO, Tryte::ZERO]);
667            assert_eq!(w.value(), n as i64);
668            assert_eq!(w.tryte(0), t);
669        }
670    }
671}