1use super::{
4 contracts::{BackendFault, InputError},
5 ordinary::OneShotError,
6 specifications::{Base64, Codec, CodecSettings, EncodePadding},
7};
8
9#[derive(Clone, Copy, Debug, Eq, PartialEq)]
11#[non_exhaustive]
12pub enum InPlaceError {
13 InputLengthExceedsBuffer {
15 input_len: usize,
17 buffer_len: usize,
19 },
20 LengthOverflow,
22 OutputTooSmall {
24 required: usize,
26 available: usize,
28 },
29 StagingTooSmall {
31 required: usize,
33 available: usize,
35 },
36 OverlappingBuffers,
38 AddressRangeOverflow,
40 SecretPolicyUnsupported,
42 Input(InputError),
44 PositionOverflow,
46 InvalidSecretInput,
48 Backend(BackendFault),
50}
51
52impl core::fmt::Display for InPlaceError {
53 fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
54 match self {
55 Self::InputLengthExceedsBuffer {
56 input_len,
57 buffer_len,
58 } => write!(
59 formatter,
60 "in-place input length {input_len} exceeds buffer length {buffer_len}"
61 ),
62 Self::LengthOverflow => formatter.write_str("in-place output length overflows usize"),
63 Self::OutputTooSmall {
64 required,
65 available,
66 } => write!(
67 formatter,
68 "in-place output requires {required} bytes; buffer has {available}"
69 ),
70 Self::StagingTooSmall {
71 required,
72 available,
73 } => write!(
74 formatter,
75 "secret decode staging requires {required} bytes; staging has {available}"
76 ),
77 Self::OverlappingBuffers => {
78 formatter.write_str("caller buffer and private staging overlap")
79 }
80 Self::AddressRangeOverflow => {
81 formatter.write_str("caller byte range address overflows usize")
82 }
83 Self::SecretPolicyUnsupported => {
84 formatter.write_str("codec policy is not eligible for secret decoding")
85 }
86 Self::Input(error) => error.fmt(formatter),
87 Self::PositionOverflow => formatter.write_str("base64 source position overflows usize"),
88 Self::InvalidSecretInput => formatter.write_str("invalid secret base64 input"),
89 Self::Backend(fault) => write!(formatter, "base64 backend fault: {}", fault.as_str()),
90 }
91 }
92}
93
94#[cfg(feature = "std")]
95impl std::error::Error for InPlaceError {}
96
97impl<S: Codec> Base64<S> {
98 pub fn encode_in_place(
104 &self,
105 buffer: &mut [u8],
106 input_len: usize,
107 ) -> Result<usize, InPlaceError> {
108 require_input_prefix(input_len, buffer.len())?;
109 let required = self.encoded_len(input_len).map_err(map_one_shot_error)?;
110 if required > buffer.len() {
111 return Err(InPlaceError::OutputTooSmall {
112 required,
113 available: buffer.len(),
114 });
115 }
116 encode_reverse(self.settings(), buffer, input_len, required);
117 Ok(required)
118 }
119
120 pub fn decode_in_place(
127 &self,
128 buffer: &mut [u8],
129 input_len: usize,
130 ) -> Result<usize, InPlaceError> {
131 require_input_prefix(input_len, buffer.len())?;
132 let required = self
133 .decoded_len(&buffer[..input_len])
134 .map_err(map_one_shot_error)?;
135 decode_forward(self.settings(), buffer, input_len);
136 Ok(required)
137 }
138}
139
140pub(super) fn require_input_prefix(
141 input_len: usize,
142 buffer_len: usize,
143) -> Result<(), InPlaceError> {
144 if input_len > buffer_len {
145 Err(InPlaceError::InputLengthExceedsBuffer {
146 input_len,
147 buffer_len,
148 })
149 } else {
150 Ok(())
151 }
152}
153
154#[cfg(feature = "secrets")]
155pub(super) fn require_disjoint_slices(left: &[u8], right: &[u8]) -> Result<(), InPlaceError> {
156 require_disjoint_ranges(
157 left.as_ptr() as usize,
158 left.len(),
159 right.as_ptr() as usize,
160 right.len(),
161 )
162}
163
164#[cfg(any(feature = "secrets", test, kani))]
165fn require_disjoint_ranges(
166 left_start: usize,
167 left_len: usize,
168 right_start: usize,
169 right_len: usize,
170) -> Result<(), InPlaceError> {
171 let left_end = left_start
172 .checked_add(left_len)
173 .ok_or(InPlaceError::AddressRangeOverflow)?;
174 let right_end = right_start
175 .checked_add(right_len)
176 .ok_or(InPlaceError::AddressRangeOverflow)?;
177 if left_len != 0 && right_len != 0 && left_start < right_end && right_start < left_end {
178 Err(InPlaceError::OverlappingBuffers)
179 } else {
180 Ok(())
181 }
182}
183
184#[cfg(kani)]
185pub(crate) fn require_in_place_disjoint_ranges_for_proof(
186 left_start: usize,
187 left_len: usize,
188 right_start: usize,
189 right_len: usize,
190) -> Result<(), InPlaceError> {
191 require_disjoint_ranges(left_start, left_len, right_start, right_len)
192}
193
194#[cfg(test)]
195pub(super) fn require_disjoint_ranges_for_test(
196 left_start: usize,
197 left_len: usize,
198 right_start: usize,
199 right_len: usize,
200) -> Result<(), InPlaceError> {
201 require_disjoint_ranges(left_start, left_len, right_start, right_len)
202}
203
204fn map_one_shot_error(error: OneShotError) -> InPlaceError {
205 match error {
206 OneShotError::LengthOverflow => InPlaceError::LengthOverflow,
207 OneShotError::Input(error) => InPlaceError::Input(error),
208 OneShotError::PositionOverflow => InPlaceError::PositionOverflow,
209 OneShotError::OutputTooSmall {
210 required,
211 available,
212 } => InPlaceError::OutputTooSmall {
213 required,
214 available,
215 },
216 OneShotError::Backend(fault) => InPlaceError::Backend(fault),
217 OneShotError::AllocationLimitExceeded { .. } | OneShotError::AllocationFailed { .. } => {
218 InPlaceError::Backend(BackendFault::ImpossibleState)
219 }
220 }
221}
222
223fn encode_reverse(settings: CodecSettings, buffer: &mut [u8], input_len: usize, output_len: usize) {
224 let alphabet = settings.alphabet().as_array();
225 let mut read = input_len;
226 let mut write = output_len;
227 let tail = input_len % 3;
228
229 if tail != 0 {
230 read -= tail;
231 let first = buffer[read];
232 let second = if tail == 2 { buffer[read + 1] } else { 0 };
233 let tail_len = encoded_tail_len(tail, settings.encode_padding() == EncodePadding::Padded);
234 write -= tail_len;
235 buffer[write] = alphabet[usize::from(first >> 2)];
236 buffer[write + 1] = alphabet[usize::from(((first & 3) << 4) | (second >> 4))];
237 if tail == 2 {
238 buffer[write + 2] = alphabet[usize::from((second & 15) << 2)];
239 if tail_len == 4 {
240 buffer[write + 3] = b'=';
241 }
242 } else if tail_len == 4 {
243 buffer[write + 2] = b'=';
244 buffer[write + 3] = b'=';
245 }
246 }
247
248 while read != 0 {
249 read -= 3;
250 write -= 4;
251 let first = buffer[read];
252 let second = buffer[read + 1];
253 let third = buffer[read + 2];
254 buffer[write] = alphabet[usize::from(first >> 2)];
255 buffer[write + 1] = alphabet[usize::from(((first & 3) << 4) | (second >> 4))];
256 buffer[write + 2] = alphabet[usize::from(((second & 15) << 2) | (third >> 6))];
257 buffer[write + 3] = alphabet[usize::from(third & 63)];
258 }
259}
260
261fn decode_forward(settings: CodecSettings, buffer: &mut [u8], input_len: usize) {
262 let mut read = 0;
263 let mut write = 0;
264 while input_len - read >= 4 {
265 let first = decode_value(settings, buffer[read]);
266 let second = decode_value(settings, buffer[read + 1]);
267 let third_byte = buffer[read + 2];
268 let fourth_byte = buffer[read + 3];
269 let produced = quantum_decoded_len(third_byte == b'=', fourth_byte == b'=');
270 buffer[write] = (first << 2) | (second >> 4);
271 if produced >= 2 {
272 let third = decode_value(settings, third_byte);
273 buffer[write + 1] = (second << 4) | (third >> 2);
274 if produced == 3 {
275 buffer[write + 2] = (third << 6) | decode_value(settings, fourth_byte);
276 }
277 }
278 write += produced;
279 read += 4;
280 }
281
282 let tail = input_len - read;
283 let produced = tail_decoded_len(tail);
284 if produced != 0 {
285 let first = decode_value(settings, buffer[read]);
286 let second = decode_value(settings, buffer[read + 1]);
287 buffer[write] = (first << 2) | (second >> 4);
288 if produced == 2 {
289 buffer[write + 1] = (second << 4) | (decode_value(settings, buffer[read + 2]) >> 2);
290 }
291 }
292}
293
294pub(crate) const fn encoded_tail_len(remainder: usize, padded: bool) -> usize {
296 if remainder == 0 {
297 0
298 } else if padded {
299 4
300 } else {
301 remainder + 1
302 }
303}
304
305pub(crate) const fn quantum_decoded_len(third_is_padding: bool, fourth_is_padding: bool) -> usize {
307 if third_is_padding {
308 1
309 } else if fourth_is_padding {
310 2
311 } else {
312 3
313 }
314}
315
316pub(crate) const fn tail_decoded_len(remainder: usize) -> usize {
318 if remainder >= 2 { remainder - 1 } else { 0 }
319}
320
321fn decode_value(settings: CodecSettings, byte: u8) -> u8 {
322 settings.alphabet().decode_byte(byte).unwrap_or(0)
323}