@@ -5,20 +5,31 @@ use crate::worker::BoxedWorker;
55use crate :: worker:: NeedsDbReconnect ;
66use crate :: worker:: WorkerResult ;
77use crate :: worker:: { Worker , WorkerFactory , WorkerKind } ;
8- use futures:: StreamExt ;
98use futures:: TryFutureExt ;
109use openmls_traits:: storage:: StorageProvider ;
11- use std:: time:: Duration ;
1210use thiserror:: Error ;
11+ use tracing:: Instrument ;
1312use xmtp_configuration:: CREATE_PQ_KEY_PACKAGE_EXTENSION ;
1413use xmtp_db:: prelude:: * ;
1514use xmtp_db:: {
1615 MlsProviderExt , StorageError ,
16+ encrypted_store:: key_package_history:: StoredKeyPackageHistoryEntry ,
1717 sql_key_store:: { KEY_PACKAGE_REFERENCES , KEY_PACKAGE_WRAPPER_PRIVATE_KEY } ,
1818} ;
1919
20- /// Interval at which the KeyPackagesCleanerWorker runs to delete expired messages.
21- pub const INTERVAL_DURATION : Duration = Duration :: from_secs ( 5 ) ;
20+ #[ derive( Debug , PartialEq , Eq ) ]
21+ enum WakePlan {
22+ RunNow ,
23+ SleepUntil ( i64 ) , // absolute deadline ns
24+ }
25+
26+ fn plan ( next_rotation : Option < i64 > , now : i64 ) -> WakePlan {
27+ match next_rotation {
28+ None => WakePlan :: RunNow , // NULL = due now
29+ Some ( at) if at <= now => WakePlan :: RunNow ,
30+ Some ( at) => WakePlan :: SleepUntil ( at) ,
31+ }
32+ }
2233
2334#[ derive( Clone ) ]
2435pub struct Factory < Context > {
@@ -116,21 +127,77 @@ where
116127 Context : XmtpSharedContext + ' static ,
117128{
118129 async fn run ( & mut self ) -> Result < ( ) , KeyPackagesCleanerError > {
119- let ( base, jitter) = self
120- . context
121- . worker_interval ( WorkerKind :: KeyPackageCleaner , INTERVAL_DURATION ) ;
122- let mut intervals = xmtp_common:: time:: jittered_interval_stream ( base, jitter) ;
123- while ( intervals. next ( ) . await ) . is_some ( ) {
124- self . tick ( ) . await ?;
130+ let receiver = self . context . key_package_channels ( ) . receiver . clone ( ) ;
131+ let mut receiver = receiver. lock ( ) . await ;
132+ loop {
133+ // Drain any pending re-arm signals before computing the plan so we
134+ // don't lose a wakeup that arrived while we were working.
135+ while receiver. try_recv ( ) . is_ok ( ) { }
136+
137+ match plan (
138+ self . context
139+ . db ( )
140+ . next_key_package_rotation_ns ( )
141+ . map_err ( KeyPackagesCleanerError :: Metadata ) ?,
142+ xmtp_common:: time:: now_ns ( ) ,
143+ ) {
144+ WakePlan :: RunNow => {
145+ self . maintain ( ) . await ?;
146+ }
147+ WakePlan :: SleepUntil ( deadline) => {
148+ let dur = std:: time:: Duration :: from_nanos (
149+ ( deadline - xmtp_common:: time:: now_ns ( ) ) . max ( 0 ) as u64 ,
150+ ) ;
151+ tokio:: select! {
152+ // Re-arm signal: recompute the deadline, do NOT run work.
153+ // `None` means every sender was dropped (context torn down) —
154+ // stop rather than busy-spin on a closed channel.
155+ msg = receiver. recv( ) => { if msg. is_none( ) { return Ok ( ( ) ) ; } }
156+ // Deadline elapsed: time to do maintenance.
157+ ( ) = xmtp_common:: time:: sleep( dur) => {
158+ self . maintain( ) . await ?;
159+ }
160+ }
161+ }
162+ }
125163 }
126- Ok ( ( ) )
127164 }
128165
129- #[ tracing:: instrument( skip_all, fields( worker = ?self . kind( ) , operation = "worker_turn" ) ) ]
130- async fn tick ( & mut self ) -> Result < ( ) , KeyPackagesCleanerError > {
131- self . delete_expired_key_packages ( ) ?;
132- self . rotate_last_key_package_if_needed ( ) . await ?;
133- Ok ( ( ) )
166+ /// One maintenance pass. Guards read OUTSIDE the span; if nothing is due,
167+ /// returns immediately with no span. Otherwise a single `worker_turn` span
168+ /// wraps the work so tracing records the full operation (including errors).
169+ async fn maintain ( & mut self ) -> Result < ( ) , KeyPackagesCleanerError > {
170+ let expired = self
171+ . context
172+ . db ( )
173+ . get_expired_key_packages ( )
174+ . map_err ( KeyPackagesCleanerError :: Fetch ) ?;
175+ let rotate_due = self
176+ . context
177+ . db ( )
178+ . is_identity_needs_rotation ( )
179+ . map_err ( KeyPackagesCleanerError :: Metadata ) ?;
180+
181+ if expired. is_empty ( ) && !rotate_due {
182+ return Ok ( ( ) ) ;
183+ }
184+
185+ let span = tracing:: info_span!(
186+ "worker_turn" ,
187+ worker = ?self . kind( ) ,
188+ operation = "worker_turn"
189+ ) ;
190+ async {
191+ if !expired. is_empty ( ) {
192+ self . delete_key_packages ( expired) ?;
193+ }
194+ if rotate_due {
195+ self . rotate_last_key_package_if_needed ( ) . await ?;
196+ }
197+ Ok :: < ( ) , KeyPackagesCleanerError > ( ( ) )
198+ }
199+ . instrument ( span)
200+ . await
134201 }
135202
136203 /// Delete a key package from the local database.
@@ -156,65 +223,110 @@ where
156223 Ok ( ( ) )
157224 }
158225
159- /// Delete all the expired keys
160- fn delete_expired_key_packages ( & mut self ) -> Result < ( ) , KeyPackagesCleanerError > {
226+ /// Delete an already-fetched list of expired key packages from local state
227+ /// and the database history table.
228+ fn delete_key_packages (
229+ & mut self ,
230+ expired : Vec < StoredKeyPackageHistoryEntry > ,
231+ ) -> Result < ( ) , KeyPackagesCleanerError > {
161232 let conn = self . context . db ( ) ;
162-
163- // Propagate (don't swallow): a swallowed fetch error never triggered the supervisor's
164- // reconnect path, so a pool outage retried silently every 5s.
165- let expired_kps = conn
166- . get_expired_key_packages ( )
167- . map_err ( KeyPackagesCleanerError :: Fetch ) ?;
168- if expired_kps. is_empty ( ) {
169- return Ok ( ( ) ) ;
170- }
171-
172- tracing:: info!( "Deleting {} expired key packages" , expired_kps. len( ) ) ;
173- // Delete from local db
174- for kp in & expired_kps {
233+ for kp in & expired {
175234 self . delete_key_package (
176235 kp. key_package_hash_ref . clone ( ) ,
177236 kp. post_quantum_public_key . clone ( ) ,
178237 )
179238 . map_err ( KeyPackagesCleanerError :: DeleteKeyPackage ) ?;
180239 }
181-
182- // Delete from database using the max expired ID
183- if let Some ( max_id) = expired_kps. iter ( ) . map ( |kp| kp. id ) . max ( ) {
240+ if let Some ( max_id) = expired. iter ( ) . map ( |kp| kp. id ) . max ( ) {
184241 conn. delete_key_package_history_up_to_id ( max_id)
185242 . map_err ( KeyPackagesCleanerError :: Deletion ) ?;
186243 tracing:: info!(
187- "Deleted {} expired key packages (up to ID {}) from local DB and state " ,
188- expired_kps . len( ) ,
244+ "Deleted {} expired key packages (up to ID {})" ,
245+ expired . len( ) ,
189246 max_id
190247 ) ;
191248 }
192-
193249 Ok ( ( ) )
194250 }
195251
196- /// Check if we need to rotate the keys and upload new keypackage if the las one rotate in has passed
252+ /// Upload a fresh key package if the current one has passed its rotation deadline.
197253 async fn rotate_last_key_package_if_needed ( & mut self ) -> Result < ( ) , KeyPackagesCleanerError > {
198- let conn = self . context . db ( ) ;
254+ tracing:: info!( "Rotating key package" ) ;
255+ self . context
256+ . identity ( )
257+ . rotate_and_upload_key_package (
258+ self . context . api ( ) ,
259+ self . context . mls_storage ( ) ,
260+ CREATE_PQ_KEY_PACKAGE_EXTENSION ,
261+ )
262+ . await
263+ . map_err ( KeyPackagesCleanerError :: Rotation ) ?;
264+ tracing:: info!( "Key package rotation successful" ) ;
265+ Ok ( ( ) )
266+ }
267+ }
199268
200- if conn
201- . is_identity_needs_rotation ( )
202- . map_err ( KeyPackagesCleanerError :: Metadata ) ?
203- {
204- tracing:: info!( "Rotating key package" ) ;
205- self . context
206- . identity ( )
207- . rotate_and_upload_key_package (
208- self . context . api ( ) ,
209- self . context . mls_storage ( ) ,
210- CREATE_PQ_KEY_PACKAGE_EXTENSION ,
211- )
212- . await
213- . map_err ( KeyPackagesCleanerError :: Rotation ) ?;
214- tracing:: info!( "Key package rotation successful" ) ;
215- return Ok ( ( ) ) ;
269+ #[ cfg( test) ]
270+ mod tests {
271+ #[ test]
272+ fn plan_cases ( ) {
273+ assert_eq ! ( super :: plan( None , 100 ) , super :: WakePlan :: RunNow ) ;
274+ assert_eq ! ( super :: plan( Some ( 99 ) , 100 ) , super :: WakePlan :: RunNow ) ;
275+ assert_eq ! ( super :: plan( Some ( 100 ) , 100 ) , super :: WakePlan :: RunNow ) ;
276+ assert_eq ! (
277+ super :: plan( Some ( 101 ) , 100 ) ,
278+ super :: WakePlan :: SleepUntil ( 101 )
279+ ) ;
280+ }
281+
282+ /// Requires a live XMTP node. Verifies that `queue_key_rotation` marks the
283+ /// identity for rotation and wakes the KeyPackageCleaner worker, which then
284+ /// uploads a new key package within ~12 s.
285+ #[ xmtp_common:: test( unwrap_try = true ) ]
286+ async fn queue_key_rotation_wakes_worker_and_rotates ( ) {
287+ use crate :: tester;
288+
289+ tester ! ( client) ;
290+
291+ let installation_id = client. installation_public_key ( ) . to_vec ( ) ;
292+
293+ // Snapshot the current on-network init key.
294+ // `VerifiedKeyPackageV2::hpke_init_key()` returns an owned Vec<u8>.
295+ let before = {
296+ let mut map = client
297+ . get_key_packages_for_installation_ids ( vec ! [ installation_id. clone( ) ] )
298+ . await ?;
299+ let entry = map
300+ . remove ( & installation_id)
301+ . expect ( "installation not found in response" ) ?;
302+ entry. hpke_init_key ( )
303+ } ;
304+
305+ // Queue a rotation and wake the worker.
306+ client. queue_key_rotation ( ) . await ?;
307+
308+ // Poll up to 12 s (60 × 200 ms) for the on-network key to change.
309+ // Uses `xmtp_common::time::sleep` so the test compiles for wasm too
310+ // (std::time::Instant panics on wasm).
311+ let mut after = None ;
312+ for _ in 0 ..60u32 {
313+ xmtp_common:: time:: sleep ( std:: time:: Duration :: from_millis ( 200 ) ) . await ;
314+ let mut map = client
315+ . get_key_packages_for_installation_ids ( vec ! [ installation_id. clone( ) ] )
316+ . await ?;
317+ let entry = map
318+ . remove ( & installation_id)
319+ . expect ( "installation not found in response" ) ?;
320+ let key = entry. hpke_init_key ( ) ;
321+ if key != before {
322+ after = Some ( key) ;
323+ break ;
324+ }
216325 }
217326
218- Ok ( ( ) )
327+ assert ! (
328+ after. is_some( ) ,
329+ "key package did not rotate after queue_key_rotation + worker wake"
330+ ) ;
219331 }
220332}
0 commit comments