Skip to main content

sz_orm_crypto/
lib.rs

1//! # SZ-ORM Crypto — 加密工具
2//!
3//! 提供常用密码学原语:AES-256-GCM 对称加密、HMAC-SHA256 消息认证码、
4//! PBKDF2 密钥派生与 SHA-256 哈希,所有实现基于 RustCrypto,保证常数时间比较。
5//!
6//! ## 主要函数
7//!
8//! - [`sha256`] / [`sha256_hex`] — SHA-256 哈希
9//! - AES-256-GCM 加解密
10//! - HMAC-SHA256 与 PBKDF2
11
12use std::collections::HashMap;
13
14use aes_gcm::aead::{Aead, KeyInit};
15use aes_gcm::{Aes256Gcm, Key, Nonce};
16use hmac::{Hmac, Mac};
17use pbkdf2::pbkdf2_hmac;
18use rand::rngs::OsRng;
19use rand::RngCore;
20use sha2::{Digest, Sha256};
21use subtle::ConstantTimeEq;
22
23type HmacSha256 = Hmac<Sha256>;
24
25// ============================================================================
26// SHA-256 (基于 RustCrypto sha2 crate, FIPS 180-4)
27// ============================================================================
28
29/// 计算 SHA-256 哈希(基于 RustCrypto sha2)
30pub fn sha256(data: &[u8]) -> [u8; 32] {
31    let mut hasher = Sha256::new();
32    hasher.update(data);
33    let result = hasher.finalize();
34    let mut out = [0u8; 32];
35    out.copy_from_slice(&result);
36    out
37}
38
39/// 计算 SHA-256 并返回十六进制字符串
40pub fn sha256_hex(data: &[u8]) -> String {
41    sha256(data).iter().map(|b| format!("{:02x}", b)).collect()
42}
43
44/// HMAC-SHA256 (RFC 2104, 基于 RustCrypto hmac crate)
45pub fn hmac_sha256(key: &[u8], message: &[u8]) -> [u8; 32] {
46    // HMAC-SHA256 按 RFC 2104 接受任意长度 key,RustCrypto 的 new_from_slice 对 HMAC 永远返回 Ok。
47    // 用 match 处理避免 panic,虽然 Err 分支不可达(RustCrypto 不变量保证)。
48    let mut mac = match <HmacSha256 as Mac>::new_from_slice(key) {
49        Ok(m) => m,
50        Err(_) => {
51            // 不可达分支:HMAC 规范允许任意 key 长度,RustCrypto 内部会先 hash 过长 key。
52            // 为安全起见返回全零(调用方在正常路径下永远不会命中此分支)。
53            return [0u8; 32];
54        }
55    };
56    mac.update(message);
57    let result = mac.finalize().into_bytes();
58    let mut out = [0u8; 32];
59    out.copy_from_slice(&result);
60    out
61}
62
63/// HMAC-SHA256 十六进制字符串
64pub fn hmac_sha256_hex(key: &[u8], message: &[u8]) -> String {
65    hmac_sha256(key, message)
66        .iter()
67        .map(|b| format!("{:02x}", b))
68        .collect()
69}
70
71/// 常量时间比较(基于 subtle crate),避免时序攻击
72fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
73    a.ct_eq(b).into()
74}
75
76// ============================================================================
77// 加密器
78// ============================================================================
79
80pub trait Crypter: Send + Sync {
81    fn encrypt(&self, plaintext: &[u8]) -> Result<Vec<u8>, CryptoError>;
82    fn decrypt(&self, ciphertext: &[u8]) -> Result<Vec<u8>, CryptoError>;
83}
84
85/// AES-256-GCM 加密器(密码学安全)
86///
87/// 使用 AES-256-GCM AEAD 算法,每次加密生成随机 12 字节 nonce。
88/// 密文格式:`nonce(12) || ciphertext || tag(16)`(由 aes-gcm crate 内部处理)。
89pub struct AesGcmCrypter {
90    cipher: Aes256Gcm,
91}
92
93impl AesGcmCrypter {
94    /// 从 32 字节密钥创建
95    pub fn new(key: &[u8; 32]) -> Self {
96        let key = Key::<Aes256Gcm>::from_slice(key);
97        Self {
98            cipher: Aes256Gcm::new(key),
99        }
100    }
101
102    /// 从任意长度密钥字符串创建(SHA-256 派生 32 字节密钥)
103    pub fn from_key_str(key: &str) -> Self {
104        let hash = sha256(key.as_bytes());
105        Self::new(&hash)
106    }
107
108    fn random_nonce() -> [u8; 12] {
109        let mut nonce = [0u8; 12];
110        OsRng.fill_bytes(&mut nonce);
111        nonce
112    }
113}
114
115impl Crypter for AesGcmCrypter {
116    fn encrypt(&self, plaintext: &[u8]) -> Result<Vec<u8>, CryptoError> {
117        self.encrypt_with_aad(plaintext, &[])
118    }
119
120    fn decrypt(&self, ciphertext: &[u8]) -> Result<Vec<u8>, CryptoError> {
121        self.decrypt_with_aad(ciphertext, &[])
122    }
123}
124
125impl AesGcmCrypter {
126    /// AES-GCM 认证加密(带附加认证数据 AAD)
127    ///
128    /// 密文格式:`nonce(12) || ciphertext || tag(16)`
129    /// AAD(Additional Authenticated Data)不包含在密文中,但参与认证标签计算,
130    /// 解密时必须提供相同的 AAD 才能成功。
131    pub fn encrypt_with_aad(&self, plaintext: &[u8], aad: &[u8]) -> Result<Vec<u8>, CryptoError> {
132        let nonce_bytes = Self::random_nonce();
133        let nonce = Nonce::from_slice(&nonce_bytes);
134        let payload = aes_gcm::aead::Payload {
135            msg: plaintext,
136            aad,
137        };
138        let ciphertext = self
139            .cipher
140            .encrypt(nonce, payload)
141            .map_err(|e| CryptoError::EncryptionFailed(e.to_string()))?;
142        let mut result = Vec::with_capacity(12 + ciphertext.len());
143        result.extend_from_slice(&nonce_bytes);
144        result.extend_from_slice(&ciphertext);
145        Ok(result)
146    }
147
148    /// AES-GCM 认证解密(带附加认证数据 AAD)
149    ///
150    /// 必须提供与加密时相同的 AAD,否则认证标签校验失败。
151    pub fn decrypt_with_aad(&self, ciphertext: &[u8], aad: &[u8]) -> Result<Vec<u8>, CryptoError> {
152        if ciphertext.len() < 12 {
153            return Err(CryptoError::DecryptionFailed(
154                "Ciphertext too short".to_string(),
155            ));
156        }
157        let nonce = Nonce::from_slice(&ciphertext[..12]);
158        let encrypted = &ciphertext[12..];
159        let payload = aes_gcm::aead::Payload {
160            msg: encrypted,
161            aad,
162        };
163        self.cipher
164            .decrypt(nonce, payload)
165            .map_err(|e| CryptoError::DecryptionFailed(e.to_string()))
166    }
167}
168
169// ============================================================================
170// 密码哈希
171// ============================================================================
172
173pub trait PasswordHasher: Send + Sync {
174    fn hash(&self, password: &str) -> Result<String, CryptoError>;
175    fn verify(&self, password: &str, hash: &str) -> Result<bool, CryptoError>;
176}
177
178/// PBKDF2-HMAC-SHA256 密码哈希器(基于 RustCrypto pbkdf2 crate)
179///
180/// 使用 PBKDF2-HMAC-SHA256 算法(RFC 8018)。
181/// 哈希格式:`$<iterations>$<salt_hex>$<hash_hex>`
182pub struct Pbkdf2Hasher {
183    iterations: u32,
184}
185
186impl Pbkdf2Hasher {
187    const DEFAULT_ITERATIONS: u32 = 100_000;
188    const SALT_LEN: usize = 16;
189    const HASH_LEN: usize = 32;
190
191    pub fn new() -> Self {
192        Self {
193            iterations: Self::DEFAULT_ITERATIONS,
194        }
195    }
196
197    pub fn with_iterations(iterations: u32) -> Self {
198        Self {
199            iterations: iterations.max(1),
200        }
201    }
202
203    fn compute_hash(password: &str, salt: &[u8], iterations: u32) -> [u8; Self::HASH_LEN] {
204        let mut out = [0u8; Self::HASH_LEN];
205        pbkdf2_hmac::<Sha256>(password.as_bytes(), salt, iterations, &mut out);
206        out
207    }
208}
209
210impl Default for Pbkdf2Hasher {
211    fn default() -> Self {
212        Self::new()
213    }
214}
215
216impl PasswordHasher for Pbkdf2Hasher {
217    fn hash(&self, password: &str) -> Result<String, CryptoError> {
218        if password.is_empty() {
219            return Err(CryptoError::InvalidHash(
220                "Password cannot be empty".to_string(),
221            ));
222        }
223        let salt = random_bytes(Self::SALT_LEN);
224        let hash = Self::compute_hash(password, &salt, self.iterations);
225        Ok(format!(
226            "${}${}${}",
227            self.iterations,
228            hex_encode(&salt),
229            hex_encode(&hash)
230        ))
231    }
232
233    fn verify(&self, password: &str, hash: &str) -> Result<bool, CryptoError> {
234        if !hash.starts_with('$') {
235            return Err(CryptoError::InvalidHash("Invalid hash format".to_string()));
236        }
237        let parts: Vec<&str> = hash[1..].splitn(3, '$').collect();
238        if parts.len() != 3 {
239            return Err(CryptoError::InvalidHash("Invalid hash format".to_string()));
240        }
241        let iterations: u32 = parts[0]
242            .parse()
243            .map_err(|_| CryptoError::InvalidHash("Invalid iterations".to_string()))?;
244        let salt = hex_decode(parts[1])
245            .map_err(|_| CryptoError::InvalidHash("Invalid salt hex".to_string()))?;
246        let expected_hash = hex_decode(parts[2])
247            .map_err(|_| CryptoError::InvalidHash("Invalid hash hex".to_string()))?;
248        let computed = Self::compute_hash(password, &salt, iterations);
249        Ok(constant_time_eq(&computed, &expected_hash))
250    }
251}
252
253// ============================================================================
254// API 签名
255// ============================================================================
256
257pub trait ApiSigner: Send + Sync {
258    fn sign(&self, params: &HashMap<String, String>, secret: &str) -> String;
259    fn verify(&self, params: &HashMap<String, String>, secret: &str, signature: &str) -> bool;
260}
261
262/// HMAC-SHA256 API 签名器
263///
264/// 对参数按字典序排序后拼接成 query string,再用 HMAC-SHA256 签名。
265pub struct HmacSigner;
266
267impl HmacSigner {
268    pub fn new() -> Self {
269        Self
270    }
271
272    fn compute_signature(params: &HashMap<String, String>, secret: &str) -> String {
273        let mut sorted: Vec<_> = params.iter().collect();
274        sorted.sort_by(|a, b| a.0.cmp(b.0));
275
276        let query_string: String = sorted
277            .iter()
278            .map(|(k, v)| format!("{}={}", k, v))
279            .collect::<Vec<_>>()
280            .join("&");
281
282        hmac_sha256_hex(secret.as_bytes(), query_string.as_bytes())
283    }
284}
285
286impl Default for HmacSigner {
287    fn default() -> Self {
288        Self::new()
289    }
290}
291
292impl ApiSigner for HmacSigner {
293    fn sign(&self, params: &HashMap<String, String>, secret: &str) -> String {
294        Self::compute_signature(params, secret)
295    }
296
297    fn verify(&self, params: &HashMap<String, String>, secret: &str, signature: &str) -> bool {
298        let computed = Self::compute_signature(params, secret);
299        constant_time_eq(computed.as_bytes(), signature.as_bytes())
300    }
301}
302
303// ============================================================================
304// RSA-OAEP 非对称加密
305// ============================================================================
306
307use rsa::oaep::Oaep;
308use rsa::{RsaPrivateKey, RsaPublicKey};
309use sha2::Sha256 as RsaSha256;
310
311/// RSA-OAEP 非对称加密器(基于 RustCrypto `rsa` crate)
312///
313/// 使用 RSA-OAEP with SHA-256 和 MGF1-SHA256 填充方案。
314/// 公钥加密,私钥解密,适用于小数据(如密钥交换、短消息加密)。
315pub struct RsaOaepCrypter {
316    public_key: RsaPublicKey,
317    private_key: RsaPrivateKey,
318}
319
320impl RsaOaepCrypter {
321    /// 生成新的 RSA 密钥对(指定位数,推荐 2048 或 3072)
322    pub fn generate(key_bits: usize) -> Result<Self, CryptoError> {
323        let mut rng = OsRng;
324        let private_key = RsaPrivateKey::new(&mut rng, key_bits)
325            .map_err(|e| CryptoError::InvalidKey(e.to_string()))?;
326        let public_key = RsaPublicKey::from(&private_key);
327        Ok(Self {
328            public_key,
329            private_key,
330        })
331    }
332
333    /// 从已有密钥对创建
334    pub fn from_keys(public_key: RsaPublicKey, private_key: RsaPrivateKey) -> Self {
335        Self {
336            public_key,
337            private_key,
338        }
339    }
340
341    /// 返回公钥引用
342    pub fn public_key(&self) -> &RsaPublicKey {
343        &self.public_key
344    }
345
346    /// 返回私钥引用
347    pub fn private_key(&self) -> &RsaPrivateKey {
348        &self.private_key
349    }
350
351    /// 使用公钥加密数据(RSA-OAEP with SHA-256)
352    pub fn encrypt(&self, plaintext: &[u8]) -> Result<Vec<u8>, CryptoError> {
353        let mut rng = OsRng;
354        let padding = Oaep::new::<RsaSha256>();
355        self.public_key
356            .encrypt(&mut rng, padding, plaintext)
357            .map_err(|e| CryptoError::EncryptionFailed(e.to_string()))
358    }
359
360    /// 使用私钥解密数据(RSA-OAEP with SHA-256)
361    pub fn decrypt(&self, ciphertext: &[u8]) -> Result<Vec<u8>, CryptoError> {
362        let padding = Oaep::new::<RsaSha256>();
363        self.private_key
364            .decrypt(padding, ciphertext)
365            .map_err(|e| CryptoError::DecryptionFailed(e.to_string()))
366    }
367}
368
369impl Crypter for RsaOaepCrypter {
370    fn encrypt(&self, plaintext: &[u8]) -> Result<Vec<u8>, CryptoError> {
371        self.encrypt(plaintext)
372    }
373
374    fn decrypt(&self, ciphertext: &[u8]) -> Result<Vec<u8>, CryptoError> {
375        self.decrypt(ciphertext)
376    }
377}
378
379// ============================================================================
380// HMAC-SHA256 签名验证器
381// ============================================================================
382
383/// 签名验证器 trait:提供消息签名与验证接口
384pub trait SignatureVerifier: Send + Sync {
385    /// 对消息生成签名
386    fn sign(&self, message: &[u8]) -> Vec<u8>;
387    /// 验证消息签名(常量时间比较)
388    fn verify(&self, message: &[u8], signature: &[u8]) -> bool;
389}
390
391/// HMAC-SHA256 签名验证器
392///
393/// 使用 HMAC-SHA256 算法对消息签名,验证时采用常量时间比较防止时序攻击。
394pub struct HmacSignatureVerifier {
395    key: Vec<u8>,
396}
397
398impl HmacSignatureVerifier {
399    /// 创建签名验证器,从任意长度密钥派生
400    pub fn new(key: &[u8]) -> Self {
401        Self { key: key.to_vec() }
402    }
403
404    /// 从字符串密钥创建
405    pub fn from_key_str(key: &str) -> Self {
406        Self::new(key.as_bytes())
407    }
408}
409
410impl SignatureVerifier for HmacSignatureVerifier {
411    fn sign(&self, message: &[u8]) -> Vec<u8> {
412        hmac_sha256(&self.key, message).to_vec()
413    }
414
415    fn verify(&self, message: &[u8], signature: &[u8]) -> bool {
416        let expected = self.sign(message);
417        constant_time_eq(&expected, signature)
418    }
419}
420
421// ============================================================================
422// 密钥轮换(Key Rotation)
423// ============================================================================
424
425/// 密钥版本:保存密钥及其版本号和创建时间
426#[derive(Clone)]
427struct KeyVersion {
428    version: u32,
429    key: Vec<u8>,
430    created_at: u64,
431}
432
433/// 密钥轮换管理器
434///
435/// 管理多个版本的密钥,支持:
436/// - 轮换生成新版本密钥
437/// - 用最新密钥签名
438/// - 用任意历史密钥验证(向后兼容)
439/// - 自动淘汰过期密钥
440pub struct KeyRotationManager {
441    keys: Vec<KeyVersion>,
442    current_version: u32,
443    max_versions: usize,
444}
445
446impl KeyRotationManager {
447    /// 创建密钥轮换管理器,指定最大保留版本数
448    pub fn new(max_versions: usize) -> Self {
449        Self {
450            keys: vec![],
451            current_version: 0,
452            max_versions: max_versions.max(1),
453        }
454    }
455
456    /// 初始化首个密钥版本
457    pub fn with_initial_key(key: Vec<u8>) -> Self {
458        let mut mgr = Self::new(3);
459        mgr.rotate_key(key);
460        mgr
461    }
462
463    /// 轮换到新密钥,返回新版本号
464    pub fn rotate_key(&mut self, new_key: Vec<u8>) -> u32 {
465        self.current_version += 1;
466        let now = current_timestamp_secs();
467        self.keys.push(KeyVersion {
468            version: self.current_version,
469            key: new_key,
470            created_at: now,
471        });
472        // 淘汰过期版本
473        while self.keys.len() > self.max_versions {
474            self.keys.remove(0);
475        }
476        self.current_version
477    }
478
479    /// 使用当前(最新)密钥签名
480    pub fn sign(&self, message: &[u8]) -> (u32, Vec<u8>) {
481        if let Some(kv) = self.keys.last() {
482            let sig = hmac_sha256(&kv.key, message).to_vec();
483            (kv.version, sig)
484        } else {
485            (0, vec![])
486        }
487    }
488
489    /// 验证签名(尝试所有保留的密钥版本)
490    pub fn verify(&self, message: &[u8], version: u32, signature: &[u8]) -> bool {
491        for kv in &self.keys {
492            if kv.version == version {
493                let expected = hmac_sha256(&kv.key, message);
494                return constant_time_eq(&expected, signature);
495            }
496        }
497        false
498    }
499
500    /// 返回当前密钥版本号
501    pub fn current_version(&self) -> u32 {
502        self.current_version
503    }
504
505    /// 返回保留的密钥版本数量
506    pub fn version_count(&self) -> usize {
507        self.keys.len()
508    }
509
510    /// 返回所有保留的版本号
511    pub fn versions(&self) -> Vec<u32> {
512        self.keys.iter().map(|kv| kv.version).collect()
513    }
514
515    /// 返回指定版本密钥的创建时间(Unix 秒),不存在返回 None
516    pub fn key_created_at(&self, version: u32) -> Option<u64> {
517        self.keys
518            .iter()
519            .find(|kv| kv.version == version)
520            .map(|kv| kv.created_at)
521    }
522}
523
524fn current_timestamp_secs() -> u64 {
525    use std::time::{SystemTime, UNIX_EPOCH};
526    SystemTime::now()
527        .duration_since(UNIX_EPOCH)
528        .unwrap_or_default()
529        .as_secs()
530}
531
532// ============================================================================
533// 密钥版本管理与轮换(并发安全)
534// ============================================================================
535
536use std::sync::RwLock;
537use std::time::Duration;
538
539/// 默认密钥轮换间隔:90 天
540const DEFAULT_ROTATION_INTERVAL_SECS: u64 = 90 * 24 * 60 * 60;
541
542/// 带版本的密钥
543#[derive(Debug, Clone)]
544pub struct VersionedKey {
545    /// 密钥版本
546    pub version: u32,
547    /// 密钥字节
548    pub key: Vec<u8>,
549    /// 创建时间
550    pub created_at: std::time::SystemTime,
551}
552
553/// 密钥管理器(支持轮换)
554///
555/// 维护一个当前活跃密钥和最多 3 个历史密钥(用于解密过渡期),
556/// 支持按时间间隔自动轮换检查。所有字段使用 `RwLock` 保护,可安全跨线程共享。
557pub struct KeyManager {
558    /// 当前活跃密钥
559    current: RwLock<VersionedKey>,
560    /// 旧密钥列表(用于解密过渡期)
561    previous: RwLock<Vec<VersionedKey>>,
562    /// 密钥轮换间隔
563    rotation_interval: Duration,
564    /// 上次轮换时间
565    last_rotation: RwLock<std::time::SystemTime>,
566}
567
568impl KeyManager {
569    /// 创建密钥管理器,使用给定的初始密钥(版本号从 1 开始)
570    pub fn new(initial_key: Vec<u8>) -> Self {
571        let now = std::time::SystemTime::now();
572        Self {
573            current: RwLock::new(VersionedKey {
574                version: 1,
575                key: initial_key,
576                created_at: now,
577            }),
578            previous: RwLock::new(Vec::new()),
579            rotation_interval: Duration::from_secs(DEFAULT_ROTATION_INTERVAL_SECS),
580            last_rotation: RwLock::new(now),
581        }
582    }
583
584    /// 设置轮换间隔
585    pub fn with_rotation_interval(mut self, interval: Duration) -> Self {
586        self.rotation_interval = interval;
587        self
588    }
589
590    /// 轮换密钥
591    pub fn rotate(&self, new_key: Vec<u8>) -> Result<(), CryptoError> {
592        let mut current = self.current.write().expect("KeyManager lock poisoned");
593        let mut previous = self.previous.write().expect("KeyManager lock poisoned");
594
595        // 将当前密钥移入旧密钥列表
596        previous.push(current.clone());
597
598        // 保留最近 3 个旧密钥
599        if previous.len() > 3 {
600            previous.remove(0);
601        }
602
603        // 设置新密钥
604        *current = VersionedKey {
605            version: current.version + 1,
606            key: new_key,
607            created_at: std::time::SystemTime::now(),
608        };
609
610        *self
611            .last_rotation
612            .write()
613            .expect("KeyManager last_rotation lock poisoned") = std::time::SystemTime::now();
614        Ok(())
615    }
616
617    /// 检查是否需要轮换
618    pub fn needs_rotation(&self) -> bool {
619        let last = *self
620            .last_rotation
621            .read()
622            .expect("KeyManager last_rotation lock poisoned");
623        std::time::SystemTime::now()
624            .duration_since(last)
625            .map(|d| d >= self.rotation_interval)
626            .unwrap_or(false)
627    }
628
629    /// 获取当前密钥
630    pub fn current_key(&self) -> VersionedKey {
631        self.current
632            .read()
633            .expect("KeyManager current lock poisoned")
634            .clone()
635    }
636
637    /// 按版本查找密钥
638    pub fn key_by_version(&self, version: u32) -> Option<VersionedKey> {
639        if self
640            .current
641            .read()
642            .expect("KeyManager current lock poisoned")
643            .version
644            == version
645        {
646            return Some(
647                self.current
648                    .read()
649                    .expect("KeyManager current lock poisoned")
650                    .clone(),
651            );
652        }
653        self.previous
654            .read()
655            .expect("KeyManager previous lock poisoned")
656            .iter()
657            .find(|k| k.version == version)
658            .cloned()
659    }
660
661    /// 返回保留的旧密钥数量
662    pub fn previous_count(&self) -> usize {
663        self.previous
664            .read()
665            .expect("KeyManager previous lock poisoned")
666            .len()
667    }
668}
669
670// ============================================================================
671// 辅助函数
672// ============================================================================
673
674fn hex_encode(bytes: &[u8]) -> String {
675    bytes.iter().map(|b| format!("{:02x}", b)).collect()
676}
677
678fn hex_decode(hex: &str) -> Result<Vec<u8>, ()> {
679    if !hex.len().is_multiple_of(2) {
680        return Err(());
681    }
682    (0..hex.len())
683        .step_by(2)
684        .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).map_err(|_| ()))
685        .collect()
686}
687
688fn random_bytes(len: usize) -> Vec<u8> {
689    let mut result = vec![0u8; len];
690    OsRng.fill_bytes(&mut result);
691    result
692}
693
694// ============================================================================
695// 错误类型
696// ============================================================================
697
698#[derive(Debug)]
699pub enum CryptoError {
700    EncryptionFailed(String),
701    DecryptionFailed(String),
702    InvalidKey(String),
703    InvalidNonce(String),
704    InvalidHash(String),
705    SigningFailed(String),
706}
707
708impl std::fmt::Display for CryptoError {
709    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
710        match self {
711            CryptoError::EncryptionFailed(msg) => write!(f, "Encryption failed: {}", msg),
712            CryptoError::DecryptionFailed(msg) => write!(f, "Decryption failed: {}", msg),
713            CryptoError::InvalidKey(msg) => write!(f, "Invalid key: {}", msg),
714            CryptoError::InvalidNonce(msg) => write!(f, "Invalid nonce: {}", msg),
715            CryptoError::InvalidHash(msg) => write!(f, "Invalid hash: {}", msg),
716            CryptoError::SigningFailed(msg) => write!(f, "Signing failed: {}", msg),
717        }
718    }
719}
720
721impl std::error::Error for CryptoError {}
722
723impl serde::Serialize for CryptoError {
724    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
725    where
726        S: serde::Serializer,
727    {
728        serializer.serialize_str(&self.to_string())
729    }
730}
731
732// ============================================================================
733// 测试
734// ============================================================================
735
736#[cfg(test)]
737mod tests {
738    use super::*;
739
740    // --- SHA-256 标准测试向量 (FIPS 180-2 / NIST) ---
741
742    #[test]
743    fn test_sha256_empty() {
744        assert_eq!(
745            sha256_hex(b""),
746            "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
747        );
748    }
749
750    #[test]
751    fn test_sha256_abc() {
752        assert_eq!(
753            sha256_hex(b"abc"),
754            "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad"
755        );
756    }
757
758    #[test]
759    fn test_sha256_hello() {
760        assert_eq!(
761            sha256_hex(b"hello"),
762            "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824"
763        );
764    }
765
766    #[test]
767    fn test_sha256_long_message() {
768        assert_eq!(
769            sha256_hex(b"abcdbcdecdefdefgefghfghighijhijkijkljklmklmnlmnomnopnopq"),
770            "248d6a61d20638b8e5c026930c3e6039a33ce45964ff2167f6ecedd419db06c1"
771        );
772    }
773
774    #[test]
775    fn test_sha256_deterministic() {
776        assert_eq!(sha256_hex(b"test"), sha256_hex(b"test"));
777        assert_ne!(sha256_hex(b"test"), sha256_hex(b"Test"));
778    }
779
780    // --- HMAC-SHA256 测试向量 (RFC 4231) ---
781
782    #[test]
783    fn test_hmac_sha256_rfc4231_case1() {
784        let key = vec![0x0bu8; 20];
785        let result = hmac_sha256_hex(&key, b"Hi There");
786        assert_eq!(
787            result,
788            "b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7"
789        );
790    }
791
792    #[test]
793    fn test_hmac_sha256_rfc4231_case2() {
794        let result = hmac_sha256_hex(b"Jefe", b"what do ya want for nothing?");
795        assert_eq!(
796            result,
797            "5bdcc146bf60754e6a042426089575c75a003f089d2739839dec58b964ec3843"
798        );
799    }
800
801    #[test]
802    fn test_hmac_sha256_long_key() {
803        let key = vec![0xaau8; 130];
804        let result = hmac_sha256_hex(&key, b"test message");
805        assert_eq!(result.len(), 64);
806        let short_key = vec![0xaau8; 32];
807        let result_short = hmac_sha256_hex(&short_key, b"test message");
808        assert_ne!(result, result_short);
809    }
810
811    #[test]
812    fn test_hmac_sha256_different_messages() {
813        let key = b"secret";
814        assert_ne!(hmac_sha256_hex(key, b"msg1"), hmac_sha256_hex(key, b"msg2"));
815    }
816
817    // --- AesGcmCrypter 测试 ---
818
819    #[test]
820    fn test_aes_gcm_roundtrip() {
821        let key = [0x42u8; 32];
822        let crypter = AesGcmCrypter::new(&key);
823        let plaintext = b"Hello, World!";
824        let encrypted = crypter.encrypt(plaintext).unwrap();
825        let decrypted = crypter.decrypt(&encrypted).unwrap();
826        assert_eq!(decrypted, plaintext);
827    }
828
829    #[test]
830    fn test_aes_gcm_random_nonce_per_encryption() {
831        let key = [0x42u8; 32];
832        let crypter = AesGcmCrypter::new(&key);
833        let plaintext = b"same plaintext";
834        let encrypted1 = crypter.encrypt(plaintext).unwrap();
835        let encrypted2 = crypter.encrypt(plaintext).unwrap();
836        assert_ne!(encrypted1, encrypted2, "随机 nonce 应使密文不同");
837        assert_eq!(crypter.decrypt(&encrypted1).unwrap(), plaintext);
838        assert_eq!(crypter.decrypt(&encrypted2).unwrap(), plaintext);
839    }
840
841    #[test]
842    fn test_aes_gcm_from_key_str() {
843        let crypter = AesGcmCrypter::from_key_str("my-secret-key");
844        let plaintext = b"data to encrypt";
845        let encrypted = crypter.encrypt(plaintext).unwrap();
846        let decrypted = crypter.decrypt(&encrypted).unwrap();
847        assert_eq!(decrypted, plaintext);
848    }
849
850    #[test]
851    fn test_aes_gcm_short_ciphertext() {
852        let key = [0x42u8; 32];
853        let crypter = AesGcmCrypter::new(&key);
854        assert!(crypter.decrypt(&[0u8; 8]).is_err());
855    }
856
857    #[test]
858    fn test_aes_gcm_empty_plaintext() {
859        let key = [0x42u8; 32];
860        let crypter = AesGcmCrypter::new(&key);
861        let encrypted = crypter.encrypt(b"").unwrap();
862        // nonce(12) + tag(16) = 28
863        assert_eq!(encrypted.len(), 28);
864        let decrypted = crypter.decrypt(&encrypted).unwrap();
865        assert_eq!(decrypted, b"");
866    }
867
868    #[test]
869    fn test_aes_gcm_tampered_ciphertext() {
870        let key = [0x42u8; 32];
871        let crypter = AesGcmCrypter::new(&key);
872        let encrypted = crypter.encrypt(b"sensitive data").unwrap();
873        let mut tampered = encrypted.clone();
874        tampered[15] ^= 0x01;
875        assert!(crypter.decrypt(&tampered).is_err());
876    }
877
878    // --- Pbkdf2Hasher 测试 ---
879
880    #[test]
881    fn test_pbkdf2_hasher_hash_format() {
882        let hasher = Pbkdf2Hasher::new();
883        let hash = hasher.hash("password123").unwrap();
884        assert!(hash.starts_with('$'));
885        let parts: Vec<&str> = hash[1..].splitn(3, '$').collect();
886        assert_eq!(parts.len(), 3);
887        assert_eq!(parts[0].parse::<u32>().unwrap(), 100_000);
888        // salt 32 hex chars (16 bytes)
889        assert_eq!(parts[1].len(), 32);
890        // hash 64 hex chars (32 bytes)
891        assert_eq!(parts[2].len(), 64);
892    }
893
894    #[test]
895    fn test_pbkdf2_hasher_verify_correct() {
896        let hasher = Pbkdf2Hasher::new();
897        let hash = hasher.hash("password123").unwrap();
898        assert!(hasher.verify("password123", &hash).unwrap());
899    }
900
901    #[test]
902    fn test_pbkdf2_hasher_verify_wrong() {
903        let hasher = Pbkdf2Hasher::new();
904        let hash = hasher.hash("password123").unwrap();
905        assert!(!hasher.verify("wrongpassword", &hash).unwrap());
906    }
907
908    #[test]
909    fn test_pbkdf2_hasher_different_passwords_different_hashes() {
910        let hasher = Pbkdf2Hasher::new();
911        let h1 = hasher.hash("pass1").unwrap();
912        let h2 = hasher.hash("pass2").unwrap();
913        assert_ne!(h1, h2);
914    }
915
916    #[test]
917    fn test_pbkdf2_hasher_same_password_different_salts() {
918        let hasher = Pbkdf2Hasher::new();
919        let h1 = hasher.hash("same").unwrap();
920        let h2 = hasher.hash("same").unwrap();
921        assert_ne!(h1, h2);
922        assert!(hasher.verify("same", &h1).unwrap());
923        assert!(hasher.verify("same", &h2).unwrap());
924    }
925
926    #[test]
927    fn test_pbkdf2_hasher_invalid_format() {
928        let hasher = Pbkdf2Hasher::new();
929        assert!(hasher.verify("password", "invalid-hash").is_err());
930        assert!(hasher.verify("password", "$abc").is_err());
931        assert!(hasher.verify("password", "$abc$def").is_err());
932    }
933
934    #[test]
935    fn test_pbkdf2_hasher_with_iterations() {
936        let hasher = Pbkdf2Hasher::with_iterations(1000);
937        let hash = hasher.hash("secret").unwrap();
938        let parts: Vec<&str> = hash[1..].splitn(3, '$').collect();
939        assert_eq!(parts[0], "1000");
940        assert!(hasher.verify("secret", &hash).unwrap());
941    }
942
943    #[test]
944    fn test_pbkdf2_hasher_empty_password() {
945        let hasher = Pbkdf2Hasher::new();
946        assert!(hasher.hash("").is_err());
947    }
948
949    // --- HmacSigner 测试 ---
950
951    #[test]
952    fn test_hmac_signer_sign_not_empty() {
953        let signer = HmacSigner::new();
954        let mut params = HashMap::new();
955        params.insert("name".to_string(), "test".to_string());
956        let signature = signer.sign(&params, "secret123");
957        assert_eq!(signature.len(), 64);
958    }
959
960    #[test]
961    fn test_hmac_signer_verify_correct() {
962        let signer = HmacSigner::new();
963        let mut params = HashMap::new();
964        params.insert("name".to_string(), "test".to_string());
965        params.insert("age".to_string(), "25".to_string());
966
967        let signature = signer.sign(&params, "mysecret");
968        assert!(signer.verify(&params, "mysecret", &signature));
969    }
970
971    #[test]
972    fn test_hmac_signer_verify_wrong_secret() {
973        let signer = HmacSigner::new();
974        let mut params = HashMap::new();
975        params.insert("name".to_string(), "test".to_string());
976        let signature = signer.sign(&params, "correctsecret");
977        assert!(!signer.verify(&params, "wrongsecret", &signature));
978    }
979
980    #[test]
981    fn test_hmac_signer_verify_wrong_signature() {
982        let signer = HmacSigner::new();
983        let mut params = HashMap::new();
984        params.insert("name".to_string(), "test".to_string());
985        let valid_sig = signer.sign(&params, "secret");
986        let tampered = if let Some(stripped) = valid_sig.strip_prefix('0') {
987            format!("1{}", stripped)
988        } else {
989            format!("0{}", &valid_sig[1..])
990        };
991        assert!(!signer.verify(&params, "secret", &tampered));
992    }
993
994    #[test]
995    fn test_hmac_signer_different_params_different_signatures() {
996        let signer = HmacSigner::new();
997        let mut params1 = HashMap::new();
998        params1.insert("a".to_string(), "1".to_string());
999
1000        let mut params2 = HashMap::new();
1001        params2.insert("b".to_string(), "2".to_string());
1002
1003        let sig1 = signer.sign(&params1, "secret");
1004        let sig2 = signer.sign(&params2, "secret");
1005        assert_ne!(sig1, sig2);
1006    }
1007
1008    #[test]
1009    fn test_hmac_signer_param_order_independent() {
1010        let signer = HmacSigner::new();
1011        let mut params1 = HashMap::new();
1012        params1.insert("b".to_string(), "2".to_string());
1013        params1.insert("a".to_string(), "1".to_string());
1014
1015        let mut params2 = HashMap::new();
1016        params2.insert("a".to_string(), "1".to_string());
1017        params2.insert("b".to_string(), "2".to_string());
1018
1019        let sig1 = signer.sign(&params1, "secret");
1020        let sig2 = signer.sign(&params2, "secret");
1021        assert_eq!(sig1, sig2);
1022    }
1023
1024    #[test]
1025    fn test_hmac_signer_empty_params() {
1026        let signer = HmacSigner::new();
1027        let params = HashMap::new();
1028        let sig = signer.sign(&params, "secret");
1029        assert_eq!(sig.len(), 64);
1030        assert!(signer.verify(&params, "secret", &sig));
1031    }
1032
1033    // --- 辅助函数测试 ---
1034
1035    #[test]
1036    fn test_random_bytes_length() {
1037        assert_eq!(random_bytes(0).len(), 0);
1038        assert_eq!(random_bytes(16).len(), 16);
1039        assert_eq!(random_bytes(100).len(), 100);
1040    }
1041
1042    #[test]
1043    fn test_random_bytes_random() {
1044        let a = random_bytes(32);
1045        let b = random_bytes(32);
1046        assert_ne!(a, b, "随机字节序列应不同");
1047    }
1048
1049    #[test]
1050    fn test_constant_time_eq() {
1051        assert!(constant_time_eq(b"abc", b"abc"));
1052        assert!(!constant_time_eq(b"abc", b"abd"));
1053        assert!(!constant_time_eq(b"abc", b"ab"));
1054        assert!(!constant_time_eq(b"abc", b"abcd"));
1055        assert!(constant_time_eq(b"", b""));
1056    }
1057
1058    #[test]
1059    fn test_hex_encode_decode_roundtrip() {
1060        let original = vec![0x00, 0xff, 0xab, 0x42];
1061        let encoded = hex_encode(&original);
1062        let decoded = hex_decode(&encoded).unwrap();
1063        assert_eq!(decoded, original);
1064    }
1065
1066    #[test]
1067    fn test_hex_decode_invalid() {
1068        assert!(hex_decode("abc").is_err());
1069        assert!(hex_decode("xy").is_err());
1070    }
1071
1072    // ===== AES-GCM AAD 测试 =====
1073
1074    #[test]
1075    fn test_aes_gcm_aad_roundtrip() {
1076        let key = [0x42u8; 32];
1077        let crypter = AesGcmCrypter::new(&key);
1078        let plaintext = b"sensitive data";
1079        let aad = b"associated metadata";
1080        let encrypted = crypter.encrypt_with_aad(plaintext, aad).unwrap();
1081        let decrypted = crypter.decrypt_with_aad(&encrypted, aad).unwrap();
1082        assert_eq!(decrypted, plaintext);
1083    }
1084
1085    #[test]
1086    fn test_aes_gcm_aad_wrong_aad_fails() {
1087        let key = [0x42u8; 32];
1088        let crypter = AesGcmCrypter::new(&key);
1089        let plaintext = b"sensitive data";
1090        let aad = b"correct aad";
1091        let encrypted = crypter.encrypt_with_aad(plaintext, aad).unwrap();
1092        // 使用错误的 AAD 解密应失败
1093        let result = crypter.decrypt_with_aad(&encrypted, b"wrong aad");
1094        assert!(result.is_err());
1095    }
1096
1097    #[test]
1098    fn test_aes_gcm_aad_empty_aad_equivalent_to_no_aad() {
1099        let key = [0x42u8; 32];
1100        let crypter = AesGcmCrypter::new(&key);
1101        let plaintext = b"test data";
1102        // 空 AAD 等价于无 AAD
1103        let encrypted_no_aad = crypter.encrypt(plaintext).unwrap();
1104        let encrypted_empty_aad = crypter.encrypt_with_aad(plaintext, b"").unwrap();
1105        // 两者都应能解密
1106        assert_eq!(crypter.decrypt(&encrypted_no_aad).unwrap(), plaintext);
1107        assert_eq!(
1108            crypter.decrypt_with_aad(&encrypted_empty_aad, b"").unwrap(),
1109            plaintext
1110        );
1111    }
1112
1113    #[test]
1114    fn test_aes_gcm_aad_tampered_ciphertext_fails() {
1115        let key = [0x42u8; 32];
1116        let crypter = AesGcmCrypter::new(&key);
1117        let encrypted = crypter.encrypt_with_aad(b"data", b"aad").unwrap();
1118        let mut tampered = encrypted.clone();
1119        tampered[15] ^= 0x01;
1120        assert!(crypter.decrypt_with_aad(&tampered, b"aad").is_err());
1121    }
1122
1123    #[test]
1124    fn test_aes_gcm_aad_empty_plaintext() {
1125        let key = [0x42u8; 32];
1126        let crypter = AesGcmCrypter::new(&key);
1127        let encrypted = crypter.encrypt_with_aad(b"", b"aad").unwrap();
1128        // nonce(12) + tag(16) = 28
1129        assert_eq!(encrypted.len(), 28);
1130        let decrypted = crypter.decrypt_with_aad(&encrypted, b"aad").unwrap();
1131        assert_eq!(decrypted, b"");
1132    }
1133
1134    // ===== RSA-OAEP 测试 =====
1135
1136    #[test]
1137    fn test_rsa_oaep_roundtrip() {
1138        let crypter = RsaOaepCrypter::generate(2048).expect("RSA key generation");
1139        let plaintext = b"Hello, RSA-OAEP!";
1140        let encrypted = crypter.encrypt(plaintext).unwrap();
1141        let decrypted = crypter.decrypt(&encrypted).unwrap();
1142        assert_eq!(decrypted, plaintext);
1143    }
1144
1145    #[test]
1146    fn test_rsa_oaep_different_ciphertexts_same_plaintext() {
1147        let crypter = RsaOaepCrypter::generate(2048).unwrap();
1148        let plaintext = b"same message";
1149        let enc1 = crypter.encrypt(plaintext).unwrap();
1150        let enc2 = crypter.encrypt(plaintext).unwrap();
1151        // OAEP 使用随机填充,相同明文应产生不同密文
1152        assert_ne!(enc1, enc2);
1153        // 但两者都能正确解密
1154        assert_eq!(crypter.decrypt(&enc1).unwrap(), plaintext);
1155        assert_eq!(crypter.decrypt(&enc2).unwrap(), plaintext);
1156    }
1157
1158    #[test]
1159    fn test_rsa_oaep_empty_plaintext() {
1160        let crypter = RsaOaepCrypter::generate(2048).unwrap();
1161        let encrypted = crypter.encrypt(b"").unwrap();
1162        let decrypted = crypter.decrypt(&encrypted).unwrap();
1163        assert_eq!(decrypted, b"");
1164    }
1165
1166    #[test]
1167    fn test_rsa_oaep_tampered_ciphertext_fails() {
1168        let crypter = RsaOaepCrypter::generate(2048).unwrap();
1169        let encrypted = crypter.encrypt(b"secret").unwrap();
1170        let mut tampered = encrypted.clone();
1171        tampered[0] ^= 0x01;
1172        assert!(crypter.decrypt(&tampered).is_err());
1173    }
1174
1175    #[test]
1176    fn test_rsa_oaep_max_message_length() {
1177        // 2048-bit RSA-OAEP with SHA-256: max message = 2048/8 - 2*32 - 2 = 190 bytes
1178        let crypter = RsaOaepCrypter::generate(2048).unwrap();
1179        let plaintext = vec![0xABu8; 190];
1180        let encrypted = crypter.encrypt(&plaintext).unwrap();
1181        let decrypted = crypter.decrypt(&encrypted).unwrap();
1182        assert_eq!(decrypted, plaintext);
1183    }
1184
1185    #[test]
1186    fn test_rsa_oaep_oversized_message_fails() {
1187        let crypter = RsaOaepCrypter::generate(2048).unwrap();
1188        // 超过最大消息长度(190 字节 + 1)
1189        let plaintext = vec![0xABu8; 191];
1190        assert!(crypter.encrypt(&plaintext).is_err());
1191    }
1192
1193    #[test]
1194    fn test_rsa_oaep_from_keys() {
1195        let crypter1 = RsaOaepCrypter::generate(2048).unwrap();
1196        let crypter2 = RsaOaepCrypter::from_keys(
1197            crypter1.public_key().clone(),
1198            crypter1.private_key().clone(),
1199        );
1200        let plaintext = b"test from_keys";
1201        let encrypted = crypter2.encrypt(plaintext).unwrap();
1202        let decrypted = crypter2.decrypt(&encrypted).unwrap();
1203        assert_eq!(decrypted, plaintext);
1204    }
1205
1206    #[test]
1207    fn test_rsa_oaep_crypter_trait() {
1208        let crypter = RsaOaepCrypter::generate(2048).unwrap();
1209        let plaintext = b"trait test";
1210        let encrypted = Crypter::encrypt(&crypter, plaintext).unwrap();
1211        let decrypted = Crypter::decrypt(&crypter, &encrypted).unwrap();
1212        assert_eq!(decrypted, plaintext);
1213    }
1214
1215    // ===== HMAC 签名验证器测试 =====
1216
1217    #[test]
1218    fn test_hmac_signature_verifier_sign_verify() {
1219        let verifier = HmacSignatureVerifier::new(b"my-secret-key");
1220        let message = b"important message";
1221        let signature = verifier.sign(message);
1222        assert_eq!(signature.len(), 32);
1223        assert!(verifier.verify(message, &signature));
1224    }
1225
1226    #[test]
1227    fn test_hmac_signature_verifier_wrong_message() {
1228        let verifier = HmacSignatureVerifier::new(b"key");
1229        let signature = verifier.sign(b"message1");
1230        assert!(!verifier.verify(b"message2", &signature));
1231    }
1232
1233    #[test]
1234    fn test_hmac_signature_verifier_wrong_signature() {
1235        let verifier = HmacSignatureVerifier::new(b"key");
1236        let signature = verifier.sign(b"message");
1237        let mut tampered = signature.clone();
1238        tampered[0] ^= 0x01;
1239        assert!(!verifier.verify(b"message", &tampered));
1240    }
1241
1242    #[test]
1243    fn test_hmac_signature_verifier_from_key_str() {
1244        let verifier = HmacSignatureVerifier::from_key_str("string-key");
1245        let message = b"test";
1246        let sig = verifier.sign(message);
1247        assert!(verifier.verify(message, &sig));
1248    }
1249
1250    #[test]
1251    fn test_hmac_signature_verifier_different_keys_different_signatures() {
1252        let v1 = HmacSignatureVerifier::new(b"key1");
1253        let v2 = HmacSignatureVerifier::new(b"key2");
1254        let message = b"same message";
1255        let sig1 = v1.sign(message);
1256        let sig2 = v2.sign(message);
1257        assert_ne!(sig1, sig2);
1258    }
1259
1260    #[test]
1261    fn test_hmac_signature_verifier_empty_message() {
1262        let verifier = HmacSignatureVerifier::new(b"key");
1263        let sig = verifier.sign(b"");
1264        assert_eq!(sig.len(), 32);
1265        assert!(verifier.verify(b"", &sig));
1266    }
1267
1268    #[test]
1269    fn test_hmac_signature_verifier_wrong_length_signature() {
1270        let verifier = HmacSignatureVerifier::new(b"key");
1271        // 长度不对的签名应验证失败
1272        assert!(!verifier.verify(b"message", b"short"));
1273        assert!(!verifier.verify(b"message", &[]));
1274    }
1275
1276    // ===== 密钥轮换测试 =====
1277
1278    #[test]
1279    fn test_key_rotation_initial_key() {
1280        let mgr = KeyRotationManager::with_initial_key(b"key-v1".to_vec());
1281        assert_eq!(mgr.current_version(), 1);
1282        assert_eq!(mgr.version_count(), 1);
1283        assert_eq!(mgr.versions(), vec![1]);
1284    }
1285
1286    #[test]
1287    fn test_key_rotation_sign_verify_current() {
1288        let mgr = KeyRotationManager::with_initial_key(b"secret-key".to_vec());
1289        let message = b"test message";
1290        let (version, signature) = mgr.sign(message);
1291        assert_eq!(version, 1);
1292        assert!(mgr.verify(message, version, &signature));
1293    }
1294
1295    #[test]
1296    fn test_key_rotation_old_version_still_valid() {
1297        let mut mgr = KeyRotationManager::with_initial_key(b"key-v1".to_vec());
1298        let message = b"persistent message";
1299        let (v1, sig1) = mgr.sign(message);
1300        // 轮换到新密钥
1301        mgr.rotate_key(b"key-v2".to_vec());
1302        let (v2, sig2) = mgr.sign(message);
1303        assert_eq!(v1, 1);
1304        assert_eq!(v2, 2);
1305        // 旧版本签名仍应验证通过
1306        assert!(mgr.verify(message, v1, &sig1));
1307        // 新版本签名也应验证通过
1308        assert!(mgr.verify(message, v2, &sig2));
1309    }
1310
1311    #[test]
1312    fn test_key_rotation_max_versions_evicts_oldest() {
1313        let mut mgr = KeyRotationManager::new(2);
1314        mgr.rotate_key(b"key-v1".to_vec());
1315        mgr.rotate_key(b"key-v2".to_vec());
1316        assert_eq!(mgr.version_count(), 2);
1317        // 第三次轮换应淘汰 v1
1318        mgr.rotate_key(b"key-v3".to_vec());
1319        assert_eq!(mgr.version_count(), 2);
1320        assert_eq!(mgr.versions(), vec![2, 3]);
1321        assert!(!mgr.versions().contains(&1));
1322    }
1323
1324    #[test]
1325    fn test_key_rotation_old_version_evicted_fails_verify() {
1326        let mut mgr = KeyRotationManager::new(2);
1327        mgr.rotate_key(b"key-v1".to_vec());
1328        let message = b"test";
1329        let (v1, sig1) = mgr.sign(message);
1330        mgr.rotate_key(b"key-v2".to_vec());
1331        mgr.rotate_key(b"key-v3".to_vec());
1332        // v1 已被淘汰,验证应失败
1333        assert!(!mgr.verify(message, v1, &sig1));
1334    }
1335
1336    #[test]
1337    fn test_key_rotation_wrong_version_fails() {
1338        let mgr = KeyRotationManager::with_initial_key(b"key".to_vec());
1339        let message = b"test";
1340        let (_, signature) = mgr.sign(message);
1341        // 使用不存在的版本号验证应失败
1342        assert!(!mgr.verify(message, 999, &signature));
1343    }
1344
1345    #[test]
1346    fn test_key_rotation_multiple_rotations() {
1347        let mut mgr = KeyRotationManager::new(5);
1348        for i in 1..=4 {
1349            let key = format!("key-v{}", i);
1350            let version = mgr.rotate_key(key.as_bytes().to_vec());
1351            assert_eq!(version, i as u32);
1352        }
1353        assert_eq!(mgr.current_version(), 4);
1354        assert_eq!(mgr.version_count(), 4);
1355        assert_eq!(mgr.versions(), vec![1, 2, 3, 4]);
1356    }
1357
1358    #[test]
1359    fn test_key_rotation_empty_manager_sign_returns_zero() {
1360        let mgr = KeyRotationManager::new(3);
1361        let (version, sig) = mgr.sign(b"message");
1362        assert_eq!(version, 0);
1363        assert!(sig.is_empty());
1364    }
1365
1366    #[test]
1367    fn test_key_rotation_verify_with_wrong_signature() {
1368        let mgr = KeyRotationManager::with_initial_key(b"key".to_vec());
1369        let message = b"test";
1370        let (version, _) = mgr.sign(message);
1371        let wrong_sig = vec![0u8; 32];
1372        assert!(!mgr.verify(message, version, &wrong_sig));
1373    }
1374
1375    #[test]
1376    fn test_key_rotation_max_versions_min_one() {
1377        // max_versions = 0 应被提升为 1
1378        let mut mgr = KeyRotationManager::new(0);
1379        mgr.rotate_key(b"k1".to_vec());
1380        mgr.rotate_key(b"k2".to_vec());
1381        assert_eq!(mgr.version_count(), 1);
1382        assert_eq!(mgr.versions(), vec![2]);
1383    }
1384
1385    // ===== KeyManager(并发密钥轮换)测试 =====
1386
1387    #[test]
1388    fn test_key_manager_initial_key() {
1389        let mgr = KeyManager::new(b"initial-key".to_vec());
1390        let current = mgr.current_key();
1391        assert_eq!(current.version, 1);
1392        assert_eq!(current.key, b"initial-key");
1393        assert_eq!(mgr.previous_count(), 0);
1394    }
1395
1396    #[test]
1397    fn test_key_manager_rotate_increments_version() {
1398        let mgr = KeyManager::new(b"v1".to_vec());
1399        assert!(mgr.rotate(b"v2".to_vec()).is_ok());
1400        let current = mgr.current_key();
1401        assert_eq!(current.version, 2);
1402        assert_eq!(current.key, b"v2");
1403        assert_eq!(mgr.previous_count(), 1);
1404    }
1405
1406    #[test]
1407    fn test_key_manager_key_by_version_current() {
1408        let mgr = KeyManager::new(b"v1".to_vec());
1409        let found = mgr.key_by_version(1).expect("v1 should exist");
1410        assert_eq!(found.key, b"v1");
1411    }
1412
1413    #[test]
1414    fn test_key_manager_key_by_version_previous() {
1415        let mgr = KeyManager::new(b"v1".to_vec());
1416        mgr.rotate(b"v2".to_vec()).unwrap();
1417        // 旧版本仍可查找
1418        let old = mgr.key_by_version(1).expect("v1 should still be retained");
1419        assert_eq!(old.key, b"v1");
1420        // 新版本也可查找
1421        let new = mgr.key_by_version(2).expect("v2 should exist");
1422        assert_eq!(new.key, b"v2");
1423    }
1424
1425    #[test]
1426    fn test_key_manager_key_by_version_not_found() {
1427        let mgr = KeyManager::new(b"v1".to_vec());
1428        assert!(mgr.key_by_version(999).is_none());
1429    }
1430
1431    #[test]
1432    fn test_key_manager_retains_at_most_three_previous() {
1433        let mgr = KeyManager::new(b"v1".to_vec());
1434        mgr.rotate(b"v2".to_vec()).unwrap();
1435        mgr.rotate(b"v3".to_vec()).unwrap();
1436        mgr.rotate(b"v4".to_vec()).unwrap();
1437        // 3 次轮换后 previous 应有 3 个,再轮换一次应淘汰最早的
1438        assert_eq!(mgr.previous_count(), 3);
1439        mgr.rotate(b"v5".to_vec()).unwrap();
1440        assert_eq!(mgr.previous_count(), 3);
1441        // v1 应已被淘汰
1442        assert!(mgr.key_by_version(1).is_none());
1443        // v2 仍应存在
1444        assert!(mgr.key_by_version(2).is_some());
1445        // 当前版本为 5
1446        assert_eq!(mgr.current_key().version, 5);
1447    }
1448
1449    #[test]
1450    fn test_key_manager_needs_rotation_false_initially() {
1451        let mgr = KeyManager::new(b"k".to_vec());
1452        // 刚创建不应需要轮换
1453        assert!(!mgr.needs_rotation());
1454    }
1455
1456    #[test]
1457    fn test_key_manager_needs_rotation_true_after_interval() {
1458        let mgr = KeyManager::new(b"k".to_vec()).with_rotation_interval(Duration::from_millis(0));
1459        // 间隔为 0,应立即需要轮换
1460        std::thread::sleep(Duration::from_millis(1));
1461        assert!(mgr.needs_rotation());
1462    }
1463
1464    #[test]
1465    fn test_key_manager_with_rotation_interval() {
1466        let mgr = KeyManager::new(b"k".to_vec()).with_rotation_interval(Duration::from_secs(60));
1467        assert!(!mgr.needs_rotation());
1468    }
1469
1470    #[test]
1471    fn test_key_manager_rotate_resets_last_rotation() {
1472        let mgr = KeyManager::new(b"k".to_vec()).with_rotation_interval(Duration::from_millis(1));
1473        std::thread::sleep(Duration::from_millis(5));
1474        assert!(mgr.needs_rotation());
1475        mgr.rotate(b"k2".to_vec()).unwrap();
1476        // 轮换后应不再立即需要轮换
1477        assert!(!mgr.needs_rotation());
1478    }
1479
1480    #[test]
1481    fn test_key_manager_concurrent_access() {
1482        use std::sync::Arc;
1483        use std::thread;
1484        let mgr = Arc::new(KeyManager::new(b"base".to_vec()));
1485        let mut handles = vec![];
1486        // 并发读
1487        for _ in 0..4 {
1488            let m = mgr.clone();
1489            handles.push(thread::spawn(move || {
1490                let _ = m.current_key();
1491                let _ = m.previous_count();
1492            }));
1493        }
1494        for h in handles {
1495            h.join().expect("thread panicked");
1496        }
1497        // 并发读不应改变状态
1498        assert_eq!(mgr.current_key().version, 1);
1499    }
1500}