1use std::collections::HashMap;
2use std::sync::{Arc, Mutex};
3
4use sframe::frame::{
5 EncryptedFrameView, FrameValidation, MediaFrameView, MonotonicCounter,
6 ReplayAttackProtectionStore,
7};
8use sframe::header::KeyId;
9use sframe::key::{DecryptionKey, EncryptionKey};
10
11use crate::error::SFrameError;
12use crate::header::SFrameHeader;
13use crate::kdf::CipherSuite;
14
15const REPLAY_WINDOW: u64 = 1024;
17
18struct EncryptorState {
27 key: EncryptionKey,
28 counter: MonotonicCounter,
29}
30
31#[derive(Clone)]
39pub struct SFrameEncryptor {
40 state: Arc<Mutex<EncryptorState>>,
41 kid: KeyId,
42}
43
44impl SFrameEncryptor {
45 pub(crate) fn new(base_key: &[u8; 32], kid: u64, suite: CipherSuite) -> Self {
46 let key = EncryptionKey::derive_from(suite.to_sframe(), kid, base_key)
47 .expect("key derivation from a 32-byte base key never fails");
48 Self {
49 state: Arc::new(Mutex::new(EncryptorState {
50 key,
51 counter: MonotonicCounter::default(),
52 })),
53 kid,
54 }
55 }
56
57 pub fn encrypt(&mut self, plaintext: &[u8], extra_aad: &[u8]) -> Result<Vec<u8>, SFrameError> {
68 let mut state = self
69 .state
70 .lock()
71 .expect("sframe encryptor state mutex poisoned");
72 if state.counter.is_exhausted() {
74 return Err(SFrameError::CounterExhausted(self.kid));
75 }
76 let frame = MediaFrameView::with_meta_data(&mut state.counter, plaintext, extra_aad);
77 let encrypted = frame
78 .encrypt(&state.key)
79 .map_err(|_| SFrameError::Encrypt)?;
80
81 Ok(encrypted.as_ref()[extra_aad.len()..].to_vec())
84 }
85
86 #[cfg(test)]
88 fn with_counter_at(base_key: &[u8; 32], kid: u64, suite: CipherSuite, start: u64) -> Self {
89 let encryptor = Self::new(base_key, kid, suite);
90 encryptor.state.lock().unwrap().counter =
91 MonotonicCounter::with_start_value(start, u64::MAX);
92 encryptor
93 }
94
95 pub fn counter(&self) -> u64 {
97 self.state
98 .lock()
99 .expect("sframe encryptor state mutex poisoned")
100 .counter
101 .current()
102 }
103
104 pub fn kid(&self) -> u64 {
106 self.kid
107 }
108}
109
110pub struct SFrameDecryptor {
119 base_key: [u8; 32],
120 epoch: u64,
121 suite: CipherSuite,
122 keys: HashMap<KeyId, DecryptionKey>,
124 replay: ReplayAttackProtectionStore,
126}
127
128impl SFrameDecryptor {
129 pub(crate) fn new(base_key: [u8; 32], epoch: u64, suite: CipherSuite) -> Self {
130 Self {
131 base_key,
132 epoch,
133 suite,
134 keys: HashMap::new(),
135 replay: ReplayAttackProtectionStore::with_tolerance(REPLAY_WINDOW),
136 }
137 }
138
139 pub fn decrypt(
143 &mut self,
144 payload: &[u8],
145 extra_aad: &[u8],
146 ) -> Result<(Vec<u8>, u32), SFrameError> {
147 let view = EncryptedFrameView::try_with_meta_data(&payload, &extra_aad)
148 .map_err(|e| SFrameError::Header(e.to_string()))?;
149
150 let kid = view.header().key_id();
151 let ctr = view.header().counter();
152
153 if SFrameHeader::epoch_from_kid(kid) != SFrameHeader::epoch_lsb(self.epoch) {
154 return Err(SFrameError::UnknownKid(kid));
155 }
156 let leaf = SFrameHeader::leaf_from_kid(kid);
157
158 self.replay
160 .inspect(view.header())
161 .map_err(|_| SFrameError::Replay { kid, ctr })?;
162
163 let suite = self.suite;
164 let base_key = self.base_key;
165 let key = self.keys.entry(kid).or_insert_with(|| {
166 DecryptionKey::derive_from(suite.to_sframe(), kid, base_key)
167 .expect("key derivation from a 32-byte base key never fails")
168 });
169
170 let media = view.decrypt(key).map_err(|_| SFrameError::Decrypt)?;
171
172 self.replay
174 .validate(view.header())
175 .map_err(|_| SFrameError::Replay { kid, ctr })?;
176
177 Ok((media.payload().to_vec(), leaf))
178 }
179
180 pub fn reset(&mut self) {
182 self.keys.clear();
183 self.replay = ReplayAttackProtectionStore::with_tolerance(REPLAY_WINDOW);
184 }
185}
186
187#[cfg(test)]
188mod tests {
189 use super::*;
190
191 const KID: u64 = 0xabc;
192
193 #[test]
194 fn last_counter_value_still_encrypts() {
195 let mut enc =
196 SFrameEncryptor::with_counter_at(&[7u8; 32], KID, CipherSuite::Aes128Gcm, u64::MAX);
197
198 assert!(enc.encrypt(b"last frame", b"").is_ok());
199 }
200
201 #[test]
202 fn exhausted_counter_errors_instead_of_panicking() {
203 let mut enc =
205 SFrameEncryptor::with_counter_at(&[7u8; 32], KID, CipherSuite::Aes128Gcm, u64::MAX);
206 enc.encrypt(b"last frame", b"").unwrap();
207
208 assert!(matches!(
209 enc.encrypt(b"one too many", b""),
210 Err(SFrameError::CounterExhausted(KID))
211 ));
212
213 assert!(enc.encrypt(b"and another", b"").is_err());
215 }
216}