use std::fs;
use std::path::{Path, PathBuf};
use base64::prelude::*;
use rand::RngExt;
use sha2::{Digest, Sha256};
use crate::error::AuthError;
use super::cipher;
const KEYRING_SERVICE: &str = "wecom-cli";
const KEYRING_USER_PREFIX: &str = "encryption-key";
pub fn encryption_key_path(dir: &Path) -> PathBuf {
dir.join(".encryption_key")
}
pub(crate) fn encode_key(key: &[u8; 32]) -> String {
BASE64_STANDARD.encode(key)
}
pub(crate) fn decode_key(s: &str) -> Result<[u8; 32], AuthError> {
let bytes = BASE64_STANDARD
.decode(s)
.map_err(|e| AuthError::Crypto(format!("加密密钥无效,base64 decode error: {e}")))?;
<[u8; 32]>::try_from(bytes.as_slice())
.map_err(|_| AuthError::Crypto("Invalid encryption key length".into()))
}
pub fn generate_random_key() -> [u8; 32] {
rand::rng().random()
}
pub(crate) fn keyring_user_for(dir: &Path) -> String {
let dir = normalize_path(&std::path::absolute(dir).unwrap_or_else(|_| dir.to_path_buf()));
let digest = Sha256::digest(dir.to_string_lossy().as_bytes());
format!("{KEYRING_USER_PREFIX}:{}", hex::encode(digest))
}
fn normalize_path(path: &Path) -> PathBuf {
let mut out = PathBuf::new();
for comp in path
.components()
.filter(|c| *c != std::path::Component::CurDir)
{
match comp {
std::path::Component::ParentDir => {
out.pop();
}
other => out.push(other.as_os_str()),
}
}
out
}
pub(crate) fn load_key_from_keyring(dir: &Path) -> Option<[u8; 32]> {
let user = keyring_user_for(dir);
let entry = keyring::Entry::new(KEYRING_SERVICE, &user).ok()?;
let b64 = entry.get_password().ok()?;
decode_key(b64.trim()).ok()
}
#[allow(clippy::disallowed_methods)]
pub(crate) fn load_key_from_file(dir: &Path) -> Option<[u8; 32]> {
let contents = fs::read_to_string(encryption_key_path(dir)).ok()?;
decode_key(contents.trim()).ok()
}
pub(crate) fn save_key(dir: &Path, key: &[u8; 32], use_keyring: bool) -> Result<(), AuthError> {
let b64 = encode_key(key);
let key_path = encryption_key_path(dir);
super::atomic_write(&key_path, b64.as_bytes(), 0o600)?;
if !use_keyring {
return Ok(());
}
let user = keyring_user_for(dir);
if let Err(e) =
keyring::Entry::new(KEYRING_SERVICE, &user).and_then(|entry| entry.set_password(&b64))
{
tracing::warn!(error = %e, "keyring unavailable, encryption key stored in file only");
}
Ok(())
}
pub fn encrypt_data<T: serde::Serialize + ?Sized>(
data: &T,
key: &[u8; 32],
) -> Result<Vec<u8>, AuthError> {
let json = serde_json::to_vec(data)
.map_err(|e| AuthError::Crypto(format!("JSON serialize error: {e:#}")))?;
cipher::encrypt(key, &json)
}
pub fn decrypt_data<T: serde::de::DeserializeOwned>(
data: &[u8],
key: &[u8; 32],
) -> Result<T, AuthError> {
let decrypted = cipher::decrypt(key, data)?;
serde_json::from_slice(&decrypted)
.map_err(|e| AuthError::Crypto(format!("JSON deserialize error: {e:#}")))
}
pub(crate) fn try_decrypt_data<T: serde::de::DeserializeOwned>(
dir: &Path,
use_keyring: bool,
data: &[u8],
) -> Result<T, AuthError> {
if let Some(key) = load_key_from_file(dir) {
if let Ok(result) = decrypt_data::<T>(data, &key) {
return Ok(result);
}
tracing::debug!("File key failed to decrypt, falling back to keyring key");
}
let key = load_key_from_keyring(dir)
.filter(|_| use_keyring)
.ok_or_else(|| AuthError::Crypto("解密数据失败(未找到有效密钥)".into()))?;
decrypt_data(data, &key)
}
#[cfg(test)]
mod tests {
use super::*;
use serde::{Deserialize, Serialize};
#[test]
fn encode_decode_roundtrip() {
let key = generate_random_key();
let encoded = encode_key(&key);
let decoded = decode_key(&encoded).unwrap();
assert_eq!(key, decoded);
}
#[test]
fn decode_key_handles_edge_cases() {
assert!(decode_key("not-valid-base64!!!").is_err());
let short = base64::prelude::BASE64_STANDARD.encode([0u8; 16]);
assert!(decode_key(&short).is_err());
let key = generate_random_key();
let encoded = format!(" {} \n", encode_key(&key));
let decoded = decode_key(encoded.trim()).unwrap();
assert_eq!(key, decoded);
}
#[test]
fn random_key_has_expected_properties() {
let key = generate_random_key();
assert_eq!(key.len(), 32);
let another = generate_random_key();
assert_ne!(key, another);
}
#[test]
fn keyring_user_for_deterministic_and_isolated() {
let dir = Path::new("/tmp/wecom-sandbox-a");
assert_eq!(keyring_user_for(dir), keyring_user_for(dir));
let b = keyring_user_for(Path::new("/tmp/wecom-sandbox-b"));
assert_ne!(keyring_user_for(dir), b);
}
#[test]
fn keyring_user_for_normalizes_and_formats() {
let base = keyring_user_for(Path::new("/tmp/wecom-sandbox-a"));
assert_eq!(base, keyring_user_for(Path::new("/tmp/./wecom-sandbox-a")));
assert_eq!(
base,
keyring_user_for(Path::new("/tmp/wecom-sandbox-b/../wecom-sandbox-a"))
);
let suffix = base.strip_prefix("encryption-key:").unwrap();
assert_eq!(suffix.len(), 64);
assert!(
suffix
.bytes()
.all(|b| b.is_ascii_hexdigit() && !b.is_ascii_uppercase()),
"expected lowercase hex, got: {suffix}"
);
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
struct TestPayload {
name: String,
value: u64,
}
#[test]
fn encrypt_decrypt_data_roundtrips() {
let key = generate_random_key();
let payload = TestPayload {
name: "test".into(),
value: 42,
};
let encrypted = encrypt_data(&payload, &key).unwrap();
let decrypted: TestPayload = decrypt_data(&encrypted, &key).unwrap();
assert_eq!(payload, decrypted);
let items = vec![
TestPayload {
name: "a".into(),
value: 1,
},
TestPayload {
name: "b".into(),
value: 2,
},
];
let encrypted = encrypt_data(&items, &key).unwrap();
let decrypted: Vec<TestPayload> = decrypt_data(&encrypted, &key).unwrap();
assert_eq!(items, decrypted);
let empty: Vec<TestPayload> = vec![];
let encrypted = encrypt_data(&empty, &key).unwrap();
let decrypted: Vec<TestPayload> = decrypt_data(&encrypted, &key).unwrap();
assert!(decrypted.is_empty());
}
#[test]
fn decrypt_data_rejects_invalid() {
let key1 = generate_random_key();
let key2 = generate_random_key();
let payload = TestPayload {
name: "secret".into(),
value: 99,
};
let encrypted = encrypt_data(&payload, &key1).unwrap();
assert!(decrypt_data::<TestPayload>(&encrypted, &key2).is_err());
let key = generate_random_key();
assert!(decrypt_data::<TestPayload>(b"garbage", &key).is_err());
}
}