Skip to content

Commit 3eac7ed

Browse files
committed
feat(xmtp_mls): KeyPackageCleaner sleeps to exact rotation deadline (kill 5s poll); span-gated maintain
1 parent 331b99a commit 3eac7ed

1 file changed

Lines changed: 168 additions & 56 deletions

File tree

crates/xmtp_mls/src/worker/key_package_cleaner.rs

Lines changed: 168 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -5,20 +5,31 @@ use crate::worker::BoxedWorker;
55
use crate::worker::NeedsDbReconnect;
66
use crate::worker::WorkerResult;
77
use crate::worker::{Worker, WorkerFactory, WorkerKind};
8-
use futures::StreamExt;
98
use futures::TryFutureExt;
109
use openmls_traits::storage::StorageProvider;
11-
use std::time::Duration;
1210
use thiserror::Error;
11+
use tracing::Instrument;
1312
use xmtp_configuration::CREATE_PQ_KEY_PACKAGE_EXTENSION;
1413
use xmtp_db::prelude::*;
1514
use 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)]
2435
pub 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

Comments
 (0)