use crate::{option::*, protection_profile::*};
use shared::error::{Error, Result};
pub const LABEL_EXTRACTOR_DTLS_SRTP: &str = "EXTRACTOR-dtls_srtp";
#[derive(Default, Debug, Clone)]
pub struct SessionKeys {
pub local_master_key: Vec<u8>,
pub local_master_salt: Vec<u8>,
pub remote_master_key: Vec<u8>,
pub remote_master_salt: Vec<u8>,
}
#[derive(Default)]
pub struct Config {
pub keys: SessionKeys,
pub profile: ProtectionProfile,
pub local_rtp_options: Option<ContextOption>,
pub remote_rtp_options: Option<ContextOption>,
pub local_rtcp_options: Option<ContextOption>,
pub remote_rtcp_options: Option<ContextOption>,
}
impl Config {
#[must_use]
pub fn keying_material_len(&self) -> usize {
let key_len = self.profile.key_len();
let salt_len = self.profile.salt_len();
(key_len * 2) + (salt_len * 2)
}
pub fn set_session_keys_from_keying_material(
&mut self,
keying_material: &[u8],
is_client: bool,
) -> Result<()> {
let expected = self.keying_material_len();
if keying_material.len() != expected {
return Err(Error::Other(format!(
"invalid DTLS-SRTP keying material length: expected {expected}, got {}",
keying_material.len()
)));
}
let key_len = self.profile.key_len();
let salt_len = self.profile.salt_len();
let mut offset = 0;
let client_write_key = keying_material[offset..offset + key_len].to_vec();
offset += key_len;
let server_write_key = keying_material[offset..offset + key_len].to_vec();
offset += key_len;
let client_write_salt = keying_material[offset..offset + salt_len].to_vec();
offset += salt_len;
let server_write_salt = keying_material[offset..offset + salt_len].to_vec();
if is_client {
self.keys.local_master_key = client_write_key;
self.keys.local_master_salt = client_write_salt;
self.keys.remote_master_key = server_write_key;
self.keys.remote_master_salt = server_write_salt;
} else {
self.keys.local_master_key = server_write_key;
self.keys.local_master_salt = server_write_salt;
self.keys.remote_master_key = client_write_key;
self.keys.remote_master_salt = client_write_salt;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn material(len: usize) -> Vec<u8> {
(0..len).map(|value| value as u8).collect()
}
#[test]
fn rejects_keying_material_with_the_wrong_length() {
let mut config = Config {
profile: ProtectionProfile::Aes128CmHmacSha1_80,
..Default::default()
};
let expected = config.keying_material_len();
for actual in [expected - 1, expected + 1] {
let error = config
.set_session_keys_from_keying_material(&material(actual), true)
.unwrap_err();
assert!(
error
.to_string()
.contains(&format!("expected {expected}, got {actual}"))
);
}
}
#[test]
fn assigns_client_and_server_material_by_role() -> Result<()> {
let mut client = Config {
profile: ProtectionProfile::Aes128CmHmacSha1_80,
..Default::default()
};
let bytes = material(client.keying_material_len());
client.set_session_keys_from_keying_material(&bytes, true)?;
let mut server = Config {
profile: client.profile,
..Default::default()
};
server.set_session_keys_from_keying_material(&bytes, false)?;
assert_eq!(client.keys.local_master_key, server.keys.remote_master_key);
assert_eq!(
client.keys.local_master_salt,
server.keys.remote_master_salt
);
assert_eq!(client.keys.remote_master_key, server.keys.local_master_key);
assert_eq!(
client.keys.remote_master_salt,
server.keys.local_master_salt
);
Ok(())
}
}