hopper_runtime/
enum_byte.rs1use crate::error::ProgramError;
38use crate::pod::{Pod, Zeroable};
39use crate::result::ProgramResult;
40use core::marker::PhantomData;
41
42pub trait UnitEnum: Copy + Sized {
46 fn to_byte(self) -> u8;
48
49 fn from_byte(byte: u8) -> Option<Self>;
51}
52
53#[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 #[inline(always)]
74 pub fn new(value: E) -> Self {
75 Self {
76 byte: value.to_byte(),
77 _enum: PhantomData,
78 }
79 }
80
81 #[inline(always)]
84 pub const fn from_raw(byte: u8) -> Self {
85 Self {
86 byte,
87 _enum: PhantomData,
88 }
89 }
90
91 #[inline(always)]
94 pub fn get(&self) -> Result<E, ProgramError> {
95 E::from_byte(self.byte).ok_or(ProgramError::InvalidAccountData)
96 }
97
98 #[inline(always)]
100 pub fn validate(&self) -> ProgramResult {
101 self.get().map(|_| ())
102 }
103
104 #[inline(always)]
107 pub fn is(&self, value: E) -> bool {
108 self.byte == value.to_byte()
109 }
110
111 #[inline(always)]
113 pub fn set(&mut self, value: E) {
114 self.byte = value.to_byte();
115 }
116
117 #[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
156unsafe impl<E: UnitEnum> Zeroable for EnumByte<E> {}
161unsafe 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 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}