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(self.0.wrapping_neg())
44 }
45
46 #[allow(clippy::should_implement_trait)]
52 #[inline]
53 pub fn not(self) -> Self {
54 Choice(self.0 ^ 1)
55 }
56
57 #[inline]
59 pub fn and(self, other: Self) -> Self {
60 Choice(self.0 & other.0)
61 }
62
63 #[inline]
65 pub fn or(self, other: Self) -> Self {
66 Choice(self.0 | other.0)
67 }
68}
69
70impl From<Choice> for bool {
71 #[inline]
72 fn from(c: Choice) -> bool {
73 c.unwrap_u8() == 1
74 }
75}
76
77#[inline]
82pub fn eq(a: &[u8], b: &[u8]) -> Choice {
83 if a.len() != b.len() {
84 return Choice::FALSE;
85 }
86 let mut acc: u8 = 0;
87 for i in 0..a.len() {
88 acc |= a[i] ^ b[i];
89 }
90 Choice::from_u8(acc).not()
91}
92
93#[inline]
97#[must_use = "this is the result of a cryptographic verification; discarding it accepts everything"]
98pub fn verify(expected: &[u8], actual: &[u8]) -> bool {
99 eq(expected, actual).into()
100}
101
102#[inline]
104pub fn select_u8(c: Choice, a: u8, b: u8) -> u8 {
105 let m = c.mask();
106 b ^ (m & (a ^ b))
107}
108
109#[inline]
111pub fn select_u32(c: Choice, a: u32, b: u32) -> u32 {
112 let m = (c.unwrap_u8() as u32).wrapping_neg();
113 b ^ (m & (a ^ b))
114}
115
116#[inline]
118pub fn select_u64(c: Choice, a: u64, b: u64) -> u64 {
119 let m = (c.unwrap_u8() as u64).wrapping_neg();
120 b ^ (m & (a ^ b))
121}
122
123#[inline]
127pub fn cswap(c: Choice, a: &mut [u8], b: &mut [u8]) {
128 debug_assert_eq!(a.len(), b.len());
129 let m = c.mask();
130 let n = core::cmp::min(a.len(), b.len());
131 for i in 0..n {
132 let t = m & (a[i] ^ b[i]);
133 a[i] ^= t;
134 b[i] ^= t;
135 }
136}
137
138#[inline]
140pub fn cmov(c: Choice, dst: &mut [u8], src: &[u8]) {
141 debug_assert_eq!(dst.len(), src.len());
142 let m = c.mask();
143 let n = core::cmp::min(dst.len(), src.len());
144 for i in 0..n {
145 dst[i] ^= m & (dst[i] ^ src[i]);
146 }
147}
148
149pub fn lt_be(a: &[u8], b: &[u8]) -> Choice {
151 debug_assert_eq!(a.len(), b.len());
152 let mut borrow: u16 = 0;
153 for i in (0..a.len()).rev() {
154 let d = (a[i] as u16).wrapping_sub(b[i] as u16).wrapping_sub(borrow);
155 borrow = (d >> 8) & 1;
156 }
157 Choice::from_u8(borrow as u8)
158}
159
160#[inline]
162pub fn is_zero(x: &[u8]) -> Choice {
163 let mut acc = 0u8;
164 for &b in x {
165 acc |= b;
166 }
167 Choice::from_u8(acc).not()
168}
169
170#[cfg(test)]
171mod tests {
172 use super::*;
173
174 #[test]
175 fn eq_matches_semantics() {
176 assert!(bool::from(eq(b"abc", b"abc")));
177 assert!(!bool::from(eq(b"abc", b"abd")));
178 assert!(!bool::from(eq(b"abc", b"ab")));
179 assert!(bool::from(eq(b"", b"")));
180 }
181
182 #[test]
183 fn select_picks_correct_branch() {
184 assert_eq!(select_u8(Choice::TRUE, 0xAA, 0x55), 0xAA);
185 assert_eq!(select_u8(Choice::FALSE, 0xAA, 0x55), 0x55);
186 assert_eq!(select_u32(Choice::TRUE, 1, 2), 1);
187 assert_eq!(select_u64(Choice::FALSE, 1, 2), 2);
188 }
189
190 #[test]
191 fn cswap_is_conditional() {
192 let (mut a, mut b) = ([1u8, 2, 3], [4u8, 5, 6]);
193 cswap(Choice::FALSE, &mut a, &mut b);
194 assert_eq!((a, b), ([1, 2, 3], [4, 5, 6]));
195 cswap(Choice::TRUE, &mut a, &mut b);
196 assert_eq!((a, b), ([4, 5, 6], [1, 2, 3]));
197 }
198
199 #[test]
200 fn cmov_is_conditional() {
201 let mut dst = [0u8; 4];
202 cmov(Choice::FALSE, &mut dst, &[9, 9, 9, 9]);
203 assert_eq!(dst, [0, 0, 0, 0]);
204 cmov(Choice::TRUE, &mut dst, &[9, 9, 9, 9]);
205 assert_eq!(dst, [9, 9, 9, 9]);
206 }
207
208 #[test]
209 fn lt_be_orders_correctly() {
210 assert!(bool::from(lt_be(&[0, 1], &[0, 2])));
211 assert!(!bool::from(lt_be(&[0, 2], &[0, 1])));
212 assert!(!bool::from(lt_be(&[0, 2], &[0, 2])));
213 assert!(bool::from(lt_be(&[0x00, 0xFF], &[0x01, 0x00])));
214 }
215
216 #[test]
217 fn is_zero_detects_all_zero() {
218 assert!(bool::from(is_zero(&[0, 0, 0])));
219 assert!(!bool::from(is_zero(&[0, 1, 0])));
220 }
221}