1use std::collections::{HashMap, hash_map::Entry};
21use std::hash::Hash;
22
23use crate::{
24 errors::CryptoError, exchange::prekey::SignedPrekey, signatures::keypair::PublicKeyBytes,
25};
26
27pub struct PrekeyBundle {
29 pub identity: PublicKeyBytes,
32 pub prekey: SignedPrekey,
34}
35
36struct Account {
37 identity: PublicKeyBytes,
38 last_resort: SignedPrekey,
39 one_time: Vec<SignedPrekey>,
40}
41
42pub struct PrekeyServer<K> {
44 accounts: HashMap<K, Account>,
45}
46
47impl<K: Eq + Hash> PrekeyServer<K> {
48 pub fn new() -> Self {
50 Self {
51 accounts: HashMap::new(),
52 }
53 }
54
55 pub fn publish(
72 &mut self,
73 user: K,
74 identity: PublicKeyBytes,
75 last_resort: SignedPrekey,
76 one_time: Vec<SignedPrekey>,
77 ) -> Result<(), CryptoError> {
78 if !last_resort.verify(&identity) || !one_time.iter().all(|p| p.verify(&identity)) {
79 return Err(CryptoError::InvalidSignature);
80 }
81 match self.accounts.entry(user) {
82 Entry::Occupied(mut entry) => {
83 let account = entry.get_mut();
84 if account.identity != identity {
85 return Err(CryptoError::InvalidKey);
86 }
87 account.last_resort = last_resort;
88 account.one_time.extend(one_time);
89 }
90 Entry::Vacant(entry) => {
91 entry.insert(Account {
92 identity,
93 last_resort,
94 one_time,
95 });
96 }
97 }
98 Ok(())
99 }
100
101 pub fn fetch(&mut self, user: &K) -> Option<PrekeyBundle> {
108 let account = self.accounts.get_mut(user)?;
109 let prekey = account
110 .one_time
111 .pop()
112 .unwrap_or_else(|| account.last_resort.clone());
113 Some(PrekeyBundle {
114 identity: account.identity,
115 prekey,
116 })
117 }
118
119 pub fn one_time_remaining(&self, user: &K) -> usize {
121 self.accounts.get(user).map_or(0, |a| a.one_time.len())
122 }
123}
124
125impl<K: Eq + Hash> Default for PrekeyServer<K> {
126 fn default() -> Self {
127 Self::new()
128 }
129}
130
131#[cfg(test)]
132mod tests {
133 use super::*;
134 use crate::{
135 exchange::pair::KEMPair,
136 signatures::keypair::{SignerPair, ViewOperations},
137 };
138
139 fn prekey(owner: &mut SignerPair) -> SignedPrekey {
140 SignedPrekey::new(owner, &KEMPair::create()).unwrap()
141 }
142
143 #[test]
144 fn test_one_time_prekeys_are_handed_out_once() {
145 let mut bob = SignerPair::create();
146 let last_resort = prekey(&mut bob);
147 let one_time = vec![prekey(&mut bob), prekey(&mut bob)];
148 let mut ids: Vec<_> = one_time.iter().map(|p| p.id()).collect();
149
150 let mut server = PrekeyServer::new();
151 server
152 .publish("bob", *bob.pub_key_bytes(), last_resort.clone(), one_time)
153 .unwrap();
154 assert_eq!(server.one_time_remaining(&"bob"), 2);
155
156 let mut handed_out = vec![
157 server.fetch(&"bob").unwrap().prekey.id(),
158 server.fetch(&"bob").unwrap().prekey.id(),
159 ];
160 handed_out.sort();
161 ids.sort();
162 assert_eq!(handed_out, ids);
163 assert_eq!(server.one_time_remaining(&"bob"), 0);
164
165 let bundle = server.fetch(&"bob").unwrap();
167 assert_eq!(bundle.prekey.id(), last_resort.id());
168 assert_eq!(bundle.identity, *bob.pub_key_bytes());
169 assert!(server.fetch(&"bob").is_some());
170
171 assert!(server.fetch(&"carol").is_none());
172 assert_eq!(server.one_time_remaining(&"carol"), 0);
173 }
174
175 #[test]
176 fn test_publish_rejects_foreign_prekeys_and_identity_change() {
177 let mut bob = SignerPair::create();
178 let mut mallory = SignerPair::create();
179 let mut server = PrekeyServer::new();
180
181 let result = server.publish("bob", *bob.pub_key_bytes(), prekey(&mut mallory), vec![]);
183 assert!(matches!(result, Err(CryptoError::InvalidSignature)));
184
185 let result = server.publish(
187 "bob",
188 *bob.pub_key_bytes(),
189 prekey(&mut bob),
190 vec![prekey(&mut mallory)],
191 );
192 assert!(matches!(result, Err(CryptoError::InvalidSignature)));
193 assert!(server.fetch(&"bob").is_none());
194
195 server
196 .publish("bob", *bob.pub_key_bytes(), prekey(&mut bob), vec![])
197 .unwrap();
198
199 let result = server.publish(
201 "bob",
202 *mallory.pub_key_bytes(),
203 prekey(&mut mallory),
204 vec![],
205 );
206 assert!(matches!(result, Err(CryptoError::InvalidKey)));
207
208 let new_last_resort = prekey(&mut bob);
210 server
211 .publish(
212 "bob",
213 *bob.pub_key_bytes(),
214 new_last_resort.clone(),
215 vec![prekey(&mut bob)],
216 )
217 .unwrap();
218 assert_eq!(server.one_time_remaining(&"bob"), 1);
219 server.fetch(&"bob").unwrap();
220 assert_eq!(
221 server.fetch(&"bob").unwrap().prekey.id(),
222 new_last_resort.id()
223 );
224 }
225
226 #[test]
227 fn test_default_is_empty() {
228 let mut server = PrekeyServer::<u32>::default();
229 assert!(server.fetch(&1).is_none());
230 }
231}