use super::error::PayloadError;
const MAX_KEY_LENGTH_BYTES: usize = u16::MAX as usize / 8;
pub fn construct_payload(
key: &[u8],
masked_key_length: usize,
cipher_block_length: usize,
random_seed: &[u8],
) -> Result<Vec<u8>, PayloadError> {
let key_len = key.len();
let key_length_bits = key_len
.checked_mul(8)
.and_then(|value| u16::try_from(value).ok())
.ok_or(PayloadError::KeyTooLong {
max: MAX_KEY_LENGTH_BYTES,
actual: key_len,
})?;
let padding_length = calculate_padding_length(key_len, masked_key_length, cipher_block_length)?;
let payload_capacity = key_len
.checked_add(2)
.and_then(|value| value.checked_add(padding_length))
.ok_or(PayloadError::InvalidTotalPayloadLength)?;
let mut payload = Vec::with_capacity(payload_capacity);
payload.extend_from_slice(&key_length_bits.to_be_bytes());
payload.extend_from_slice(key);
if random_seed.len() < padding_length {
return Err(PayloadError::RandomSeedTooShort {
required: padding_length,
actual: random_seed.len(),
});
}
payload.extend_from_slice(&random_seed[..padding_length]);
Ok(payload)
}
pub fn extract_key_from_payload(payload: &[u8]) -> Result<Vec<u8>, PayloadError> {
const KEY_LENGTH_FIELD_SIZE: usize = 2;
if payload.len() < KEY_LENGTH_FIELD_SIZE {
return Err(PayloadError::PayloadTooShort {
minimum: KEY_LENGTH_FIELD_SIZE,
actual: payload.len(),
});
}
let key_length_bits = u16::from_be_bytes([payload[0], payload[1]]);
let key_length_bytes = (key_length_bits / 8) as usize;
let required_length = KEY_LENGTH_FIELD_SIZE
.checked_add(key_length_bytes)
.ok_or(PayloadError::InvalidTotalPayloadLength)?;
if payload.len() < required_length {
return Err(PayloadError::PayloadTooShortForKey {
required: required_length,
actual: payload.len(),
});
}
Ok(payload[KEY_LENGTH_FIELD_SIZE..required_length].to_vec())
}
pub fn calculate_padding_length(
key_len: usize,
masked_key_length: usize,
cipher_block_length: usize,
) -> Result<usize, PayloadError> {
if cipher_block_length == 0 {
return Err(PayloadError::InvalidCipherBlockLength);
}
let raw_key_section_length = 2usize
.checked_add(key_len)
.ok_or(PayloadError::InvalidTotalPayloadLength)?;
let effective_key_length = std::cmp::max(key_len, masked_key_length);
let length_to_round = 2usize
.checked_add(effective_key_length)
.and_then(|value| value.checked_add(cipher_block_length - 1))
.ok_or(PayloadError::InvalidTotalPayloadLength)?;
let block_count = length_to_round / cipher_block_length;
let total_payload_length = block_count
.checked_mul(cipher_block_length)
.ok_or(PayloadError::InvalidTotalPayloadLength)?;
if total_payload_length < raw_key_section_length {
return Err(PayloadError::InvalidTotalPayloadLength);
}
Ok(total_payload_length - raw_key_section_length)
}
#[test]
fn test_calculate_padding_length_zero_block_length() {
let result = calculate_padding_length(16, 0, 0);
assert_eq!(result, Err(PayloadError::InvalidCipherBlockLength));
}
#[test]
fn test_construct_payload_random_seed_too_short() {
let key = [0u8; 16];
let result = construct_payload(&key, 0, 16, &[]);
assert!(matches!(
result,
Err(PayloadError::RandomSeedTooShort {
required,
actual: 0,
}) if required > 0
));
}