Skip to main content

moq_secure/
wire.rs

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
14// magic(4) | version(1) | key_id(1) | ctr(8) |
15// n_signed(1) | sig_flag(1) | encryption_type(1)
16pub 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    /// For encrypted frames, this is ciphertext without the AEAD tag.
124    ///
125    /// For unencrypted frames, this is:
126    ///
127    /// pad_len(4-byte big endian) || padding || plaintext
128    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        /*
320         * Validate the signing state and verify signatures first.
321         * The lease is updated only after decryption and padding
322         * validation succeed.
323         */
324        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}