2020-08-14 16:20:42 -04:00
|
|
|
// Copyright (c) Facebook, Inc. and its affiliates.
|
|
|
|
|
//
|
|
|
|
|
// This source code is licensed under the MIT license found in the
|
|
|
|
|
// LICENSE file in the root directory of this source tree.
|
2020-10-09 10:51:58 -07:00
|
|
|
|
2020-08-14 16:20:42 -04:00
|
|
|
use crate::errors::PakeError;
|
|
|
|
|
|
2020-11-12 10:00:09 -08:00
|
|
|
// Corresponds to the I2OSP() function from RFC8017
|
2021-07-08 11:04:32 -07:00
|
|
|
pub(crate) fn i2osp(input: usize, length: usize) -> Result<Vec<u8>, PakeError> {
|
|
|
|
|
let sizeof_usize = std::mem::size_of::<usize>();
|
|
|
|
|
|
|
|
|
|
// Check if input >= 256^length
|
|
|
|
|
if (sizeof_usize as u32 - input.leading_zeros() / 8) > length as u32 {
|
|
|
|
|
return Err(PakeError::SerializationError);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if length <= sizeof_usize {
|
|
|
|
|
return Ok((&input.to_be_bytes()[sizeof_usize - length..]).to_vec());
|
2020-11-12 10:00:09 -08:00
|
|
|
}
|
|
|
|
|
|
|
|
|
|
let mut output = vec![0u8; length];
|
|
|
|
|
output.splice(
|
2021-07-08 11:04:32 -07:00
|
|
|
length - sizeof_usize..length,
|
2020-11-12 10:00:09 -08:00
|
|
|
input.to_be_bytes().iter().cloned(),
|
|
|
|
|
);
|
2021-07-08 11:04:32 -07:00
|
|
|
Ok(output)
|
2020-08-14 16:20:42 -04:00
|
|
|
}
|
|
|
|
|
|
2020-11-12 10:00:09 -08:00
|
|
|
// Corresponds to the OS2IP() function from RFC8017
|
|
|
|
|
pub(crate) fn os2ip(input: &[u8]) -> Result<usize, PakeError> {
|
|
|
|
|
if input.len() > std::mem::size_of::<usize>() {
|
2020-08-14 16:20:42 -04:00
|
|
|
return Err(PakeError::SerializationError);
|
|
|
|
|
}
|
|
|
|
|
|
2020-11-12 10:00:09 -08:00
|
|
|
let mut output_array = [0u8; std::mem::size_of::<usize>()];
|
|
|
|
|
output_array[std::mem::size_of::<usize>() - input.len()..].copy_from_slice(input);
|
|
|
|
|
Ok(usize::from_be_bytes(output_array))
|
|
|
|
|
}
|
2020-11-12 09:56:06 -05:00
|
|
|
|
2020-11-12 10:00:09 -08:00
|
|
|
// Computes I2OSP(len(input), max_bytes) || input
|
2021-07-08 11:04:32 -07:00
|
|
|
pub(crate) fn serialize(input: &[u8], max_bytes: usize) -> Result<Vec<u8>, PakeError> {
|
|
|
|
|
Ok([&i2osp(input.len(), max_bytes)?, input].concat())
|
2020-11-12 10:00:09 -08:00
|
|
|
}
|
2020-11-12 09:56:06 -05:00
|
|
|
|
2020-11-12 10:00:09 -08:00
|
|
|
// Tokenizes an input of the format I2OSP(len(input), max_bytes) || input, outputting
|
|
|
|
|
// (input, remainder)
|
2020-11-16 11:49:27 -08:00
|
|
|
pub(crate) fn tokenize(input: &[u8], size_bytes: usize) -> Result<(Vec<u8>, Vec<u8>), PakeError> {
|
2020-11-12 10:00:09 -08:00
|
|
|
if size_bytes > std::mem::size_of::<usize>() || input.len() < size_bytes {
|
2020-11-12 09:56:06 -05:00
|
|
|
return Err(PakeError::SerializationError);
|
|
|
|
|
}
|
|
|
|
|
|
2020-11-12 10:00:09 -08:00
|
|
|
let size = os2ip(&input[..size_bytes])?;
|
2020-08-14 16:20:42 -04:00
|
|
|
if size_bytes + size > input.len() {
|
|
|
|
|
return Err(PakeError::SerializationError);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
Ok((
|
|
|
|
|
input[size_bytes..size_bytes + size].to_vec(),
|
|
|
|
|
input[size_bytes + size..].to_vec(),
|
|
|
|
|
))
|
|
|
|
|
}
|
|
|
|
|
|
2021-06-10 21:48:38 +02:00
|
|
|
/// Inner macro used for deriving `serde`'s `Serialize` and `Deserialize` traits.
|
|
|
|
|
macro_rules! impl_serialize_and_deserialize_for {
|
|
|
|
|
($t:ident) => {
|
|
|
|
|
#[cfg(feature = "serialize")]
|
|
|
|
|
impl<CS: CipherSuite> serde::Serialize for $t<CS> {
|
|
|
|
|
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
|
|
|
|
where
|
|
|
|
|
S: serde::Serializer,
|
|
|
|
|
{
|
|
|
|
|
if serializer.is_human_readable() {
|
|
|
|
|
serializer.serialize_str(&base64::encode(&self.serialize()))
|
|
|
|
|
} else {
|
|
|
|
|
serializer.serialize_bytes(&self.serialize())
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
#[cfg(feature = "serialize")]
|
|
|
|
|
impl<'de, CS: CipherSuite> serde::Deserialize<'de> for $t<CS> {
|
|
|
|
|
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
|
|
|
|
where
|
|
|
|
|
D: serde::Deserializer<'de>,
|
|
|
|
|
{
|
|
|
|
|
if deserializer.is_human_readable() {
|
|
|
|
|
let s = <&str>::deserialize(deserializer)?;
|
|
|
|
|
$t::<CS>::deserialize(&base64::decode(s).map_err(serde::de::Error::custom)?)
|
|
|
|
|
.map_err(serde::de::Error::custom)
|
|
|
|
|
} else {
|
|
|
|
|
struct ByteVisitor<CS: CipherSuite> {
|
|
|
|
|
marker: std::marker::PhantomData<CS>,
|
|
|
|
|
}
|
|
|
|
|
impl<'de, CS: CipherSuite> serde::de::Visitor<'de> for ByteVisitor<CS> {
|
|
|
|
|
type Value = $t<CS>;
|
|
|
|
|
fn expecting(
|
|
|
|
|
&self,
|
|
|
|
|
formatter: &mut std::fmt::Formatter,
|
|
|
|
|
) -> std::fmt::Result {
|
|
|
|
|
formatter.write_str(std::concat!(
|
|
|
|
|
"the byte representation of a ",
|
|
|
|
|
std::stringify!($t)
|
|
|
|
|
))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
fn visit_bytes<E>(self, value: &[u8]) -> Result<Self::Value, E>
|
|
|
|
|
where
|
|
|
|
|
E: serde::de::Error,
|
|
|
|
|
{
|
|
|
|
|
$t::<CS>::deserialize(value).map_err(|_| {
|
|
|
|
|
serde::de::Error::invalid_value(
|
|
|
|
|
serde::de::Unexpected::Bytes(value),
|
|
|
|
|
&std::concat!(
|
|
|
|
|
"invalid byte sequence for ",
|
|
|
|
|
std::stringify!($t)
|
|
|
|
|
),
|
|
|
|
|
)
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
deserializer.deserialize_bytes(ByteVisitor::<CS> {
|
|
|
|
|
marker: std::marker::PhantomData,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
};
|
|
|
|
|
}
|
|
|
|
|
|
2020-08-14 16:20:42 -04:00
|
|
|
#[cfg(test)]
|
|
|
|
|
mod tests;
|
2021-07-08 11:04:32 -07:00
|
|
|
|
|
|
|
|
#[cfg(test)]
|
|
|
|
|
mod unit_tests {
|
|
|
|
|
use super::*;
|
|
|
|
|
|
|
|
|
|
// Test the error condition for I2OSP
|
|
|
|
|
#[test]
|
|
|
|
|
fn test_i2osp_err_check() {
|
|
|
|
|
assert!(i2osp(0, 1).is_ok());
|
|
|
|
|
|
|
|
|
|
assert!(i2osp(255, 1).is_ok());
|
|
|
|
|
assert!(i2osp(256, 1).is_err());
|
|
|
|
|
assert!(i2osp(257, 1).is_err());
|
|
|
|
|
|
|
|
|
|
assert!(i2osp(256 * 256 - 1, 2).is_ok());
|
|
|
|
|
assert!(i2osp(256 * 256, 2).is_err());
|
|
|
|
|
assert!(i2osp(256 * 256 + 1, 2).is_err());
|
|
|
|
|
}
|
|
|
|
|
}
|