use crate::InklogError;
use base64::{Engine as _, engine::general_purpose};
use pbkdf2::pbkdf2_hmac;
use rand::Rng;
use sha2::Sha256;
use zeroize::Zeroizing;
pub fn get_encryption_key(env_var: &str) -> Result<[u8; 32], InklogError> {
let env_value = Zeroizing::new(std::env::var(env_var).map_err(|_| {
InklogError::ConfigError(
"Encryption key environment variable not set. Please configure INKLOG_ENCRYPTION_KEY."
.to_string(),
)
})?);
let raw_bytes = env_value.as_bytes();
if raw_bytes.len() == 32 {
let mut result = [0u8; 32];
result.copy_from_slice(raw_bytes);
return Ok(result);
}
if let Ok(decoded) = general_purpose::STANDARD.decode(env_value.as_str()) {
if decoded.len() == 32 {
let mut result = [0u8; 32];
result.copy_from_slice(&decoded);
return Ok(result);
}
return Err(InklogError::ConfigError(format!(
"Encryption key from Base64 must be exactly 32 bytes (256 bits), got {} bytes. \
Please provide a valid 32-byte key encoded in Base64.",
decoded.len()
)));
}
if !raw_bytes.is_empty() && raw_bytes.len() < 128 {
let (key, _salt) = derive_key_from_password(env_value.as_str(), None)?;
return Ok(key);
}
Err(InklogError::ConfigError(format!(
"Encryption key must be exactly 32 bytes (256 bits) for raw keys, or a password string (1-127 chars) for key derivation. Got {} bytes. \
Please provide a valid 32-byte key in raw or Base64 format, or use a password string.",
raw_bytes.len()
)))
}
pub fn derive_key_from_password(
password: &str,
salt: Option<&[u8]>,
) -> Result<([u8; 32], Vec<u8>), InklogError> {
let mut key = [0u8; 32];
let salt: Vec<u8> = match salt {
Some(s) => s.to_vec(),
None => {
let mut salt_bytes = vec![0u8; 16];
rand::rng().fill_bytes(&mut salt_bytes);
salt_bytes
}
};
pbkdf2_hmac::<Sha256>(
password.as_bytes(),
&salt,
100_000, &mut key,
);
Ok((key, salt))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_get_encryption_key_from_base64() {
let key_b64 = "MDEyMzQ1Njc4OTAxMjM0NTY3ODkwMTIzNDU2Nzg5MDE=";
unsafe {
std::env::set_var("INKLOG_TEST_KEY", key_b64);
}
let result = get_encryption_key("INKLOG_TEST_KEY");
unsafe {
std::env::remove_var("INKLOG_TEST_KEY");
}
assert!(result.is_ok());
let key = result.unwrap();
assert_eq!(key.len(), 32);
}
#[test]
fn test_get_encryption_key_from_raw_bytes() {
let key_raw = "abcdefghijklmnopqrstuvwxyz123456";
unsafe {
std::env::set_var("INKLOG_TEST_KEY", key_raw);
}
let result = get_encryption_key("INKLOG_TEST_KEY");
unsafe {
std::env::remove_var("INKLOG_TEST_KEY");
}
assert!(result.is_ok());
let key = result.unwrap();
assert_eq!(key.len(), 32);
}
#[test]
fn test_get_encryption_key_missing() {
unsafe {
std::env::remove_var("INKLOG_NONEXISTENT_KEY");
}
let result = get_encryption_key("INKLOG_NONEXISTENT_KEY");
assert!(result.is_err());
}
#[test]
fn test_derive_key_from_password() {
let result = derive_key_from_password("test_password", Some(b"test_salt"));
assert!(result.is_ok());
let (key, salt) = result.unwrap();
assert_eq!(key.len(), 32);
assert_eq!(salt, b"test_salt");
}
#[test]
fn test_derive_key_deterministic() {
let result1 = derive_key_from_password("password", Some(b"salt"));
let result2 = derive_key_from_password("password", Some(b"salt"));
assert!(result1.is_ok());
assert!(result2.is_ok());
assert_eq!(result1.unwrap().0, result2.unwrap().0);
}
#[test]
fn test_derive_key_different_salts() {
let result1 = derive_key_from_password("password", Some(b"salt1"));
let result2 = derive_key_from_password("password", Some(b"salt2"));
assert!(result1.is_ok());
assert!(result2.is_ok());
assert_ne!(result1.unwrap().0, result2.unwrap().0);
}
#[test]
fn test_derive_key_different_passwords() {
let result1 = derive_key_from_password("password1", Some(b"salt"));
let result2 = derive_key_from_password("password2", Some(b"salt"));
assert!(result1.is_ok());
assert!(result2.is_ok());
assert_ne!(result1.unwrap().0, result2.unwrap().0);
}
#[test]
fn test_derive_key_with_random_salt() {
let (key1, salt1) = derive_key_from_password("test_password", None).unwrap();
assert_eq!(key1.len(), 32);
assert_eq!(salt1.len(), 16);
let (key2, salt2) = derive_key_from_password("test_password", None).unwrap();
assert_ne!(salt1, salt2); assert_ne!(key1, key2); }
#[test]
fn test_get_encryption_key_from_password() {
unsafe {
std::env::set_var("INKLOG_TEST_PWD_DERIVE", "my_password");
}
let result = get_encryption_key("INKLOG_TEST_PWD_DERIVE");
unsafe {
std::env::remove_var("INKLOG_TEST_PWD_DERIVE");
}
assert!(result.is_ok());
let key = result.unwrap();
assert_eq!(key.len(), 32);
}
#[test]
fn test_get_encryption_key_base64_wrong_length() {
use base64::{Engine as _, engine::general_purpose};
let key_16_bytes = [0u8; 16];
let b64 = general_purpose::STANDARD.encode(key_16_bytes);
unsafe {
std::env::set_var("INKLOG_TEST_B64_WRONG_LEN", &b64);
}
let result = get_encryption_key("INKLOG_TEST_B64_WRONG_LEN");
unsafe {
std::env::remove_var("INKLOG_TEST_B64_WRONG_LEN");
}
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("32 bytes") || err_msg.contains("256 bits"));
}
#[test]
fn test_get_encryption_key_too_long_input() {
let long_password = "a".repeat(128);
unsafe {
std::env::set_var("INKLOG_TEST_TOO_LONG", &long_password);
}
let result = get_encryption_key("INKLOG_TEST_TOO_LONG");
unsafe {
std::env::remove_var("INKLOG_TEST_TOO_LONG");
}
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(err_msg.contains("32 bytes") || err_msg.contains("password"));
}
#[test]
fn test_get_encryption_key_empty_string() {
unsafe {
std::env::set_var("INKLOG_TEST_EMPTY", "");
}
let result = get_encryption_key("INKLOG_TEST_EMPTY");
unsafe {
std::env::remove_var("INKLOG_TEST_EMPTY");
}
assert!(result.is_err());
}
#[test]
fn test_derive_key_with_empty_password() {
let result = derive_key_from_password("", Some(b"salt"));
assert!(result.is_ok());
let (key, _) = result.unwrap();
assert_eq!(key.len(), 32);
}
#[test]
fn test_derive_key_with_long_salt() {
let long_salt = vec![0u8; 64];
let result = derive_key_from_password("password", Some(&long_salt));
assert!(result.is_ok());
let (key, salt) = result.unwrap();
assert_eq!(key.len(), 32);
assert_eq!(salt.len(), 64);
}
#[test]
fn test_get_encryption_key_long_non_base64_input() {
let long_non_base64 = "!".repeat(128);
unsafe {
std::env::set_var("INKLOG_TEST_LONG_NON_B64", &long_non_base64);
}
let result = get_encryption_key("INKLOG_TEST_LONG_NON_B64");
unsafe {
std::env::remove_var("INKLOG_TEST_LONG_NON_B64");
}
assert!(result.is_err());
let err_msg = format!("{}", result.unwrap_err());
assert!(
err_msg.contains("32 bytes") || err_msg.contains("password"),
"error should mention 32 bytes or password, got: {}",
err_msg
);
}
}