Skip to main content

pb_mapper_protocol/secure/
frame.rs

1//! Protocol-v2 key derivation and directional authenticated frame codecs.
2//!
3//! ```text
4//! credential + connection salt -> HKDF -> c2s key / s2c key
5//! plaintext -> counter + length + AEAD(AAD) -> encrypted frame
6//! encrypted frame -> bound length -> verify counter/tag -> plaintext
7//! ```
8//!
9//! Counters are monotonic per direction and are included in both the nonce and AAD.
10//! The initial reader can impose a smaller pre-authentication limit before allocating
11//! a body; continuation frames retain the normal protocol maximum.
12
13use pb_mapper_auth::KeyId;
14
15use super::*;
16
17#[derive(Clone)]
18pub(super) struct V2Material {
19    pub(super) key_id: KeyId,
20    pub(super) flags: u8,
21    pub(super) salt: [u8; CONNECTION_SALT_LEN],
22    pub(super) client_to_server: AesKeyType,
23    pub(super) server_to_client: AesKeyType,
24}
25
26pub struct V2MessageReader<'a, T: AsyncReadExt + Unpin> {
27    reader: &'a mut T,
28    material: V2Material,
29    key: LessSafeKey,
30    direction: u8,
31    expected_counter: u64,
32    buffer: Vec<u8>,
33}
34
35impl<'a, T: AsyncReadExt + Unpin> V2MessageReader<'a, T> {
36    pub(super) fn new(
37        reader: &'a mut T,
38        material: V2Material,
39        direction: u8,
40        expected_counter: u64,
41    ) -> Result<Self> {
42        let key_bytes = direction_key(&material, direction);
43        let key = LessSafeKey::new(
44            UnboundKey::new(&AES_256_GCM, key_bytes)
45                .map_err(|_| protocol_error("invalid protocol-v2 read key"))?,
46        );
47        Ok(Self {
48            reader,
49            material,
50            key,
51            direction,
52            expected_counter,
53            buffer: Vec::new(),
54        })
55    }
56
57    pub(super) async fn read_msg_with_limit(&mut self, max_plaintext_len: u32) -> Result<&'_ [u8]> {
58        let (counter, ciphertext) =
59            read_v2_frame(self.reader, self.expected_counter, max_plaintext_len).await?;
60        self.buffer = ciphertext;
61        let datalen = u32::try_from(self.buffer.len())
62            .map_err(|_| protocol_error("protocol-v2 payload exceeds u32 length"))?;
63        let aad = frame_aad(&self.material, self.direction, counter, datalen);
64        let plain = self
65            .key
66            .open_in_place(nonce(counter), Aad::from(aad.as_slice()), &mut self.buffer)
67            .map_err(|_| protocol_error("protocol-v2 payload authentication failed"))?;
68        let plain_len = plain.len();
69        self.buffer.truncate(plain_len);
70        self.expected_counter = self
71            .expected_counter
72            .checked_add(1)
73            .ok_or_else(|| protocol_error("protocol-v2 receive counter exhausted"))?;
74        Ok(&self.buffer)
75    }
76}
77
78pub(super) async fn read_v2_frame<T: AsyncReadExt + Unpin>(
79    reader: &mut T,
80    expected_counter: u64,
81    max_plaintext_len: u32,
82) -> Result<(u64, Vec<u8>)> {
83    let counter = reader
84        .read_u64()
85        .await
86        .map_err(|error| protocol_error(format!("failed to read v2 counter: {error}")))?;
87    if counter != expected_counter {
88        return Err(protocol_error(format!(
89            "protocol-v2 counter mismatch: expected {expected_counter}, got {counter}"
90        )));
91    }
92    let datalen = reader
93        .read_u32()
94        .await
95        .map_err(|error| protocol_error(format!("failed to read v2 length: {error}")))?;
96    let max_encrypted_len = max_plaintext_len.saturating_add(AES_256_GCM.tag_len() as u32);
97    if datalen < AES_256_GCM.tag_len() as u32 || datalen > max_encrypted_len {
98        return Err(protocol_error(format!(
99            "protocol-v2 payload length {datalen} exceeds the {max_plaintext_len}-byte limit"
100        )));
101    }
102    let mut ciphertext = vec![0_u8; datalen as usize];
103    reader
104        .read_exact(&mut ciphertext)
105        .await
106        .map_err(|error| protocol_error(format!("failed to read v2 payload: {error}")))?;
107    Ok((counter, ciphertext))
108}
109
110impl<T: AsyncReadExt + Unpin> MessageReader for V2MessageReader<'_, T> {
111    async fn read_msg(&mut self) -> Result<&'_ [u8]> {
112        self.read_msg_with_limit(MAX_MSG_LEN - AES_256_GCM.tag_len() as u32)
113            .await
114    }
115}
116
117pub struct V2MessageWriter<'a, T: AsyncWriteExt + Unpin> {
118    writer: &'a mut T,
119    material: V2Material,
120    key: LessSafeKey,
121    direction: u8,
122    counter: u64,
123}
124
125impl<'a, T: AsyncWriteExt + Unpin> V2MessageWriter<'a, T> {
126    pub(super) fn new(
127        writer: &'a mut T,
128        material: V2Material,
129        direction: u8,
130        counter: u64,
131    ) -> Result<Self> {
132        let key_bytes = direction_key(&material, direction);
133        let key = LessSafeKey::new(
134            UnboundKey::new(&AES_256_GCM, key_bytes)
135                .map_err(|_| protocol_error("invalid protocol-v2 write key"))?,
136        );
137        Ok(Self {
138            writer,
139            material,
140            key,
141            direction,
142            counter,
143        })
144    }
145}
146
147impl<T: AsyncWriteExt + Unpin> MessageWriter for V2MessageWriter<'_, T> {
148    async fn write_msg(&mut self, message: &[u8]) -> Result<()> {
149        let encrypted_len = message
150            .len()
151            .checked_add(AES_256_GCM.tag_len())
152            .and_then(|len| DataLenType::try_from(len).ok())
153            .ok_or_else(|| protocol_error("protocol-v2 message is too large"))?;
154        if encrypted_len > MAX_MSG_LEN {
155            return Err(protocol_error(
156                "protocol-v2 message exceeds the maximum length",
157            ));
158        }
159        let counter = self.counter;
160        let aad = frame_aad(&self.material, self.direction, counter, encrypted_len);
161        let mut encrypted = message.to_vec();
162        self.key
163            .seal_in_place_append_tag(nonce(counter), Aad::from(aad.as_slice()), &mut encrypted)
164            .map_err(|_| protocol_error("failed to encrypt protocol-v2 message"))?;
165        self.writer
166            .write_u64(counter)
167            .await
168            .map_err(|error| protocol_error(format!("failed to write v2 frame header: {error}")))?;
169        self.writer
170            .write_u32(encrypted_len)
171            .await
172            .map_err(|error| protocol_error(format!("failed to write v2 frame header: {error}")))?;
173        self.writer
174            .write_all(&encrypted)
175            .await
176            .map_err(|error| protocol_error(format!("failed to write v2 frame body: {error}")))?;
177        self.counter = self
178            .counter
179            .checked_add(1)
180            .ok_or_else(|| protocol_error("protocol-v2 send counter exhausted"))?;
181        Ok(())
182    }
183}
184
185pub(super) fn open_v2_payload(
186    material: &V2Material,
187    direction: u8,
188    counter: u64,
189    ciphertext: &mut [u8],
190) -> Result<Vec<u8>> {
191    let key_bytes = direction_key(material, direction);
192    let key = LessSafeKey::new(
193        UnboundKey::new(&AES_256_GCM, key_bytes)
194            .map_err(|_| protocol_error("invalid protocol-v2 read key"))?,
195    );
196    let datalen = u32::try_from(ciphertext.len())
197        .map_err(|_| protocol_error("protocol-v2 payload length is invalid"))?;
198    let aad = frame_aad(material, direction, counter, datalen);
199    let plain = key
200        .open_in_place(nonce(counter), Aad::from(aad.as_slice()), ciphertext)
201        .map_err(|_| protocol_error("protocol-v2 payload authentication failed"))?;
202    Ok(plain.to_vec())
203}
204
205pub(super) fn derive_material(
206    key_id: KeyId,
207    credential_key: &AesKeyType,
208    salt_bytes: [u8; CONNECTION_SALT_LEN],
209) -> Result<V2Material> {
210    let salt = Salt::new(HKDF_SHA256, &salt_bytes);
211    let pseudo_random_key = salt.extract(credential_key);
212    let client_to_server = expand_direction(&pseudo_random_key, b"pb-mapper-v2-c2s")?;
213    let server_to_client = expand_direction(&pseudo_random_key, b"pb-mapper-v2-s2c")?;
214    Ok(V2Material {
215        key_id,
216        flags: 0,
217        salt: salt_bytes,
218        client_to_server,
219        server_to_client,
220    })
221}
222
223fn expand_direction(
224    pseudo_random_key: &ring::hkdf::Prk,
225    label: &'static [u8],
226) -> Result<AesKeyType> {
227    let info = [label];
228    let output = pseudo_random_key
229        .expand(&info, HkdfLen(32))
230        .map_err(|_| protocol_error("failed to derive protocol-v2 direction key"))?;
231    let mut key = [0_u8; 32];
232    output
233        .fill(&mut key)
234        .map_err(|_| protocol_error("failed to fill protocol-v2 direction key"))?;
235    Ok(key)
236}
237
238struct HkdfLen(usize);
239
240impl ring::hkdf::KeyType for HkdfLen {
241    fn len(&self) -> usize {
242        self.0
243    }
244}
245
246fn direction_key(material: &V2Material, direction: u8) -> &AesKeyType {
247    if direction == DIRECTION_CLIENT_TO_SERVER {
248        &material.client_to_server
249    } else {
250        &material.server_to_client
251    }
252}
253
254pub(super) fn first_prefix(material: &V2Material) -> Vec<u8> {
255    let mut prefix = Vec::with_capacity(PROTOCOL_V2_MAGIC.len() + FIRST_PREFIX_REMAINDER_LEN);
256    prefix.extend_from_slice(&PROTOCOL_V2_MAGIC);
257    prefix.push(PROTOCOL_V2_VERSION);
258    prefix.push(material.flags);
259    prefix.extend_from_slice(&0_u16.to_be_bytes());
260    prefix.extend_from_slice(&material.key_id.to_be_bytes());
261    prefix.extend_from_slice(&material.salt);
262    prefix
263}
264
265fn frame_aad(material: &V2Material, direction: u8, counter: u64, datalen: u32) -> Vec<u8> {
266    let mut aad = Vec::with_capacity(
267        PROTOCOL_V2_MAGIC.len() + FIRST_PREFIX_REMAINDER_LEN + 1 + FRAME_HEADER_LEN,
268    );
269    aad.extend_from_slice(&first_prefix(material));
270    aad.push(direction);
271    aad.extend_from_slice(&counter.to_be_bytes());
272    aad.extend_from_slice(&datalen.to_be_bytes());
273    aad
274}
275
276fn nonce(counter: u64) -> Nonce {
277    let mut bytes = [0_u8; 12];
278    bytes[4..].copy_from_slice(&counter.to_be_bytes());
279    Nonce::assume_unique_for_key(bytes)
280}