use std::io::{self, Seek, Write};
use byteorder::{BigEndian, ReadBytesExt, WriteBytesExt};
use bytes::{BufMut, Bytes, BytesMut};
use digest::DigestProcessor;
use rand::Rng;
use scuffle_bytes_util::BytesCursorExt;
use super::{RTMP_HANDSHAKE_SIZE, RtmpVersion, ServerHandshakeState, TIME_VERSION_LENGTH, current_time};
pub mod digest;
pub mod error;
pub const RTMP_SERVER_VERSION: u32 = 0x04050001;
pub const RTMP_DIGEST_LENGTH: usize = 32;
pub const RTMP_SERVER_KEY_FIRST_HALF: &[u8] = b"Genuine Adobe Flash Media Server 001";
pub const RTMP_CLIENT_KEY_FIRST_HALF: &[u8] = b"Genuine Adobe Flash Player 001";
pub const RTMP_SERVER_KEY: &[u8] = &[
0x47, 0x65, 0x6e, 0x75, 0x69, 0x6e, 0x65, 0x20, 0x41, 0x64, 0x6f, 0x62, 0x65, 0x20, 0x46, 0x6c, 0x61, 0x73, 0x68, 0x20,
0x4d, 0x65, 0x64, 0x69, 0x61, 0x20, 0x53, 0x65, 0x72, 0x76, 0x65, 0x72, 0x20, 0x30, 0x30, 0x31, 0xf0, 0xee, 0xc2, 0x4a,
0x80, 0x68, 0xbe, 0xe8, 0x2e, 0x00, 0xd0, 0xd1, 0x02, 0x9e, 0x7e, 0x57, 0x6e, 0xec, 0x5d, 0x2d, 0x29, 0x80, 0x6f, 0xab,
0x93, 0xb8, 0xe6, 0x36, 0xcf, 0xeb, 0x31, 0xae,
];
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
pub enum SchemaVersion {
Schema0,
Schema1,
}
pub struct ComplexHandshakeServer {
version: RtmpVersion,
requested_version: RtmpVersion,
state: ServerHandshakeState,
schema_version: SchemaVersion,
c1_digest: Bytes,
c1_timestamp: u32,
c1_version: u32,
}
impl Default for ComplexHandshakeServer {
fn default() -> Self {
Self {
state: ServerHandshakeState::ReadC0C1,
c1_digest: Bytes::default(),
c1_timestamp: 0,
version: RtmpVersion::Version3,
requested_version: RtmpVersion(0),
c1_version: 0,
schema_version: SchemaVersion::Schema0,
}
}
}
impl ComplexHandshakeServer {
pub fn is_finished(&self) -> bool {
self.state == ServerHandshakeState::Finish
}
pub fn handshake(&mut self, input: &mut io::Cursor<Bytes>, output: &mut Vec<u8>) -> Result<(), crate::error::RtmpError> {
match self.state {
ServerHandshakeState::ReadC0C1 => {
self.read_c0(input)?;
self.read_c1(input)?;
self.write_s0(output)?;
self.write_s1(output)?;
self.write_s2(output)?;
self.state = ServerHandshakeState::ReadC2;
}
ServerHandshakeState::ReadC2 => {
self.read_c2(input)?;
self.state = ServerHandshakeState::Finish;
}
ServerHandshakeState::Finish => {}
}
Ok(())
}
fn read_c0(&mut self, input: &mut io::Cursor<Bytes>) -> Result<(), crate::error::RtmpError> {
self.requested_version = RtmpVersion(input.read_u8()?);
self.version = RtmpVersion::Version3;
Ok(())
}
fn read_c1(&mut self, input: &mut io::Cursor<Bytes>) -> Result<(), crate::error::RtmpError> {
let c1_bytes = input.extract_bytes(RTMP_HANDSHAKE_SIZE)?;
self.c1_timestamp = (&c1_bytes[0..4]).read_u32::<BigEndian>()?;
self.c1_version = (&c1_bytes[4..8]).read_u32::<BigEndian>()?;
let data_digest = DigestProcessor::new(c1_bytes, RTMP_CLIENT_KEY_FIRST_HALF);
let (c1_digest_data, schema_version) = data_digest.read_digest()?;
self.c1_digest = c1_digest_data;
self.schema_version = schema_version;
Ok(())
}
fn read_c2(&mut self, input: &mut io::Cursor<Bytes>) -> Result<(), crate::error::RtmpError> {
input.seek_relative(RTMP_HANDSHAKE_SIZE as i64)?;
Ok(())
}
fn write_s0(&mut self, output: &mut Vec<u8>) -> Result<(), crate::error::RtmpError> {
output.write_u8(self.version.0)?;
Ok(())
}
fn write_s1(&self, output: &mut Vec<u8>) -> Result<(), crate::error::RtmpError> {
let mut writer = BytesMut::new().writer();
writer.write_u32::<BigEndian>(current_time())?;
writer.write_u32::<BigEndian>(RTMP_SERVER_VERSION)?;
let mut rng = rand::rng();
for _ in 0..RTMP_HANDSHAKE_SIZE - TIME_VERSION_LENGTH {
writer.write_u8(rng.random())?;
}
let data_digest = DigestProcessor::new(writer.into_inner().freeze(), RTMP_SERVER_KEY_FIRST_HALF);
data_digest.generate_and_fill_digest(self.schema_version)?.write_to(output)?;
Ok(())
}
fn write_s2(&self, output: &mut Vec<u8>) -> Result<(), crate::error::RtmpError> {
let start = output.len();
output.write_u32::<BigEndian>(current_time())?;
output.write_u32::<BigEndian>(self.c1_timestamp)?;
let mut rng = rand::rng();
for _ in 0..RTMP_HANDSHAKE_SIZE - RTMP_DIGEST_LENGTH - TIME_VERSION_LENGTH {
output.write_u8(rng.random())?;
}
let key_digest = DigestProcessor::new(Bytes::new(), RTMP_SERVER_KEY);
let key = key_digest.make_digest(&self.c1_digest, &[])?;
let data_digest = DigestProcessor::new(Bytes::new(), &key);
let digest = data_digest.make_digest(&output[start..start + RTMP_HANDSHAKE_SIZE - RTMP_DIGEST_LENGTH], &[])?;
output.write_all(&digest)?;
Ok(())
}
}