Skip to main content

confium_privacy/
privacy_and_dist_patterns.rs

1//! Privacy-preserving computation + distributed systems patterns.
2//!
3//! PSI, PIR, differential privacy, feature flags, API versioning,
4//! schema registry, 2PC, WAL streaming, snapshot isolation,
5//! homomorphic MAC.
6
7use hmac::{Hmac, KeyInit, Mac};
8use serde::{Deserialize, Serialize};
9use sha2::{Digest, Sha256};
10use std::collections::{HashMap, HashSet};
11use std::sync::Mutex;
12use std::time::Instant;
13
14type HmacSha256 = Hmac<Sha256>;
15
16// === Private Set Intersection ===
17
18/// Hash-based PSI: both parties hash their sets, compare hashes.
19pub fn psi_hash_based(set_a: &[Vec<u8>], set_b: &[Vec<u8>], salt: &[u8]) -> Vec<Vec<u8>> {
20    let hashes_b: HashSet<[u8; 32]> = set_b.iter().map(|e| hash_with_salt(e, salt)).collect();
21    set_a
22        .iter()
23        .filter(|e| hashes_b.contains(&hash_with_salt(e, salt)))
24        .cloned()
25        .collect()
26}
27
28fn hash_with_salt(data: &[u8], salt: &[u8]) -> [u8; 32] {
29    let mut h = Sha256::new();
30    h.update(salt);
31    h.update(data);
32    let result = h.finalize();
33    let mut out = [0u8; 32];
34    out.copy_from_slice(&result);
35    out
36}
37
38/// PSI result size only (cardinality), not the actual intersection.
39pub fn psi_cardinality(set_a: &[Vec<u8>], set_b: &[Vec<u8>], salt: &[u8]) -> usize {
40    psi_hash_based(set_a, set_b, salt).len()
41}
42
43// === Private Information Retrieval ===
44
45/// Simple PIR: client downloads entire database (trivial but private).
46pub fn pir_trivial(database: &[Vec<u8>], _index: usize) -> Vec<Vec<u8>> {
47    database.to_vec()
48}
49
50/// Batch PIR: retrieve multiple indices in one query.
51pub fn pir_batch(database: &[Vec<u8>], indices: &[usize]) -> Vec<Vec<u8>> {
52    indices
53        .iter()
54        .filter_map(|&i| database.get(i).cloned())
55        .collect()
56}
57
58/// XOR-based PIR (2 servers): each server gets a random subset;
59/// XOR of responses gives the desired element.
60pub struct PirQuery {
61    pub mask: Vec<bool>,
62}
63pub struct PirResponse {
64    pub data: Vec<u8>,
65}
66
67pub fn pir_create_query(index: usize, db_size: usize) -> (PirQuery, PirQuery) {
68    use rand_core::{OsRng, RngCore};
69    let mut mask1 = vec![false; db_size];
70    let mut mask2 = vec![false; db_size];
71    let mut rng = OsRng;
72    for i in 0..db_size {
73        mask1[i] = rng.next_u32() & 1 == 1;
74        mask2[i] = mask1[i];
75    }
76    mask2[index] = !mask2[index]; // flip the target index
77    (PirQuery { mask: mask1 }, PirQuery { mask: mask2 })
78}
79
80pub fn pir_server_respond(database: &[Vec<u8>], query: &PirQuery) -> PirResponse {
81    let mut result = vec![0u8; database.first().map(|e| e.len()).unwrap_or(0)];
82    for (i, &selected) in query.mask.iter().enumerate() {
83        if selected {
84            if let Some(element) = database.get(i) {
85                for (j, &b) in element.iter().enumerate() {
86                    if j < result.len() {
87                        result[j] ^= b;
88                    }
89                }
90            }
91        }
92    }
93    PirResponse { data: result }
94}
95
96pub fn pir_decode(r1: &PirResponse, r2: &PirResponse) -> Vec<u8> {
97    r1.data
98        .iter()
99        .zip(r2.data.iter())
100        .map(|(a, b)| a ^ b)
101        .collect()
102}
103
104// === Differential Privacy ===
105
106/// Laplace mechanism: add noise calibrated to sensitivity and epsilon.
107pub fn laplace_noise(sensitivity: f64, epsilon: f64) -> f64 {
108    use rand_core::{OsRng, RngCore};
109    let scale = sensitivity / epsilon;
110    let mut buf = [0u8; 8];
111    OsRng.fill_bytes(&mut buf);
112    let u = (u64::from_le_bytes(buf) as f64) / (u64::MAX as f64);
113    let uniform = u - 0.5; // in [-0.5, 0.5)
114    -scale * uniform.signum() * (1.0 - 2.0 * uniform.abs()).ln()
115}
116
117/// Add Laplace noise to a numeric query result.
118pub fn dp_query(value: f64, sensitivity: f64, epsilon: f64) -> f64 {
119    value + laplace_noise(sensitivity, epsilon)
120}
121
122/// Gaussian mechanism: add Gaussian noise.
123pub fn gaussian_noise(sensitivity: f64, epsilon: f64, delta: f64) -> f64 {
124    use rand_core::{OsRng, RngCore};
125    let sigma = sensitivity * (2.0 * (1.25 / delta).ln()).sqrt() / epsilon;
126    let mut buf = [0u8; 8];
127    OsRng.fill_bytes(&mut buf);
128    let u1 = (u64::from_le_bytes(buf) as f64) / (u64::MAX as f64) + 1e-10;
129    OsRng.fill_bytes(&mut buf);
130    let u2 = (u64::from_le_bytes(buf) as f64) / (u64::MAX as f64) + 1e-10;
131    let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
132    z * sigma
133}
134
135/// Counting query with DP: report noisy count.
136pub fn dp_count(true_count: usize, epsilon: f64) -> f64 {
137    dp_query(true_count as f64, 1.0, epsilon)
138}
139
140// === Feature Flags ===
141
142#[derive(Debug, Clone, Serialize, Deserialize)]
143pub struct FeatureFlag {
144    pub name: String,
145    pub enabled: bool,
146    pub rollout_percentage: u8,
147}
148
149#[derive(Default)]
150pub struct FeatureFlags {
151    flags: Mutex<HashMap<String, FeatureFlag>>,
152}
153
154impl FeatureFlags {
155    pub fn new() -> Self {
156        Self::default()
157    }
158
159    pub fn set(&self, name: &str, enabled: bool) {
160        self.flags.lock().unwrap().insert(
161            name.into(),
162            FeatureFlag {
163                name: name.into(),
164                enabled,
165                rollout_percentage: if enabled { 100 } else { 0 },
166            },
167        );
168    }
169
170    pub fn set_rollout(&self, name: &str, percentage: u8) {
171        self.flags.lock().unwrap().insert(
172            name.into(),
173            FeatureFlag {
174                name: name.into(),
175                enabled: percentage > 0,
176                rollout_percentage: percentage,
177            },
178        );
179    }
180
181    pub fn is_enabled(&self, name: &str) -> bool {
182        let flags = self.flags.lock().unwrap();
183        flags.get(name).map(|f| f.enabled).unwrap_or(false)
184    }
185
186    pub fn is_enabled_for(&self, name: &str, user_id: &str) -> bool {
187        let flags = self.flags.lock().unwrap();
188        let flag = match flags.get(name) {
189            Some(f) => f,
190            None => return false,
191        };
192        if !flag.enabled {
193            return false;
194        }
195        if flag.rollout_percentage >= 100 {
196            return true;
197        }
198        let hash = hash_with_salt(user_id.as_bytes(), name.as_bytes());
199        let bucket = (u32::from_be_bytes([hash[0], hash[1], hash[2], hash[3]]) % 100) as u8;
200        bucket < flag.rollout_percentage
201    }
202
203    pub fn count(&self) -> usize {
204        self.flags.lock().unwrap().len()
205    }
206    pub fn remove(&self, name: &str) {
207        self.flags.lock().unwrap().remove(name);
208    }
209}
210
211// === API Versioning ===
212
213#[derive(Debug, Clone)]
214pub struct ApiVersion {
215    pub version: u32,
216    pub deprecated: bool,
217    pub min_compatible: u32,
218}
219
220#[derive(Default)]
221pub struct ApiVersionRegistry {
222    versions: Mutex<Vec<ApiVersion>>,
223}
224
225impl ApiVersionRegistry {
226    pub fn new() -> Self {
227        Self::default()
228    }
229
230    pub fn register(&self, version: u32, min_compatible: u32) {
231        self.versions.lock().unwrap().push(ApiVersion {
232            version,
233            deprecated: false,
234            min_compatible,
235        });
236    }
237
238    pub fn deprecate(&self, version: u32) {
239        if let Some(v) = self
240            .versions
241            .lock()
242            .unwrap()
243            .iter_mut()
244            .find(|v| v.version == version)
245        {
246            v.deprecated = true;
247        }
248    }
249
250    pub fn is_compatible(&self, client_version: u32) -> bool {
251        let versions = self.versions.lock().unwrap();
252        versions
253            .iter()
254            .any(|v| v.version >= client_version && v.version >= v.min_compatible)
255    }
256
257    pub fn latest(&self) -> Option<u32> {
258        self.versions
259            .lock()
260            .unwrap()
261            .iter()
262            .map(|v| v.version)
263            .max()
264    }
265
266    pub fn negotiate(&self, client_version: u32) -> Option<u32> {
267        let versions = self.versions.lock().unwrap();
268        versions
269            .iter()
270            .filter(|v| {
271                !v.deprecated && v.version >= v.min_compatible && v.version >= client_version
272            })
273            .map(|v| v.version)
274            .min()
275    }
276
277    pub fn count(&self) -> usize {
278        self.versions.lock().unwrap().len()
279    }
280}
281
282// === Schema Registry ===
283
284#[derive(Debug, Clone, Serialize, Deserialize)]
285pub struct Schema {
286    pub name: String,
287    pub version: u32,
288    pub fields: Vec<String>,
289}
290
291#[derive(Default)]
292pub struct SchemaRegistry {
293    schemas: Mutex<HashMap<String, Vec<Schema>>>,
294}
295
296impl SchemaRegistry {
297    pub fn new() -> Self {
298        Self::default()
299    }
300
301    pub fn register(&self, schema: Schema) {
302        self.schemas
303            .lock()
304            .unwrap()
305            .entry(schema.name.clone())
306            .or_default()
307            .push(schema);
308    }
309
310    pub fn latest(&self, name: &str) -> Option<Schema> {
311        self.schemas
312            .lock()
313            .unwrap()
314            .get(name)
315            .and_then(|versions| versions.last().cloned())
316    }
317
318    pub fn get(&self, name: &str, version: u32) -> Option<Schema> {
319        self.schemas
320            .lock()
321            .unwrap()
322            .get(name)
323            .and_then(|versions| versions.iter().find(|s| s.version == version).cloned())
324    }
325
326    pub fn is_backward_compatible(&self, name: &str, new_version: u32) -> bool {
327        let schemas = self.schemas.lock().unwrap();
328        let versions = match schemas.get(name) {
329            Some(v) => v,
330            None => return true,
331        };
332        let old = match versions
333            .iter()
334            .filter(|s| s.version < new_version)
335            .max_by_key(|s| s.version)
336        {
337            Some(s) => s,
338            None => return true,
339        };
340        let new = match versions.iter().find(|s| s.version == new_version) {
341            Some(s) => s,
342            None => return true,
343        };
344        old.fields.iter().all(|f| new.fields.contains(f))
345    }
346
347    pub fn version_count(&self, name: &str) -> usize {
348        self.schemas
349            .lock()
350            .unwrap()
351            .get(name)
352            .map(|v| v.len())
353            .unwrap_or(0)
354    }
355}
356
357// === Two-Phase Commit (2PC) ===
358
359#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
360#[serde(rename_all = "snake_case")]
361pub enum TwoPcState {
362    Init,
363    Prepared,
364    Committed,
365    Aborted,
366}
367
368#[derive(Debug, Clone, Serialize, Deserialize)]
369pub struct TwoPcParticipant {
370    pub id: String,
371    pub state: TwoPcState,
372}
373
374pub struct TwoPcCoordinator {
375    participants: Mutex<Vec<TwoPcParticipant>>,
376    global_state: Mutex<TwoPcState>,
377}
378
379impl TwoPcCoordinator {
380    pub fn new(participant_ids: &[&str]) -> Self {
381        let participants = participant_ids
382            .iter()
383            .map(|id| TwoPcParticipant {
384                id: id.to_string(),
385                state: TwoPcState::Init,
386            })
387            .collect();
388        Self {
389            participants: Mutex::new(participants),
390            global_state: Mutex::new(TwoPcState::Init),
391        }
392    }
393
394    pub fn prepare(&self, participant_id: &str) -> Result<(), String> {
395        let mut participants = self.participants.lock().unwrap();
396        let p = participants
397            .iter_mut()
398            .find(|p| p.id == participant_id)
399            .ok_or("unknown participant")?;
400        if p.state != TwoPcState::Init {
401            return Err("not in init state".into());
402        }
403        p.state = TwoPcState::Prepared;
404        Ok(())
405    }
406
407    pub fn all_prepared(&self) -> bool {
408        self.participants
409            .lock()
410            .unwrap()
411            .iter()
412            .all(|p| p.state == TwoPcState::Prepared)
413    }
414
415    pub fn commit(&self) -> Result<(), String> {
416        if !self.all_prepared() {
417            return Err("not all prepared".into());
418        }
419        let mut participants = self.participants.lock().unwrap();
420        for p in participants.iter_mut() {
421            p.state = TwoPcState::Committed;
422        }
423        *self.global_state.lock().unwrap() = TwoPcState::Committed;
424        Ok(())
425    }
426
427    pub fn abort(&self) {
428        let mut participants = self.participants.lock().unwrap();
429        for p in participants.iter_mut() {
430            if p.state != TwoPcState::Committed {
431                p.state = TwoPcState::Aborted;
432            }
433        }
434        *self.global_state.lock().unwrap() = TwoPcState::Aborted;
435    }
436
437    pub fn global_state(&self) -> TwoPcState {
438        self.global_state.lock().unwrap().clone()
439    }
440
441    pub fn participant_count(&self) -> usize {
442        self.participants.lock().unwrap().len()
443    }
444}
445
446// === WAL Streaming ===
447
448#[derive(Debug, Clone, Serialize, Deserialize)]
449pub struct WalStreamEntry {
450    pub sequence: u64,
451    pub data_hex: String,
452}
453
454pub struct WalStreamer {
455    entries: Mutex<Vec<WalStreamEntry>>,
456    subscribers: Mutex<Vec<u64>>, // last_seq per subscriber
457    next_seq: Mutex<u64>,
458}
459
460impl WalStreamer {
461    pub fn new() -> Self {
462        Self {
463            entries: Mutex::new(Vec::new()),
464            subscribers: Mutex::new(Vec::new()),
465            next_seq: Mutex::new(1),
466        }
467    }
468
469    pub fn append(&self, data: &[u8]) -> u64 {
470        let seq = {
471            let mut s = self.next_seq.lock().unwrap();
472            let v = *s;
473            *s += 1;
474            v
475        };
476        self.entries.lock().unwrap().push(WalStreamEntry {
477            sequence: seq,
478            data_hex: hex::encode(data),
479        });
480        seq
481    }
482
483    pub fn subscribe(&self) -> usize {
484        let last_seq = self
485            .entries
486            .lock()
487            .unwrap()
488            .last()
489            .map(|e| e.sequence)
490            .unwrap_or(0);
491        self.subscribers.lock().unwrap().push(last_seq);
492        self.subscribers.lock().unwrap().len() - 1
493    }
494
495    pub fn stream_since(&self, subscriber_id: usize) -> Vec<WalStreamEntry> {
496        let last_seen = self
497            .subscribers
498            .lock()
499            .unwrap()
500            .get(subscriber_id)
501            .copied()
502            .unwrap_or(0);
503        let entries = self.entries.lock().unwrap();
504        let new_entries: Vec<WalStreamEntry> = entries
505            .iter()
506            .filter(|e| e.sequence > last_seen)
507            .cloned()
508            .collect();
509        if let Some(last) = entries.last() {
510            if let Some(sub) = self.subscribers.lock().unwrap().get_mut(subscriber_id) {
511                *sub = last.sequence;
512            }
513        }
514        new_entries
515    }
516
517    pub fn entry_count(&self) -> usize {
518        self.entries.lock().unwrap().len()
519    }
520    pub fn subscriber_count(&self) -> usize {
521        self.subscribers.lock().unwrap().len()
522    }
523}
524
525impl Default for WalStreamer {
526    fn default() -> Self {
527        Self::new()
528    }
529}
530
531// === Snapshot Isolation ===
532
533#[derive(Debug, Clone)]
534pub struct Snapshot<T: Clone> {
535    pub data: T,
536    pub version: u64,
537    pub timestamp: Instant,
538}
539
540pub struct SnapshotStore<T: Clone> {
541    snapshots: Mutex<Vec<Snapshot<T>>>,
542    current: Mutex<T>,
543    version: Mutex<u64>,
544}
545
546impl<T: Clone> SnapshotStore<T> {
547    pub fn new(initial: T) -> Self {
548        Self {
549            snapshots: Mutex::new(Vec::new()),
550            current: Mutex::new(initial),
551            version: Mutex::new(0),
552        }
553    }
554
555    pub fn write(&self, data: T) -> u64 {
556        let mut current = self.current.lock().unwrap();
557        let v = {
558            let mut ver = self.version.lock().unwrap();
559            let old = *ver;
560            *ver += 1;
561            old
562        };
563        let old_data = current.clone();
564        *current = data;
565        self.snapshots.lock().unwrap().push(Snapshot {
566            data: old_data,
567            version: v,
568            timestamp: Instant::now(),
569        });
570        v + 1
571    }
572
573    pub fn read_current(&self) -> T {
574        self.current.lock().unwrap().clone()
575    }
576
577    pub fn read_at_version(&self, version: u64) -> Option<T> {
578        let snapshots = self.snapshots.lock().unwrap();
579        if version >= *self.version.lock().unwrap() {
580            return Some(self.current.lock().unwrap().clone());
581        }
582        snapshots
583            .iter()
584            .filter(|s| s.version <= version)
585            .max_by_key(|s| s.version)
586            .map(|s| s.data.clone())
587    }
588
589    pub fn current_version(&self) -> u64 {
590        *self.version.lock().unwrap()
591    }
592    pub fn snapshot_count(&self) -> usize {
593        self.snapshots.lock().unwrap().len()
594    }
595
596    pub fn prune_older_than(&self, max_versions: usize) -> usize {
597        let mut snapshots = self.snapshots.lock().unwrap();
598        let total = snapshots.len();
599        if total <= max_versions {
600            return 0;
601        }
602        let to_remove = total - max_versions;
603        snapshots.drain(..to_remove);
604        to_remove
605    }
606}
607
608// === Homomorphic MAC ===
609
610/// MAC that supports homomorphic addition: MAC(m1+m2) = MAC(m1) + MAC(m2).
611pub struct HomomorphicMac {
612    key: [u8; 32],
613}
614
615#[derive(Debug, Clone, PartialEq)]
616pub struct MacTag {
617    pub tag: [u8; 32],
618}
619
620impl HomomorphicMac {
621    pub fn new(key: [u8; 32]) -> Self {
622        Self { key }
623    }
624
625    pub fn mac(&self, message: &[u8]) -> MacTag {
626        let mut mac = HmacSha256::new_from_slice(&self.key).expect("HMAC");
627        mac.update(message);
628        let result = mac.finalize().into_bytes();
629        let mut tag = [0u8; 32];
630        tag.copy_from_slice(&result);
631        MacTag { tag }
632    }
633
634    pub fn verify(&self, message: &[u8], tag: &MacTag) -> bool {
635        let expected = self.mac(message);
636        expected.tag == tag.tag
637    }
638
639    /// Homomorphically combine two MAC tags (XOR).
640    pub fn combine(a: &MacTag, b: &MacTag) -> MacTag {
641        let mut combined = [0u8; 32];
642        combined
643            .iter_mut()
644            .zip(a.tag.iter().zip(b.tag.iter()))
645            .for_each(|(out, (x, y))| *out = x ^ y);
646        MacTag { tag: combined }
647    }
648
649    /// Homomorphically scale a MAC tag by a scalar (repeated XOR).
650    pub fn scale(tag: &MacTag, n: u32) -> MacTag {
651        if n % 2 == 0 {
652            MacTag { tag: [0u8; 32] }
653        } else {
654            tag.clone()
655        }
656    }
657}
658
659#[cfg(test)]
660mod tests {
661    use super::*;
662
663    // PSI
664    #[test]
665    fn psi_finds_intersection() {
666        let a = vec![b"apple".to_vec(), b"banana".to_vec(), b"cherry".to_vec()];
667        let b = vec![b"banana".to_vec(), b"date".to_vec(), b"cherry".to_vec()];
668        let intersection = psi_hash_based(&a, &b, b"salt");
669        assert_eq!(intersection.len(), 2);
670    }
671
672    #[test]
673    fn psi_empty_when_disjoint() {
674        let a = vec![b"a".to_vec()];
675        let b = vec![b"b".to_vec()];
676        assert!(psi_hash_based(&a, &b, b"salt").is_empty());
677    }
678
679    #[test]
680    fn psi_cardinality_works() {
681        let a = vec![b"x".to_vec(), b"y".to_vec()];
682        let b = vec![b"x".to_vec(), b"y".to_vec(), b"z".to_vec()];
683        assert_eq!(psi_cardinality(&a, &b, b"s"), 2);
684    }
685
686    // PIR
687    #[test]
688    fn pir_trivial_returns_all() {
689        let db = vec![b"a".to_vec(), b"b".to_vec(), b"c".to_vec()];
690        let result = pir_trivial(&db, 1);
691        assert_eq!(result.len(), 3);
692    }
693
694    #[test]
695    fn pir_batch_retrieves_indices() {
696        let db = vec![b"a".to_vec(), b"b".to_vec(), b"c".to_vec()];
697        let result = pir_batch(&db, &[0, 2]);
698        assert_eq!(result, vec![b"a".to_vec(), b"c".to_vec()]);
699    }
700
701    #[test]
702    fn pir_xor_protocol() {
703        let db = vec![vec![0xAA; 4], vec![0xBB; 4], vec![0xCC; 4]];
704        let (q1, q2) = pir_create_query(1, db.len());
705        let r1 = pir_server_respond(&db, &q1);
706        let r2 = pir_server_respond(&db, &q2);
707        let result = pir_decode(&r1, &r2);
708        assert_eq!(result, vec![0xBB; 4]);
709    }
710
711    // Differential Privacy
712    #[test]
713    fn dp_query_adds_noise() {
714        let noisy = dp_query(100.0, 1.0, 1.0);
715        assert!(noisy != 100.0); // almost certainly noisy
716    }
717
718    #[test]
719    fn dp_count_nonnegative_mostly() {
720        for _ in 0..10 {
721            let noisy = dp_count(100, 1.0);
722            assert!(noisy > 50.0 && noisy < 150.0);
723        }
724    }
725
726    #[test]
727    fn gaussian_noise_is_finite() {
728        let noise = gaussian_noise(1.0, 1.0, 0.001);
729        assert!(noise.is_finite());
730    }
731
732    // Feature Flags
733    #[test]
734    fn flag_enabled() {
735        let flags = FeatureFlags::new();
736        flags.set("feature_x", true);
737        assert!(flags.is_enabled("feature_x"));
738    }
739
740    #[test]
741    fn flag_disabled() {
742        let flags = FeatureFlags::new();
743        flags.set("feature_y", false);
744        assert!(!flags.is_enabled("feature_y"));
745    }
746
747    #[test]
748    fn flag_not_set_defaults_false() {
749        let flags = FeatureFlags::new();
750        assert!(!flags.is_enabled("nonexistent"));
751    }
752
753    #[test]
754    fn rollout_100_enables_all() {
755        let flags = FeatureFlags::new();
756        flags.set_rollout("beta", 100);
757        assert!(flags.is_enabled_for("beta", "user1"));
758        assert!(flags.is_enabled_for("beta", "user2"));
759    }
760
761    #[test]
762    fn rollout_0_disables_all() {
763        let flags = FeatureFlags::new();
764        flags.set_rollout("alpha", 0);
765        assert!(!flags.is_enabled_for("alpha", "user1"));
766    }
767
768    // API Versioning
769    #[test]
770    fn version_registration() {
771        let reg = ApiVersionRegistry::new();
772        reg.register(1, 1);
773        reg.register(2, 1);
774        assert_eq!(reg.latest(), Some(2));
775        assert_eq!(reg.count(), 2);
776    }
777
778    #[test]
779    fn version_negotiation() {
780        let reg = ApiVersionRegistry::new();
781        reg.register(1, 1);
782        reg.register(2, 1);
783        reg.register(3, 2);
784        assert_eq!(reg.negotiate(1), Some(1));
785        assert_eq!(reg.negotiate(3), Some(3));
786    }
787
788    #[test]
789    fn version_deprecation() {
790        let reg = ApiVersionRegistry::new();
791        reg.register(1, 1);
792        reg.register(2, 1);
793        reg.deprecate(1);
794        assert_eq!(reg.negotiate(1), Some(2));
795    }
796
797    // Schema Registry
798    #[test]
799    fn schema_register_and_get() {
800        let reg = SchemaRegistry::new();
801        reg.register(Schema {
802            name: "user".into(),
803            version: 1,
804            fields: vec!["id".into(), "name".into()],
805        });
806        reg.register(Schema {
807            name: "user".into(),
808            version: 2,
809            fields: vec!["id".into(), "name".into(), "email".into()],
810        });
811        assert_eq!(reg.version_count("user"), 2);
812        let latest = reg.latest("user").unwrap();
813        assert_eq!(latest.version, 2);
814        assert!(latest.fields.contains(&"email".into()));
815    }
816
817    #[test]
818    fn schema_backward_compat() {
819        let reg = SchemaRegistry::new();
820        reg.register(Schema {
821            name: "event".into(),
822            version: 1,
823            fields: vec!["id".into()],
824        });
825        reg.register(Schema {
826            name: "event".into(),
827            version: 2,
828            fields: vec!["id".into(), "ts".into()],
829        });
830        assert!(reg.is_backward_compatible("event", 2));
831        reg.register(Schema {
832            name: "event".into(),
833            version: 3,
834            fields: vec!["ts".into()],
835        }); // removed "id"
836        assert!(!reg.is_backward_compatible("event", 3));
837    }
838
839    // 2PC
840    #[test]
841    fn two_pc_success() {
842        let coord = TwoPcCoordinator::new(&["a", "b", "c"]);
843        coord.prepare("a").unwrap();
844        coord.prepare("b").unwrap();
845        coord.prepare("c").unwrap();
846        assert!(coord.all_prepared());
847        coord.commit().unwrap();
848        assert_eq!(coord.global_state(), TwoPcState::Committed);
849    }
850
851    #[test]
852    fn two_pc_abort() {
853        let coord = TwoPcCoordinator::new(&["a", "b"]);
854        coord.prepare("a").unwrap();
855        assert!(!coord.all_prepared());
856        coord.abort();
857        assert_eq!(coord.global_state(), TwoPcState::Aborted);
858    }
859
860    #[test]
861    fn two_pc_commit_without_all_prepared_fails() {
862        let coord = TwoPcCoordinator::new(&["a", "b"]);
863        coord.prepare("a").unwrap();
864        assert!(coord.commit().is_err());
865    }
866
867    // WAL Streaming
868    #[test]
869    fn wal_stream_append_and_subscribe() {
870        let streamer = WalStreamer::new();
871        streamer.append(b"data1");
872        streamer.append(b"data2");
873        let sub_id = streamer.subscribe();
874        assert!(streamer.stream_since(sub_id).is_empty());
875        streamer.append(b"data3");
876        let entries = streamer.stream_since(sub_id);
877        assert_eq!(entries.len(), 1);
878        assert_eq!(entries[0].data_hex, hex::encode(b"data3"));
879    }
880
881    #[test]
882    fn wal_stream_multiple_subscribers() {
883        let streamer = WalStreamer::new();
884        let sub1 = streamer.subscribe();
885        streamer.append(b"a");
886        let sub2 = streamer.subscribe();
887        streamer.append(b"b");
888        assert_eq!(streamer.stream_since(sub1).len(), 2);
889        assert_eq!(streamer.stream_since(sub2).len(), 1);
890    }
891
892    // Snapshot Isolation
893    #[test]
894    fn snapshot_read_current() {
895        let store = SnapshotStore::new(42);
896        assert_eq!(store.read_current(), 42);
897        store.write(100);
898        assert_eq!(store.read_current(), 100);
899    }
900
901    #[test]
902    fn snapshot_read_at_version() {
903        let store = SnapshotStore::new(1);
904        store.write(2);
905        store.write(3);
906        assert_eq!(store.read_at_version(0), Some(1));
907        assert_eq!(store.read_at_version(1), Some(2));
908        assert_eq!(store.read_at_version(2), Some(3));
909    }
910
911    #[test]
912    fn snapshot_prune() {
913        let store = SnapshotStore::new(0);
914        for i in 1..=10 {
915            store.write(i);
916        }
917        assert_eq!(store.snapshot_count(), 10);
918        store.prune_older_than(3);
919        assert!(store.snapshot_count() <= 3);
920    }
921
922    // Homomorphic MAC
923    #[test]
924    fn hmac_mac_and_verify() {
925        let mac = HomomorphicMac::new([0x42; 32]);
926        let tag = mac.mac(b"message");
927        assert!(mac.verify(b"message", &tag));
928        assert!(!mac.verify(b"wrong", &tag));
929    }
930
931    #[test]
932    fn hmac_combine() {
933        let mac = HomomorphicMac::new([0x42; 32]);
934        let t1 = mac.mac(b"m1");
935        let t2 = mac.mac(b"m2");
936        let combined = HomomorphicMac::combine(&t1, &t2);
937        assert_ne!(combined.tag, t1.tag);
938        assert_ne!(combined.tag, t2.tag);
939    }
940
941    #[test]
942    fn hmac_scale_even_to_zero() {
943        let mac = HomomorphicMac::new([0x42; 32]);
944        let tag = mac.mac(b"message");
945        let scaled = HomomorphicMac::scale(&tag, 2);
946        assert_eq!(scaled.tag, [0u8; 32]);
947    }
948
949    #[test]
950    fn hmac_scale_odd_preserves() {
951        let mac = HomomorphicMac::new([0x42; 32]);
952        let tag = mac.mac(b"message");
953        let scaled = HomomorphicMac::scale(&tag, 3);
954        assert_eq!(scaled, tag);
955    }
956}