1use crate::crypto::{
2 aead_decrypt, aead_encrypt, sha256_digest, AEAD_TAG_LEN,
3};
4use crate::error::MoqSecureError;
5use crate::key_store::KeyStore;
6use ed25519_dalek::{Signer, Verifier};
7
8pub const MAGIC: [u8; 4] = *b"MOQS";
9pub const VERSION: u8 = 1;
10
11pub const SIG_SLOT_LEN: usize = 64;
12pub const PAD_LEN_FIELD_LEN: usize = 4;
13
14pub const FIXED_HEADER_LEN: usize = 4 + 1 + 1 + 8 + 1 + 1 + 1;
17
18pub const ENCRYPTION_UNENCRYPTED: u8 = 0;
19pub const ENCRYPTION_CHACHA20_POLY1305: u8 = 1;
20pub const ENCRYPTION_AES_256_GCM: u8 = 2;
21
22fn is_valid_encryption_type(value: u8) -> bool {
23 matches!(
24 value,
25 ENCRYPTION_UNENCRYPTED
26 | ENCRYPTION_CHACHA20_POLY1305
27 | ENCRYPTION_AES_256_GCM
28 )
29}
30
31fn is_encrypted(encryption_type: u8) -> bool {
32 encryption_type != ENCRYPTION_UNENCRYPTED
33}
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36#[repr(u8)]
37pub enum EncryptionType {
38 Unencrypted = ENCRYPTION_UNENCRYPTED,
39 ChaCha20Poly1305 = ENCRYPTION_CHACHA20_POLY1305,
40 Aes256Gcm = ENCRYPTION_AES_256_GCM,
41}
42
43impl TryFrom<u8> for EncryptionType {
44 type Error = MoqSecureError;
45
46 fn try_from(value: u8) -> Result<Self, Self::Error> {
47 match value {
48 ENCRYPTION_UNENCRYPTED => Ok(Self::Unencrypted),
49 ENCRYPTION_CHACHA20_POLY1305 => {
50 Ok(Self::ChaCha20Poly1305)
51 }
52 ENCRYPTION_AES_256_GCM => Ok(Self::Aes256Gcm),
53 value => Err(MoqSecureError::UnsupportedAlgorithm(value)),
54 }
55 }
56}
57
58impl From<EncryptionType> for u8 {
59 fn from(value: EncryptionType) -> Self {
60 value as u8
61 }
62}
63
64#[derive(Debug, Clone, PartialEq, Eq)]
65pub struct WireHeader {
66 pub magic: [u8; 4],
67 pub version: u8,
68 pub key_id: u8,
69 pub ctr: u64,
70 pub n_signed: u8,
71 pub sig_flag: u8,
72 pub encryption_type: u8,
73}
74
75impl WireHeader {
76 pub fn encode(&self) -> Vec<u8> {
77 let mut result = Vec::with_capacity(FIXED_HEADER_LEN);
78
79 result.extend_from_slice(&self.magic);
80 result.push(self.version);
81 result.push(self.key_id);
82 result.extend_from_slice(&self.ctr.to_be_bytes());
83 result.push(self.n_signed);
84 result.push(self.sig_flag);
85 result.push(self.encryption_type);
86
87 result
88 }
89
90 pub fn aad(&self) -> Vec<u8> {
91 self.encode()
92 }
93
94 pub fn validate(&self) -> Result<(), MoqSecureError> {
95 if self.magic != MAGIC {
96 return Err(MoqSecureError::InvalidMagic);
97 }
98
99 if self.version != VERSION {
100 return Err(MoqSecureError::UnsupportedVersion(self.version));
101 }
102
103 if self.sig_flag != 0 && self.sig_flag != 1 {
104 return Err(MoqSecureError::InvalidSigFlag(self.sig_flag));
105 }
106
107 if !is_valid_encryption_type(self.encryption_type) {
108 return Err(MoqSecureError::UnsupportedAlgorithm(
109 self.encryption_type,
110 ));
111 }
112
113 if self.n_signed == 0 && self.sig_flag != 0 {
114 return Err(MoqSecureError::SigningMismatch);
115 }
116
117 Ok(())
118 }
119}
120
121#[derive(Debug, Clone)]
122pub struct Frame {
123 pub header: WireHeader,
129
130 pub payload: Vec<u8>,
131
132 pub tag: [u8; AEAD_TAG_LEN],
133
134 pub signature: Option<[u8; SIG_SLOT_LEN]>,
135}
136
137impl Frame {
138 pub fn parse(frame: &[u8]) -> Result<Self, MoqSecureError> {
139 if frame.len() < FIXED_HEADER_LEN {
140 return Err(MoqSecureError::TruncatedFrame);
141 }
142
143 let (header_bytes, body_and_signature) =
144 frame.split_at(FIXED_HEADER_LEN);
145
146 let mut offset = 0;
147
148 let mut magic = [0u8; 4];
149 magic.copy_from_slice(&header_bytes[offset..offset + 4]);
150 offset += 4;
151
152 let version = header_bytes[offset];
153 offset += 1;
154
155 let key_id = header_bytes[offset];
156 offset += 1;
157
158 let mut ctr_bytes = [0u8; 8];
159 ctr_bytes.copy_from_slice(&header_bytes[offset..offset + 8]);
160 offset += 8;
161
162 let ctr = u64::from_be_bytes(ctr_bytes);
163
164 let n_signed = header_bytes[offset];
165 offset += 1;
166
167 let sig_flag = header_bytes[offset];
168 offset += 1;
169
170 let encryption_type = header_bytes[offset];
171
172 let header = WireHeader {
173 magic,
174 version,
175 key_id,
176 ctr,
177 n_signed,
178 sig_flag,
179 encryption_type,
180 };
181
182 header.validate()?;
183
184 let signature_len = if header.sig_flag == 1 {
185 SIG_SLOT_LEN
186 } else {
187 0
188 };
189
190 if body_and_signature.len() < signature_len {
191 return Err(MoqSecureError::TruncatedFrame);
192 }
193
194 let (body, signature_bytes) = if signature_len == 0 {
195 (body_and_signature, None)
196 } else {
197 let body_len = body_and_signature.len() - SIG_SLOT_LEN;
198 let (body, signature) =
199 body_and_signature.split_at(body_len);
200
201 (body, Some(signature))
202 };
203
204 let signature = match signature_bytes {
205 None => None,
206
207 Some(bytes) => {
208 if bytes.len() != SIG_SLOT_LEN {
209 return Err(MoqSecureError::TruncatedFrame);
210 }
211
212 let mut signature = [0u8; SIG_SLOT_LEN];
213 signature.copy_from_slice(bytes);
214
215 if signature == [0u8; SIG_SLOT_LEN] {
216 return Err(MoqSecureError::InvalidSignature);
217 }
218
219 Some(signature)
220 }
221 };
222
223 if is_encrypted(header.encryption_type) {
224 if body.len() < AEAD_TAG_LEN {
225 return Err(MoqSecureError::CiphertextTooShort);
226 }
227
228 let ciphertext_len = body.len() - AEAD_TAG_LEN;
229 let ciphertext = &body[..ciphertext_len];
230 let tag_bytes = &body[ciphertext_len..];
231
232 let tag: [u8; AEAD_TAG_LEN] = tag_bytes
233 .try_into()
234 .map_err(|_| MoqSecureError::CiphertextTooShort)?;
235
236 Ok(Self {
237 header,
238 payload: ciphertext.to_vec(),
239 tag,
240 signature,
241 })
242 } else {
243 Ok(Self {
244 header,
245 payload: body.to_vec(),
246 tag: [0u8; AEAD_TAG_LEN],
247 signature,
248 })
249 }
250 }
251
252 pub fn serialize(&self) -> Vec<u8> {
253 let encrypted = is_encrypted(self.header.encryption_type);
254
255 let signature_len = if self.header.sig_flag == 1 {
256 SIG_SLOT_LEN
257 } else {
258 0
259 };
260
261 let body_len = self.payload.len()
262 + if encrypted { AEAD_TAG_LEN } else { 0 };
263
264 let mut result = Vec::with_capacity(
265 FIXED_HEADER_LEN + body_len + signature_len,
266 );
267
268 result.extend_from_slice(&self.header.encode());
269 result.extend_from_slice(&self.payload);
270
271 if encrypted {
272 result.extend_from_slice(&self.tag);
273 }
274
275 if self.header.sig_flag == 1 {
276 if let Some(signature) = self.signature {
277 result.extend_from_slice(&signature);
278 } else {
279 result.extend_from_slice(&[0u8; SIG_SLOT_LEN]);
280 }
281 }
282
283 result
284 }
285
286 pub fn aad_bytes(&self) -> Vec<u8> {
287 self.header.aad()
288 }
289
290 pub fn digest_for_signature(&self) -> [u8; 32] {
291 let header_bytes = self.header.encode();
292 let encrypted = is_encrypted(self.header.encryption_type);
293
294 let mut data = Vec::with_capacity(
295 header_bytes.len()
296 + self.payload.len()
297 + if encrypted { AEAD_TAG_LEN } else { 0 },
298 );
299
300 data.extend_from_slice(&header_bytes);
301 data.extend_from_slice(&self.payload);
302
303 if encrypted {
304 data.extend_from_slice(&self.tag);
305 }
306
307 sha256_digest(&data)
308 }
309
310 pub fn decode_plaintext_with_key_store(
311 &self,
312 key_store: &dyn KeyStore,
313 broadcaster_public_key: &ed25519_dalek::VerifyingKey,
314 lease_remaining: &mut u8,
315 ) -> Result<Vec<u8>, MoqSecureError> {
316 let signing_enabled = self.header.n_signed > 0;
317 let signed = self.header.sig_flag == 1;
318
319 let next_lease_remaining = if !signing_enabled {
325 if self.header.sig_flag != 0 || self.signature.is_some() {
326 return Err(MoqSecureError::SigningMismatch);
327 }
328
329 None
330 } else if signed {
331 let signature_bytes = self
332 .signature
333 .ok_or(MoqSecureError::InvalidSignature)?;
334
335 let signature =
336 ed25519_dalek::Signature::from_bytes(&signature_bytes);
337
338 let digest = self.digest_for_signature();
339
340 broadcaster_public_key
341 .verify(&digest, &signature)
342 .map_err(|_| MoqSecureError::InvalidSignature)?;
343
344 Some(self.header.n_signed.saturating_sub(1))
345 } else {
346 if *lease_remaining == 0 {
347 return Err(MoqSecureError::InvalidSignature);
348 }
349
350 Some(lease_remaining.saturating_sub(1))
351 };
352
353 let padded_plaintext = if is_encrypted(self.header.encryption_type) {
354 let key = key_store
355 .aead_key(self.header.key_id)
356 .ok_or(MoqSecureError::InvalidKeyId(self.header.key_id))?;
357
358 aead_decrypt(
359 self.header.encryption_type,
360 key,
361 self.header.key_id,
362 self.header.ctr,
363 &self.aad_bytes(),
364 &self.payload,
365 &self.tag,
366 )?
367 } else {
368 self.payload.clone()
369 };
370
371 if padded_plaintext.len() < PAD_LEN_FIELD_LEN {
372 return Err(MoqSecureError::InvalidPadLength);
373 }
374
375 let mut pad_len_bytes = [0u8; PAD_LEN_FIELD_LEN];
376 pad_len_bytes.copy_from_slice(
377 &padded_plaintext[..PAD_LEN_FIELD_LEN],
378 );
379
380 let pad_len = u32::from_be_bytes(pad_len_bytes) as usize;
381
382 let content_start = PAD_LEN_FIELD_LEN
383 .checked_add(pad_len)
384 .ok_or(MoqSecureError::InvalidPadLength)?;
385
386 if content_start > padded_plaintext.len() {
387 return Err(MoqSecureError::InvalidPadLength);
388 }
389
390 let plaintext = padded_plaintext[content_start..].to_vec();
391
392 if let Some(next) = next_lease_remaining {
393 *lease_remaining = next;
394 }
395
396 Ok(plaintext)
397 }
398}
399
400pub fn encrypt_frame(
401 key_store: &dyn KeyStore,
402 broadcaster_private_key: &ed25519_dalek::SigningKey,
403 key_id: u8,
404 ctr: u64,
405 n_signed: u8,
406 maybe_sign: bool,
407 encryption_type: u8,
408 pad_len: u32,
409 plaintext: &[u8],
410) -> Result<Frame, MoqSecureError> {
411 if !is_valid_encryption_type(encryption_type) {
412 return Err(MoqSecureError::UnsupportedAlgorithm(
413 encryption_type,
414 ));
415 }
416
417 let sig_flag = if n_signed != 0 && maybe_sign {
418 1
419 } else {
420 0
421 };
422
423 let header = WireHeader {
424 magic: MAGIC,
425 version: VERSION,
426 key_id,
427 ctr,
428 n_signed,
429 sig_flag,
430 encryption_type,
431 };
432
433 header.validate()?;
434
435 let pad_len_usize = pad_len as usize;
436
437 let mut padded_plaintext = Vec::with_capacity(
438 PAD_LEN_FIELD_LEN + pad_len_usize + plaintext.len(),
439 );
440
441 padded_plaintext.extend_from_slice(&pad_len.to_be_bytes());
442 padded_plaintext.resize(
443 PAD_LEN_FIELD_LEN + pad_len_usize,
444 0,
445 );
446 padded_plaintext.extend_from_slice(plaintext);
447
448 let frame_without_signature =
449 if is_encrypted(encryption_type) {
450 let key = key_store
451 .aead_key(key_id)
452 .ok_or(MoqSecureError::InvalidKeyId(key_id))?;
453
454 let (ciphertext, tag) = aead_encrypt(
455 encryption_type,
456 key,
457 key_id,
458 ctr,
459 &header.aad(),
460 &padded_plaintext,
461 )?;
462
463 Frame {
464 header,
465 payload: ciphertext,
466 tag,
467 signature: None,
468 }
469 } else {
470 Frame {
471 header,
472 payload: padded_plaintext,
473 tag: [0u8; AEAD_TAG_LEN],
474 signature: None,
475 }
476 };
477
478 if sig_flag == 1 {
479 let digest = frame_without_signature.digest_for_signature();
480 let signature = broadcaster_private_key.sign(&digest);
481
482 Ok(Frame {
483 signature: Some(signature.to_bytes()),
484 ..frame_without_signature
485 })
486 } else {
487 Ok(frame_without_signature)
488 }
489}
490
491pub fn decrypt_frame(
492 key_store: &dyn KeyStore,
493 broadcaster_public_key: &ed25519_dalek::VerifyingKey,
494 lease_remaining: &mut u8,
495 frame_bytes: &[u8],
496) -> Result<Vec<u8>, MoqSecureError> {
497 let frame = Frame::parse(frame_bytes)?;
498
499 frame.decode_plaintext_with_key_store(
500 key_store,
501 broadcaster_public_key,
502 lease_remaining,
503 )
504}