use std::io;
use bytes::Bytes;
use hmac::{Hmac, Mac};
use sha2::Sha256;
use super::error::ComplexHandshakeError;
use super::{RTMP_DIGEST_LENGTH, SchemaVersion};
use crate::handshake::{CHUNK_LENGTH, TIME_VERSION_LENGTH};
pub struct DigestProcessor<'a> {
data: Bytes,
key: &'a [u8],
}
pub struct DigestResult {
pub left: Bytes,
pub digest: [u8; 32],
pub right: Bytes,
}
impl DigestResult {
pub fn write_to(&self, writer: &mut impl io::Write) -> io::Result<()> {
writer.write_all(&self.left)?;
writer.write_all(&self.digest)?;
writer.write_all(&self.right)?;
Ok(())
}
}
impl<'a> DigestProcessor<'a> {
pub const fn new(data: Bytes, key: &'a [u8]) -> Self {
Self { data, key }
}
pub fn read_digest(&self) -> Result<(Bytes, SchemaVersion), ComplexHandshakeError> {
if let Ok(digest) = self.generate_and_validate(SchemaVersion::Schema0) {
Ok((digest, SchemaVersion::Schema0))
} else {
let digest = self.generate_and_validate(SchemaVersion::Schema1)?;
Ok((digest, SchemaVersion::Schema1))
}
}
pub fn generate_and_fill_digest(&self, version: SchemaVersion) -> Result<DigestResult, ComplexHandshakeError> {
let (left_part, _, right_part) = self.split_message(version)?;
let computed_digest = self.make_digest(&left_part, &right_part)?;
Ok(DigestResult {
left: left_part,
digest: computed_digest,
right: right_part,
})
}
fn find_digest_offset(&self, version: SchemaVersion) -> Result<usize, ComplexHandshakeError> {
const OFFSET_LENGTH: usize = 4;
let schema_offset = match version {
SchemaVersion::Schema0 => CHUNK_LENGTH + TIME_VERSION_LENGTH,
SchemaVersion::Schema1 => TIME_VERSION_LENGTH,
};
Ok((*self.data.get(schema_offset).unwrap() as usize
+ *self.data.get(schema_offset + 1).unwrap() as usize
+ *self.data.get(schema_offset + 2).unwrap() as usize
+ *self.data.get(schema_offset + 3).unwrap() as usize)
% (CHUNK_LENGTH - RTMP_DIGEST_LENGTH - OFFSET_LENGTH)
+ schema_offset
+ OFFSET_LENGTH)
}
fn split_message(&self, version: SchemaVersion) -> Result<(Bytes, Bytes, Bytes), ComplexHandshakeError> {
let digest_offset = self.find_digest_offset(version)?;
let left_part = self.data.slice(0..digest_offset);
let digest_data = self.data.slice(digest_offset..digest_offset + RTMP_DIGEST_LENGTH);
let right_part = self.data.slice(digest_offset + RTMP_DIGEST_LENGTH..);
Ok((left_part, digest_data, right_part))
}
pub fn make_digest(&self, left: &[u8], right: &[u8]) -> Result<[u8; 32], ComplexHandshakeError> {
let mut mac = Hmac::<Sha256>::new_from_slice(self.key).unwrap();
mac.update(left);
mac.update(right);
let result = mac.finalize().into_bytes();
if result.len() != RTMP_DIGEST_LENGTH {
return Err(ComplexHandshakeError::DigestLengthNotCorrect);
}
Ok(result.into())
}
fn generate_and_validate(&self, version: SchemaVersion) -> Result<Bytes, ComplexHandshakeError> {
let (left_part, digest_data, right_part) = self.split_message(version)?;
if digest_data == self.make_digest(&left_part, &right_part)?.as_ref() {
Ok(digest_data)
} else {
Err(ComplexHandshakeError::CannotGenerate)
}
}
}