1use std::{collections::BTreeMap, sync::Arc};
2
3use tuwunel_core::{
4 Result, debug, info,
5 smallvec::SmallVec,
6 utils::{TryReadyExt, hash::sha256::Digest},
7 warn,
8};
9use tuwunel_database::{Map, Txn};
10
11use super::{
12 Reason, Verdict, clear_chains,
13 scan::{Family, Scan, short_of},
14};
15use crate::{
16 Services,
17 rooms::state_compressor::{
18 CompressedState, CompressedStateEvent, StateDiff, compress_state_event,
19 parse_compressed_state_event,
20 },
21};
22
23type Digests = BTreeMap<u64, SmallVec<[Digest; 1]>>;
28
29#[derive(Default)]
35struct Healed {
36 reinstated: usize,
37 promoted: usize,
38}
39
40struct Patched {
47 added: CompressedState,
48 removed: CompressedState,
49 added_changes: u64,
50 removed_changes: u64,
51 shrunk: usize,
52 colliding: usize,
53}
54
55#[tracing::instrument(level = "debug", skip_all)]
68pub(super) fn heal(services: &Services, scan: &Scan) -> bool {
69 if scan.unverifiable || !scan.healable() {
70 return false;
71 }
72
73 let db = &services.db;
74 let mut txn = db.txn();
75
76 let event_reverse = &db["shorteventid_eventid"];
77 let event_forward = &db["eventid_shorteventid"];
78 let events = heal_family(&mut txn, event_reverse, event_forward, &scan.events);
79
80 let statekey_reverse = &db["shortstatekey_statekey"];
81 let statekey_forward = &db["statekey_shortstatekey"];
82 let statekeys = heal_family(&mut txn, statekey_reverse, statekey_forward, &scan.statekeys);
83
84 info!(
85 reinstated_events = events.reinstated,
86 promoted_events = events.promoted,
87 reinstated_statekeys = statekeys.reinstated,
88 promoted_statekeys = statekeys.promoted,
89 "Completing torn short id writes; rescanning to re-measure what they explain."
90 );
91
92 txn.execute();
93
94 true
95}
96
97fn heal_family(txn: &mut Txn, reverse: &Map, forward: &Map, family: &Family) -> Healed {
102 if !family.healable() {
103 return Healed::default();
104 }
105
106 for (short, identity) in &family.dangling {
107 txn.insert_raw(reverse, short.to_be_bytes(), identity.as_slice());
108 }
109
110 for (short, identity) in &family.promotable {
111 txn.insert_raw(forward, identity.as_slice(), short.to_be_bytes());
112 }
113
114 Healed {
115 reinstated: family.dangling.len(),
116 promoted: family.promotable.len(),
117 }
118}
119
120#[tracing::instrument(level = "debug", skip_all)]
128pub(super) async fn repair(
129 services: &Services,
130 scan: &Scan,
131 chains_cleared_this_boot: bool,
132) -> Result<Verdict> {
133 let healable = scan.healable();
134 let family_anomalous = scan.family_anomalous();
135 let needs_chain_clear = scan.unverifiable || healable || family_anomalous;
136
137 if !chains_cleared_this_boot && needs_chain_clear {
138 warn!(
139 "Short id verdict skipped the deep auth chain census; clearing the cache before \
140 finalizing."
141 );
142
143 clear_chains(services).await?;
144 }
145
146 if scan.unverifiable {
147 return Ok(Verdict::Declined(Reason::Unverifiable));
148 }
149
150 if scan.strays > 0 {
151 warn!(
152 stray_references = scan.strays,
153 "Short room id references without a forward row exist; nothing repairs them."
154 );
155 }
156
157 if scan.dirty > 0 {
158 warn!(
159 dirty_entries = scan.dirty,
160 total_entries = scan.entries,
161 "Cached auth chains contain malformed or stale short id data; clearing the auth \
162 chain cache."
163 );
164
165 clear_chains(services).await?;
166 }
167
168 if healable {
171 warn!(
172 dangling_events = scan.events.dangling.len(),
173 dangling_statekeys = scan.statekeys.dangling.len(),
174 promotable_events = scan.events.promotable.len(),
175 promotable_statekeys = scan.statekeys.promotable.len(),
176 "Refusing the destructive short id repair; the heal passes did not settle. The \
177 residue is recorded and left in place. Please report this line upstream."
178 );
179
180 return Ok(Verdict::Declined(Reason::Healable));
181 }
182
183 if family_anomalous {
184 warn!(
185 dangling_events = scan.events.dangling.len(),
186 dangling_statekeys = scan.statekeys.dangling.len(),
187 promotable_events = scan.events.promotable.len(),
188 promotable_statekeys = scan.statekeys.promotable.len(),
189 contended_events = scan.events.contended,
190 contended_statekeys = scan.statekeys.contended,
191 unresolved_events = scan.events.unresolved,
192 unresolved_statekeys = scan.statekeys.unresolved,
193 malformed_event_keys = scan.events.malformed,
194 malformed_statekey_keys = scan.statekeys.malformed,
195 "Refusing the destructive short id repair; the family census contains an unhandled \
196 shape. The residue is recorded and left in place. Please report this line upstream."
197 );
198
199 return Ok(Verdict::Declined(Reason::FamilyAnomalous));
200 }
201
202 if scan.events.losers.is_empty() && scan.statekeys.losers.is_empty() {
203 info!("Short id mappings verified injective.");
204
205 return Ok(Verdict::Settled);
206 }
207
208 if scan.anomalous() {
209 warn!(
210 infected_parents = scan.infected_parents,
211 orphan_entries = scan.orphans,
212 malformed_diffs = scan.malformed_diffs,
213 colliding_diffs = scan.colliding_diffs,
214 "Refusing the destructive short id repair; the nonzero counts name shapes it does \
215 not handle. The residue is recorded and left in place. Please report this line \
216 upstream."
217 );
218
219 return Ok(Verdict::Declined(Reason::DeepAnomalous));
220 }
221
222 patch_statediffs(services, scan).await?;
223 move_keys(services, scan).await?;
224 delete_losers(services, scan);
225
226 Ok(Verdict::Settled)
227}
228
229#[tracing::instrument(level = "debug", skip_all)]
237async fn patch_statediffs(services: &Services, scan: &Scan) -> Result {
238 if scan.infected.is_empty() {
239 return Ok(());
240 }
241
242 let progress = &services.server.progress;
243
244 progress.begin("fix_short_injectivity: patch state diffs");
245 let digests: Digests = services.db["statehash_shortstatehash"]
246 .raw_stream()
247 .ready_try_fold(Digests::new(), |mut digests, (key, value)| {
248 progress.advance();
249
250 if let Some(state) = short_of(value).filter(|state| scan.infected.contains(state))
251 && let Ok(digest) = key.try_into()
252 {
253 digests.entry(state).or_default().push(digest);
254 }
255
256 Ok(digests)
257 })
258 .await?;
259
260 for &state in &scan.infected {
263 patch_state(services, scan, &digests, state).await?;
264 }
265
266 Ok(())
267}
268
269#[tracing::instrument(
274 level = "debug",
275 skip_all,
276 fields(
277 %state,
278 ),
279)]
280async fn patch_state(services: &Services, scan: &Scan, digests: &Digests, state: u64) -> Result {
281 let diff = services
282 .state_compressor
283 .get_statediff(state)
284 .await?;
285
286 let Patched {
287 added,
288 removed,
289 added_changes,
290 removed_changes,
291 shrunk,
292 colliding,
293 } = patch_runs(&diff, scan);
294
295 if removed_changes > 0 {
296 warn!(
297 %state,
298 entries = removed_changes,
299 "Patched ghost entries inside a removed run; the state resolves differently \
300 now that the removal matches."
301 );
302 }
303
304 if shrunk > 0 {
305 info!(
306 %state,
307 entries = shrunk,
308 "Patching converged duplicate entries; the state shrank."
309 );
310 }
311
312 if colliding > 0 {
313 warn!(
314 %state,
315 entries = colliding,
316 "Subtracted colliding entries from the removed run; each state key keeps its \
317 event rather than vanishing."
318 );
319 }
320
321 let patched = StateDiff {
322 parent: diff.parent,
323 added: Arc::new(added),
324 removed: Arc::new(removed),
325 };
326
327 let statehashes = &services.db["statehash_shortstatehash"];
328 let mut txn = services.db.txn();
329
330 services
333 .state_compressor
334 .save_statediff(&mut txn, state, &patched);
335
336 digests
337 .get(&state)
338 .into_iter()
339 .flatten()
340 .for_each(|digest| txn.del_raw(statehashes, digest));
341
342 txn.execute();
343
344 info!(
345 %state,
346 entries = added_changes.saturating_add(removed_changes),
347 "Patched stale short ids out of a compressed state."
348 );
349
350 Ok(())
351}
352
353fn patch_runs(diff: &StateDiff, scan: &Scan) -> Patched {
364 let (added, added_changes) = patch(&diff.added, scan);
365 let (patched_removed, removed_changes) = patch(&diff.removed, scan);
366
367 let shrunk = diff
368 .added
369 .len()
370 .saturating_add(diff.removed.len())
371 .saturating_sub(added.len())
372 .saturating_sub(patched_removed.len());
373
374 let removed: CompressedState = patched_removed
375 .difference(&added)
376 .copied()
377 .collect();
378
379 let colliding = patched_removed
380 .len()
381 .saturating_sub(removed.len());
382
383 Patched {
384 added,
385 removed,
386 added_changes,
387 removed_changes,
388 shrunk,
389 colliding,
390 }
391}
392
393fn patch(entries: &CompressedState, scan: &Scan) -> (CompressedState, u64) {
398 entries
399 .iter()
400 .fold((CompressedState::new(), 0_u64), |(mut patched, changes), entry| {
401 let winner = winner_of(entry, scan);
402 let changes = changes.saturating_add(u64::from(winner.is_some()));
403
404 patched.insert(winner.unwrap_or(*entry));
405
406 (patched, changes)
407 })
408}
409
410fn winner_of(entry: &CompressedStateEvent, scan: &Scan) -> Option<CompressedStateEvent> {
414 let (statekey, event) = parse_compressed_state_event(*entry);
415 let winner_statekey = scan.statekeys.winners.get(&statekey).copied();
416 let winner_event = scan.events.winners.get(&event).copied();
417
418 (winner_statekey.is_some() || winner_event.is_some()).then(|| {
419 compress_state_event(winner_statekey.unwrap_or(statekey), winner_event.unwrap_or(event))
420 })
421}
422
423#[tracing::instrument(level = "debug", skip_all)]
429async fn move_keys(services: &Services, scan: &Scan) -> Result {
430 if scan.moves.is_empty() && scan.relations.is_empty() {
431 return Ok(());
432 }
433
434 let progress = &services.server.progress;
435
436 progress.begin("fix_short_injectivity: move keys");
437
438 let states = &services.db["shorteventid_shortstatehash"];
441 let mut txn = services.db.txn();
442
443 for &loser in &scan.moves {
444 let Some(&winner) = scan.events.winners.get(&loser) else {
445 continue;
446 };
447
448 move_state_row(states, &mut txn, loser, winner).await?;
449 }
450
451 let relations = &services.db["relatesto_typed"];
452
453 for (key, loser) in &scan.relations {
454 let Some(&winner) = scan.events.winners.get(loser) else {
455 continue;
456 };
457
458 txn.insert_raw(relations, key, winner.to_be_bytes());
459 }
460
461 info!(
462 moves = scan.moves.len(),
463 relations = scan.relations.len(),
464 "Rewrote loser-keyed and loser-valued rows."
465 );
466
467 txn.execute();
468
469 Ok(())
470}
471
472async fn move_state_row(states: &Arc<Map>, txn: &mut Txn, loser: u64, winner: u64) -> Result {
477 match states.get(&winner.to_be_bytes()).await {
478 | Ok(_) =>
479 debug!(loser, winner, "Dropping a loser-keyed state row; the winner has its own."),
480 | Err(error) if error.is_not_found() => {
481 let value = states.get(&loser.to_be_bytes()).await?;
482
483 txn.insert_raw(states, winner.to_be_bytes(), &*value);
484 },
485 | Err(error) => return Err(error),
486 }
487
488 txn.del_raw(states, loser.to_be_bytes());
489
490 Ok(())
491}
492
493fn delete_losers(services: &Services, scan: &Scan) {
499 let progress = &services.server.progress;
500
501 progress.begin("fix_short_injectivity: delete losers");
502 info!(
503 stale_events = scan.events.losers.len(),
504 stale_statekeys = scan.statekeys.losers.len(),
505 "Deleting stale short id reverse rows."
506 );
507
508 let _cork = services.db.cork_and_sync();
509
510 let events = &services.db["shorteventid_eventid"];
511
512 for loser in &scan.events.losers {
513 events.remove(&loser.to_be_bytes());
514 }
515
516 let statekeys = &services.db["shortstatekey_statekey"];
517
518 for loser in &scan.statekeys.losers {
519 statekeys.remove(&loser.to_be_bytes());
520 }
521}
522
523#[cfg(test)]
524mod tests {
525 use std::{collections::BTreeMap, sync::Arc};
526
527 use super::{CompressedState, Family, Scan, StateDiff, compress_state_event, patch_runs};
528
529 fn ghost_scan(loser: u64, winner: u64) -> Scan {
530 Scan {
531 events: Family {
532 losers: vec![loser],
533 winners: BTreeMap::from([(loser, winner)]),
534 ..Default::default()
535 },
536 ..Default::default()
537 }
538 }
539
540 fn diff_of(added: &[(u64, u64)], removed: &[(u64, u64)]) -> StateDiff {
541 let compress = |entries: &[(u64, u64)]| -> CompressedState {
542 entries
543 .iter()
544 .map(|&(statekey, event)| compress_state_event(statekey, event))
545 .collect()
546 };
547
548 StateDiff {
549 parent: None,
550 added: Arc::new(compress(added)),
551 removed: Arc::new(compress(removed)),
552 }
553 }
554
555 #[test]
556 fn a_ghost_added_with_its_winner_removed_keeps_the_state_key() {
557 let scan = ghost_scan(7, 3);
558 let diff = diff_of(&[(5, 7)], &[(5, 3)]);
559
560 let patched = patch_runs(&diff, &scan);
561 let winner_entry = compress_state_event(5, 3);
562
563 assert_eq!(patched.colliding, 1);
564 assert_eq!(patched.shrunk, 0);
565 assert_eq!(patched.added_changes, 1);
566 assert_eq!(patched.removed_changes, 0);
567 assert!(patched.added.contains(&winner_entry));
568 assert!(patched.removed.is_empty());
569 }
570
571 #[test]
572 fn a_ghost_removed_with_its_winner_added_keeps_the_state_key() {
573 let scan = ghost_scan(7, 3);
574 let diff = diff_of(&[(5, 3)], &[(5, 7)]);
575
576 let patched = patch_runs(&diff, &scan);
577 let winner_entry = compress_state_event(5, 3);
578
579 assert_eq!(patched.colliding, 1);
580 assert_eq!(patched.shrunk, 0);
581 assert_eq!(patched.added_changes, 0);
582 assert_eq!(patched.removed_changes, 1);
583 assert!(patched.added.contains(&winner_entry));
584 assert!(patched.removed.is_empty());
585 }
586}