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