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}
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}