Skip to main content

hopper_runtime/
enum_byte.rs

1//! Unit enums in zero-copy layouts and instruction arguments.
2//!
3//! A Rust enum is not `Pod`: a `#[repr(u8)]` enum with three variants has
4//! 253 byte values that are not a valid value of the type, so overlaying
5//! it on account bytes is undefined behaviour the moment an account holds
6//! one of them. The usual workaround is a bare `u8` field and a
7//! hand-written `match`, which loses the type in the layout and lets a
8//! handler forget the validation.
9//!
10//! [`EnumByte<E>`] is the field type instead: one byte, alignment 1, every
11//! bit pattern valid as far as memory safety goes, and the enum recovered
12//! through [`EnumByte::get`], which refuses a byte that names no variant.
13//! `#[hopper::unit_enum]` implements [`UnitEnum`] for a fieldless enum, so
14//! the mapping between variants and bytes is generated, not written.
15//!
16//! ```ignore
17//! #[hopper::unit_enum]
18//! pub enum Status {
19//!     Open = 1,
20//!     Settled = 2,
21//!     Cancelled = 3,
22//! }
23//!
24//! #[hopper::state(disc = 5, version = 1)]
25//! #[derive(Clone, Copy)]
26//! #[repr(C)]
27//! pub struct Order {
28//!     pub maker: Address,
29//!     pub status: EnumByte<Status>,
30//! }
31//!
32//! if order.status.get()? == Status::Open {
33//!     order.status.set(Status::Settled);
34//! }
35//! ```
36
37use crate::error::ProgramError;
38use crate::pod::{Pod, Zeroable};
39use crate::result::ProgramResult;
40use core::marker::PhantomData;
41
42/// A fieldless enum with a one-byte representation. Implemented by
43/// `#[hopper::unit_enum]`; hand-written impls must keep `from_byte` the
44/// exact inverse of `to_byte` on every variant.
45pub trait UnitEnum: Copy + Sized {
46    /// The variant's byte.
47    fn to_byte(self) -> u8;
48
49    /// The variant this byte names, if any.
50    fn from_byte(byte: u8) -> Option<Self>;
51}
52
53/// One byte that stores a [`UnitEnum`]. `Pod`, so it can sit in a
54/// `#[hopper::state]` layout, a `#[hopper::pod]` struct, or
55/// `#[hopper::args]`; the enum is validated when it is read.
56#[repr(transparent)]
57pub struct EnumByte<E: UnitEnum> {
58    byte: u8,
59    _enum: PhantomData<E>,
60}
61
62impl<E: UnitEnum> Clone for EnumByte<E> {
63    #[inline(always)]
64    fn clone(&self) -> Self {
65        *self
66    }
67}
68
69impl<E: UnitEnum> Copy for EnumByte<E> {}
70
71impl<E: UnitEnum> EnumByte<E> {
72    /// Store `value`.
73    #[inline(always)]
74    pub fn new(value: E) -> Self {
75        Self {
76            byte: value.to_byte(),
77            _enum: PhantomData,
78        }
79    }
80
81    /// Wrap a raw byte without checking that it names a variant;
82    /// [`get`](Self::get) checks on every read.
83    #[inline(always)]
84    pub const fn from_raw(byte: u8) -> Self {
85        Self {
86            byte,
87            _enum: PhantomData,
88        }
89    }
90
91    /// The stored variant. A byte that names no variant is refused with
92    /// `InvalidAccountData`.
93    #[inline(always)]
94    pub fn get(&self) -> Result<E, ProgramError> {
95        E::from_byte(self.byte).ok_or(ProgramError::InvalidAccountData)
96    }
97
98    /// Whether the stored byte names a variant.
99    #[inline(always)]
100    pub fn validate(&self) -> ProgramResult {
101        self.get().map(|_| ())
102    }
103
104    /// Whether the stored byte is exactly `value`'s. Never fails: an
105    /// unknown byte is simply not `value`.
106    #[inline(always)]
107    pub fn is(&self, value: E) -> bool {
108        self.byte == value.to_byte()
109    }
110
111    /// Replace the stored variant.
112    #[inline(always)]
113    pub fn set(&mut self, value: E) {
114        self.byte = value.to_byte();
115    }
116
117    /// The raw byte.
118    #[inline(always)]
119    pub const fn raw(&self) -> u8 {
120        self.byte
121    }
122}
123
124impl<E: UnitEnum> From<E> for EnumByte<E> {
125    #[inline(always)]
126    fn from(value: E) -> Self {
127        Self::new(value)
128    }
129}
130
131impl<E: UnitEnum> PartialEq for EnumByte<E> {
132    #[inline(always)]
133    fn eq(&self, other: &Self) -> bool {
134        self.byte == other.byte
135    }
136}
137
138impl<E: UnitEnum> Eq for EnumByte<E> {}
139
140impl<E: UnitEnum> PartialEq<E> for EnumByte<E> {
141    #[inline(always)]
142    fn eq(&self, other: &E) -> bool {
143        self.is(*other)
144    }
145}
146
147impl<E: UnitEnum + core::fmt::Debug> core::fmt::Debug for EnumByte<E> {
148    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
149        match E::from_byte(self.byte) {
150            Some(value) => value.fmt(f),
151            None => write!(f, "EnumByte(invalid {})", self.byte),
152        }
153    }
154}
155
156// SAFETY: `EnumByte` is `repr(transparent)` over one `u8` (the marker is
157// zero-sized): alignment 1, no padding, no pointers, and every byte value
158// is a valid `EnumByte`; the enum itself is only produced by `get`, which
159// validates. `Copy + Sized` holds through the manual impls above.
160unsafe impl<E: UnitEnum> Zeroable for EnumByte<E> {}
161// SAFETY: as above.
162unsafe impl<E: UnitEnum> Pod for EnumByte<E> {}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    #[derive(Clone, Copy, Debug, PartialEq, Eq)]
169    #[repr(u8)]
170    enum Status {
171        Open = 1,
172        Settled = 2,
173        Cancelled = 7,
174    }
175
176    impl UnitEnum for Status {
177        fn to_byte(self) -> u8 {
178            self as u8
179        }
180        fn from_byte(byte: u8) -> Option<Self> {
181            match byte {
182                1 => Some(Self::Open),
183                2 => Some(Self::Settled),
184                7 => Some(Self::Cancelled),
185                _ => None,
186            }
187        }
188    }
189
190    #[test]
191    fn layout_is_one_byte_with_alignment_one() {
192        assert_eq!(core::mem::size_of::<EnumByte<Status>>(), 1);
193        assert_eq!(core::mem::align_of::<EnumByte<Status>>(), 1);
194    }
195
196    #[test]
197    fn reads_validate_and_writes_store_the_variant_byte() {
198        let mut field = EnumByte::new(Status::Open);
199        assert_eq!(field.raw(), 1);
200        assert_eq!(field.get(), Ok(Status::Open));
201        assert!(field == Status::Open);
202        assert!(field.is(Status::Open) && !field.is(Status::Settled));
203        field.set(Status::Cancelled);
204        assert_eq!(field.raw(), 7);
205        assert_eq!(field.get(), Ok(Status::Cancelled));
206        assert_eq!(EnumByte::from(Status::Settled), EnumByte::from_raw(2));
207    }
208
209    #[test]
210    fn a_byte_that_names_no_variant_is_refused_not_transmuted() {
211        for byte in [0u8, 3, 6, 8, 255] {
212            let field = EnumByte::<Status>::from_raw(byte);
213            assert_eq!(field.get(), Err(ProgramError::InvalidAccountData));
214            assert!(field.validate().is_err());
215            assert!(!field.is(Status::Open));
216        }
217        assert_eq!(
218            std::format!("{:?}", EnumByte::<Status>::from_raw(9)),
219            "EnumByte(invalid 9)"
220        );
221        assert_eq!(std::format!("{:?}", EnumByte::new(Status::Open)), "Open");
222    }
223
224    #[test]
225    fn overlays_on_account_bytes() {
226        let bytes = [2u8, 7, 9];
227        // SAFETY: `EnumByte` is `repr(transparent)` over `u8`; the slice
228        // holds three initialized bytes.
229        let fields: &[EnumByte<Status>; 3] =
230            unsafe { &*(bytes.as_ptr() as *const [EnumByte<Status>; 3]) };
231        assert_eq!(fields[0].get(), Ok(Status::Settled));
232        assert_eq!(fields[1].get(), Ok(Status::Cancelled));
233        assert!(fields[2].get().is_err());
234    }
235}