pb_mapper_protocol/secure/
frame.rs1use 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}