use sha1::{Digest as Sha1Digest, Sha1};
use sha2::Sha512;
use crate::{Error, Result, AUTH_KEY_SIZE};
pub const LOCAL_ENCRYPT_SALT_SIZE: usize = 32;
pub const AES_KEY_SIZE: usize = 32;
pub const AES_BLOCK_SIZE: usize = 16;
const PBKDF2_ITERATIONS_WITH_PASSCODE: u32 = 100_000;
const PBKDF2_ITERATIONS_NO_PASSCODE: u32 = 1;
#[derive(Clone)]
pub struct AuthKey {
data: [u8; AUTH_KEY_SIZE],
}
impl AuthKey {
pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
if bytes.len() != AUTH_KEY_SIZE {
return Err(Error::invalid_format(format!(
"auth key must be {} bytes, got {}",
AUTH_KEY_SIZE,
bytes.len()
)));
}
let mut data = [0u8; AUTH_KEY_SIZE];
data.copy_from_slice(bytes);
Ok(Self { data })
}
pub fn as_bytes(&self) -> &[u8; AUTH_KEY_SIZE] {
&self.data
}
}
impl std::fmt::Debug for AuthKey {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AuthKey")
.field("len", &self.data.len())
.finish()
}
}
pub fn create_local_key(salt: &[u8], passcode: &[u8]) -> AuthKey {
let mut key_data = [0u8; AUTH_KEY_SIZE];
let mut hasher = Sha512::new();
hasher.update(salt);
hasher.update(passcode);
hasher.update(salt);
let hash_key = hasher.finalize();
let iterations = if passcode.is_empty() {
PBKDF2_ITERATIONS_NO_PASSCODE
} else {
PBKDF2_ITERATIONS_WITH_PASSCODE
};
pbkdf2::pbkdf2_hmac::<Sha512>(&hash_key, salt, iterations, &mut key_data);
AuthKey { data: key_data }
}
pub fn decrypt_local(encrypted: &[u8], key: &AuthKey) -> Result<Vec<u8>> {
if encrypted.len() <= AES_BLOCK_SIZE {
return Err(Error::invalid_format("encrypted data too short"));
}
if encrypted.len() % AES_BLOCK_SIZE != 0 {
return Err(Error::invalid_format(
"encrypted data length must be multiple of 16",
));
}
let (encrypted_key, encrypted_data) = encrypted.split_at(AES_BLOCK_SIZE);
let encrypted_key: &[u8; AES_BLOCK_SIZE] = encrypted_key
.try_into()
.map_err(|_| Error::invalid_format("invalid encrypted message key"))?;
tracing::debug!("decrypt_local: encrypted len={}", encrypted.len());
let (aes_key, aes_iv) = prepare_aes_oldmtp(key.as_bytes(), encrypted_key)?;
let decrypted = ige_decrypt(&aes_key, &aes_iv, encrypted_data);
let digest = sha1_hash(&decrypted);
let (check_hash, _) = digest.split_at(AES_BLOCK_SIZE);
tracing::debug!("Computed decrypted payload integrity check");
if check_hash != encrypted_key {
return Err(Error::ChecksumMismatch);
}
if decrypted.len() < 4 {
return Err(Error::DecryptionFailed);
}
let original_len_bytes: [u8; 4] = decrypted
.get(..4)
.ok_or(Error::DecryptionFailed)?
.try_into()
.map_err(|_| Error::DecryptionFailed)?;
let original_len = usize::try_from(u32::from_le_bytes(original_len_bytes))
.map_err(|_| Error::invalid_format("decrypted payload length is too large"))?;
let full_len = encrypted_data.len();
if original_len > decrypted.len()
|| original_len <= full_len.saturating_sub(16)
|| original_len < 4
{
return Err(Error::invalid_format(format!(
"invalid decrypted length: {}, full_len: {}, decrypted size: {}",
original_len,
full_len,
decrypted.len()
)));
}
decrypted
.get(4..original_len)
.map(ToOwned::to_owned)
.ok_or_else(|| Error::invalid_format("invalid decrypted payload bounds"))
}
fn prepare_aes_oldmtp(
auth_key: &[u8; AUTH_KEY_SIZE],
msg_key: &[u8; AES_BLOCK_SIZE],
) -> Result<([u8; AES_KEY_SIZE], [u8; AES_KEY_SIZE])> {
let auth_range = |range| {
auth_key
.get(range)
.ok_or_else(|| Error::invalid_format("auth key is too short for AES derivation"))
};
let sha1_a = sha1_hash_2(msg_key, auth_range(8..40)?);
let sha1_b = sha1_hash_3(auth_range(40..56)?, msg_key, auth_range(56..72)?);
let sha1_c = sha1_hash_2(auth_range(72..104)?, msg_key);
let sha1_d = sha1_hash_2(msg_key, auth_range(104..136)?);
let key_bytes: Vec<u8> = sha1_a
.iter()
.take(8)
.chain(sha1_b.iter().skip(8).take(12))
.chain(sha1_c.iter().skip(4).take(12))
.copied()
.collect();
let key = key_bytes
.try_into()
.map_err(|_| Error::invalid_format("failed to derive AES key"))?;
let iv_bytes: Vec<u8> = sha1_a
.iter()
.skip(8)
.take(12)
.chain(sha1_b.iter().take(8))
.chain(sha1_c.iter().skip(16).take(4))
.chain(sha1_d.iter().take(8))
.copied()
.collect();
let iv = iv_bytes
.try_into()
.map_err(|_| Error::invalid_format("failed to derive AES IV"))?;
Ok((key, iv))
}
fn ige_decrypt(key: &[u8; 32], iv: &[u8; 32], data: &[u8]) -> Vec<u8> {
use grammers_crypto::aes::ige_decrypt as grammers_ige_decrypt;
let mut decrypted = data.to_vec();
grammers_ige_decrypt(&mut decrypted, key, iv);
decrypted
}
fn sha1_hash(data: &[u8]) -> [u8; 20] {
let mut hasher = Sha1::new();
hasher.update(data);
hasher.finalize().into()
}
fn sha1_hash_2(a: &[u8], b: &[u8]) -> [u8; 20] {
let mut hasher = Sha1::new();
hasher.update(a);
hasher.update(b);
hasher.finalize().into()
}
fn sha1_hash_3(a: &[u8], b: &[u8], c: &[u8]) -> [u8; 20] {
let mut hasher = Sha1::new();
hasher.update(a);
hasher.update(b);
hasher.update(c);
hasher.finalize().into()
}
#[cfg(test)]
mod tests {
use super::*;
fn fixture_salt() -> [u8; LOCAL_ENCRYPT_SALT_SIZE] {
std::array::from_fn(|index| u8::try_from(index).unwrap_or_default())
}
#[test]
fn test_create_local_key_no_passcode() {
let salt = fixture_salt();
let passcode = b"";
let key = create_local_key(&salt, passcode);
assert_eq!(key.as_bytes().len(), AUTH_KEY_SIZE);
}
#[test]
fn test_create_local_key_with_passcode() {
let salt = fixture_salt();
let passcode = b"test";
let key = create_local_key(&salt, passcode);
assert_eq!(key.as_bytes().len(), AUTH_KEY_SIZE);
let key2 = create_local_key(&salt, passcode);
assert_eq!(key.as_bytes(), key2.as_bytes());
}
#[test]
fn test_auth_key_from_bytes() -> Result<()> {
let bytes = [0xAB; AUTH_KEY_SIZE];
let key = AuthKey::from_bytes(&bytes)?;
assert_eq!(key.as_bytes(), &bytes);
Ok(())
}
#[test]
fn test_auth_key_wrong_size() {
let bytes = [0u8; 100];
assert!(AuthKey::from_bytes(&bytes).is_err());
}
#[test]
fn test_sha1_hash() {
let data = b"hello";
let hash = sha1_hash(data);
assert_eq!(
hex::encode(hash),
"aaf4c61ddcc5e8a2dabede0f3b482cd9aea9434d"
);
}
}