1use core::hint::black_box;
10
11#[must_use = "a Choice carries the outcome of a constant-time comparison"]
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub struct Choice(u8);
18
19impl Choice {
20 pub const FALSE: Choice = Choice(0);
22 pub const TRUE: Choice = Choice(1);
24
25 #[inline]
27 pub fn from_u8(v: u8) -> Self {
28 Choice(((v | v.wrapping_neg()) >> 7) & 1)
29 }
30
31 #[inline]
36 pub fn unwrap_u8(self) -> u8 {
37 black_box(self.0)
38 }
39
40 #[inline]
42 pub fn mask(self) -> u8 {
43 black_box(black_box(self.0).wrapping_neg())
46 }
47
48 #[allow(clippy::should_implement_trait)]
54 #[inline]
55 pub fn not(self) -> Self {
56 Choice(self.0 ^ 1)
57 }
58
59 #[inline]
61 pub fn and(self, other: Self) -> Self {
62 Choice(self.0 & other.0)
63 }
64
65 #[inline]
67 pub fn or(self, other: Self) -> Self {
68 Choice(self.0 | other.0)
69 }
70}
71
72impl From<Choice> for bool {
73 #[inline]
74 fn from(c: Choice) -> bool {
75 c.unwrap_u8() == 1
76 }
77}
78
79#[inline]
84pub fn eq(a: &[u8], b: &[u8]) -> Choice {
85 if a.len() != b.len() {
86 return Choice::FALSE;
87 }
88 let mut acc: u8 = 0;
89 for i in 0..a.len() {
90 acc |= a[i] ^ b[i];
91 }
92 Choice::from_u8(acc).not()
93}
94
95#[inline]
99#[must_use = "this is the result of a cryptographic verification; discarding it accepts everything"]
100pub fn verify(expected: &[u8], actual: &[u8]) -> bool {
101 eq(expected, actual).into()
102}
103
104#[inline]
106pub fn select_u8(c: Choice, a: u8, b: u8) -> u8 {
107 let m = c.mask();
108 b ^ (m & (a ^ b))
109}
110
111#[inline]
113pub fn select_u32(c: Choice, a: u32, b: u32) -> u32 {
114 let m = (c.unwrap_u8() as u32).wrapping_neg();
115 b ^ (m & (a ^ b))
116}
117
118#[inline]
120pub fn select_u64(c: Choice, a: u64, b: u64) -> u64 {
121 let m = (c.unwrap_u8() as u64).wrapping_neg();
122 b ^ (m & (a ^ b))
123}
124
125#[inline]
129pub fn cswap(c: Choice, a: &mut [u8], b: &mut [u8]) {
130 debug_assert_eq!(a.len(), b.len());
131 let m = c.mask();
132 let n = core::cmp::min(a.len(), b.len());
133 for i in 0..n {
134 let t = m & (a[i] ^ b[i]);
135 a[i] ^= t;
136 b[i] ^= t;
137 }
138}
139
140#[inline]
142pub fn cmov(c: Choice, dst: &mut [u8], src: &[u8]) {
143 debug_assert_eq!(dst.len(), src.len());
144 let m = c.mask();
145 let n = core::cmp::min(dst.len(), src.len());
146 for i in 0..n {
147 dst[i] ^= m & (dst[i] ^ src[i]);
148 }
149}
150
151pub fn lt_be(a: &[u8], b: &[u8]) -> Choice {
153 debug_assert_eq!(a.len(), b.len());
154 let mut borrow: u16 = 0;
155 for i in (0..a.len()).rev() {
156 let d = (a[i] as u16).wrapping_sub(b[i] as u16).wrapping_sub(borrow);
157 borrow = (d >> 8) & 1;
158 }
159 Choice::from_u8(borrow as u8)
160}
161
162#[inline]
164pub fn is_zero(x: &[u8]) -> Choice {
165 let mut acc = 0u8;
166 for &b in x {
167 acc |= b;
168 }
169 Choice::from_u8(acc).not()
170}
171
172#[cfg(test)]
173mod tests {
174 use super::*;
175
176 #[test]
177 fn choice_masks_preserve_nonzero_semantics_for_every_byte() {
178 for flag in 0..=u8::MAX {
179 let choice = Choice::from_u8(flag);
180 let truth = flag != 0;
181 assert_eq!(choice.mask(), if truth { u8::MAX } else { 0 });
182 assert_eq!(
183 select_u8(choice, 0xa5, 0x5a),
184 if truth { 0xa5 } else { 0x5a }
185 );
186 let mut dst = [0x5a; 4];
187 cmov(choice, &mut dst, &[0xa5; 4]);
188 assert_eq!(dst, if truth { [0xa5; 4] } else { [0x5a; 4] });
189 let (mut a, mut b) = ([0xa5; 4], [0x5a; 4]);
190 cswap(choice, &mut a, &mut b);
191 assert_eq!(
192 (a, b),
193 if truth {
194 ([0x5a; 4], [0xa5; 4])
195 } else {
196 ([0xa5; 4], [0x5a; 4])
197 }
198 );
199 }
200 }
201
202 #[test]
203 fn eq_matches_semantics() {
204 assert!(bool::from(eq(b"abc", b"abc")));
205 assert!(!bool::from(eq(b"abc", b"abd")));
206 assert!(!bool::from(eq(b"abc", b"ab")));
207 assert!(bool::from(eq(b"", b"")));
208 }
209
210 #[test]
211 fn select_picks_correct_branch() {
212 assert_eq!(select_u8(Choice::TRUE, 0xAA, 0x55), 0xAA);
213 assert_eq!(select_u8(Choice::FALSE, 0xAA, 0x55), 0x55);
214 assert_eq!(select_u32(Choice::TRUE, 1, 2), 1);
215 assert_eq!(select_u64(Choice::FALSE, 1, 2), 2);
216 }
217
218 #[test]
219 fn cswap_is_conditional() {
220 let (mut a, mut b) = ([1u8, 2, 3], [4u8, 5, 6]);
221 cswap(Choice::FALSE, &mut a, &mut b);
222 assert_eq!((a, b), ([1, 2, 3], [4, 5, 6]));
223 cswap(Choice::TRUE, &mut a, &mut b);
224 assert_eq!((a, b), ([4, 5, 6], [1, 2, 3]));
225 }
226
227 #[test]
228 fn cmov_is_conditional() {
229 let mut dst = [0u8; 4];
230 cmov(Choice::FALSE, &mut dst, &[9, 9, 9, 9]);
231 assert_eq!(dst, [0, 0, 0, 0]);
232 cmov(Choice::TRUE, &mut dst, &[9, 9, 9, 9]);
233 assert_eq!(dst, [9, 9, 9, 9]);
234 }
235
236 #[test]
237 fn lt_be_orders_correctly() {
238 assert!(bool::from(lt_be(&[0, 1], &[0, 2])));
239 assert!(!bool::from(lt_be(&[0, 2], &[0, 1])));
240 assert!(!bool::from(lt_be(&[0, 2], &[0, 2])));
241 assert!(bool::from(lt_be(&[0x00, 0xFF], &[0x01, 0x00])));
242 }
243
244 #[test]
245 fn is_zero_detects_all_zero() {
246 assert!(bool::from(is_zero(&[0, 0, 0])));
247 assert!(!bool::from(is_zero(&[0, 1, 0])));
248 }
249}