use crate::InklogError;
use base64::{Engine as _, engine::general_purpose};
use pbkdf2::pbkdf2_hmac;
use rand::Rng;
use sha2::Sha256;
use zeroize::Zeroizing;
pub(crate) const PBKDF2_ITERATIONS: u32 = 600_000;
pub fn get_encryption_key(env_var: &str) -> Result<Zeroizing<[u8; 32]>, InklogError> {
let env_value = read_key_env_value(env_var)?;
key_from_env_value(env_var, &env_value, None)
}
pub fn get_encryption_key_with_salt(
env_var: &str,
salt: &[u8],
) -> Result<Zeroizing<[u8; 32]>, InklogError> {
let env_value = read_key_env_value(env_var)?;
key_from_env_value(env_var, &env_value, Some(salt))
}
pub fn env_key_is_password(env_var: &str) -> bool {
let Ok(value) = std::env::var(env_var) else {
return false;
};
let raw = value.as_bytes();
raw.len() != 32
&& !raw.is_empty()
&& raw.len() < 128
&& general_purpose::STANDARD.decode(value.as_str()).is_err()
}
fn read_key_env_value(env_var: &str) -> Result<Zeroizing<String>, InklogError> {
let value = std::env::var(env_var).map_err(|_| {
let mut args = fluent_bundle::FluentArgs::new();
args.set("env", env_var);
InklogError::ConfigError(crate::i18n::tr_args("config-encryption_key_not_set", args))
})?;
Ok(Zeroizing::new(value))
}
fn key_from_env_value(
env_var: &str,
env_value: &str,
salt: Option<&[u8]>,
) -> Result<Zeroizing<[u8; 32]>, InklogError> {
let raw_bytes = env_value.as_bytes();
if raw_bytes.len() == 32 {
tracing::warn!(
env = %env_var,
"32-byte input used directly as a raw encryption key; \
prefer a Base64-encoded random key or a password whose length is not 32"
);
let mut result = [0u8; 32];
result.copy_from_slice(raw_bytes);
return Ok(Zeroizing::new(result));
}
if let Ok(decoded) = general_purpose::STANDARD.decode(env_value) {
if decoded.len() == 32 {
let mut result = [0u8; 32];
result.copy_from_slice(&decoded);
return Ok(Zeroizing::new(result));
}
let mut args = fluent_bundle::FluentArgs::new();
args.set("got", decoded.len());
return Err(InklogError::ConfigError(crate::i18n::tr_args(
"config-encryption_base64_wrong_length",
args,
)));
}
if !raw_bytes.is_empty() && raw_bytes.len() < 128 {
let (key, _salt) = derive_key_from_password(env_value, salt)?;
return Ok(Zeroizing::new(key));
}
let mut args = fluent_bundle::FluentArgs::new();
args.set("got", raw_bytes.len());
Err(InklogError::ConfigError(crate::i18n::tr_args(
"config-encryption_key_wrong_length",
args,
)))
}
pub fn derive_key_from_password(
password: &str,
salt: Option<&[u8]>,
) -> Result<([u8; 32], Vec<u8>), InklogError> {
if password.len() < 12 {
let mut args = fluent_bundle::FluentArgs::new();
args.set("got", password.len());
return Err(InklogError::ConfigError(crate::i18n::tr_args(
"config-encryption_password_too_short",
args,
)));
}
if password.len() < 16 {
tracing::warn!("{}", crate::i18n::tr("warn-weak_password"));
}
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, PBKDF2_ITERATIONS, &mut key);
Ok((key, salt))
}
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
#[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("password1234", Some(b"salt"));
let result2 = derive_key_from_password("password1234", 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("password1234", Some(b"salt1"));
let result2 = derive_key_from_password("password1234", 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("password1234a", Some(b"salt"));
let result2 = derive_key_from_password("password1234b", 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_password12");
}
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_err());
let err_msg = result.unwrap_err().to_string();
assert!(err_msg.contains("at least 12 characters"));
}
#[test]
fn test_derive_key_with_short_password() {
let result = derive_key_from_password("short", Some(b"salt"));
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(err_msg.contains("at least 12 characters"));
}
#[test]
fn test_derive_key_minimum_length() {
let result = derive_key_from_password("123456789012", Some(b"salt"));
assert!(result.is_ok());
}
#[test]
fn test_derive_key_with_long_salt() {
let long_salt = vec![0u8; 64];
let result = derive_key_from_password("password1234", 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
);
}
#[test]
fn test_pbkdf2_iteration_count_is_at_least_600k() {
#[allow(clippy::assertions_on_constants)]
{
assert!(
PBKDF2_ITERATIONS >= 600_000,
"PBKDF2 iterations must stay at or above the OWASP recommendation \
(600,000), got {PBKDF2_ITERATIONS}"
);
}
let test_vector = "pbkdf2-guard-test-vector-01";
let test_salt: &[u8] = b"pbkdf2-guard-salt";
let (derived, used_salt) =
derive_key_from_password(test_vector, Some(test_salt)).expect("derive should succeed");
assert_eq!(
used_salt,
test_salt.to_vec(),
"provided salt must be used as-is"
);
let mut expected = [0u8; 32];
pbkdf2_hmac::<Sha256>(
test_vector.as_bytes(),
used_salt.as_slice(),
PBKDF2_ITERATIONS,
&mut expected,
);
assert_eq!(
derived, expected,
"derive_key_from_password must derive with PBKDF2_ITERATIONS iterations"
);
}
#[test]
fn test_minimum_password_length_is_12() {
let result = derive_key_from_password("12345678901", Some(b"salt"));
assert!(result.is_err(), "11-char password should be rejected");
let result = derive_key_from_password("123456789012", Some(b"salt"));
assert!(result.is_ok(), "12-char password should be accepted");
}
#[test]
#[serial]
fn test_get_encryption_key_with_salt_password_deterministic() {
unsafe {
std::env::set_var("INKLOG_TEST_KEY_WITH_SALT", "round-trip-password-01");
}
let salt = b"0123456789abcdef"; let key1 = get_encryption_key_with_salt("INKLOG_TEST_KEY_WITH_SALT", salt).unwrap();
let key2 = get_encryption_key_with_salt("INKLOG_TEST_KEY_WITH_SALT", salt).unwrap();
assert_eq!(
*key1, *key2,
"same password + same salt must derive the same key"
);
let key3 =
get_encryption_key_with_salt("INKLOG_TEST_KEY_WITH_SALT", b"different-salt!!").unwrap();
assert_ne!(*key1, *key3, "different salt must derive a different key");
let (expected, used_salt) =
derive_key_from_password("round-trip-password-01", Some(salt)).unwrap();
assert_eq!(used_salt, salt.to_vec());
assert_eq!(*key1, expected);
unsafe {
std::env::remove_var("INKLOG_TEST_KEY_WITH_SALT");
}
}
#[test]
#[serial]
fn test_get_encryption_key_with_salt_ignores_salt_for_base64() {
let key_bytes: [u8; 32] = core::array::from_fn(|i| (i as u8) * 7 + 3);
let key_b64 = general_purpose::STANDARD.encode(key_bytes);
unsafe {
std::env::set_var("INKLOG_TEST_KEY_WITH_SALT_B64", &key_b64);
}
let key = get_encryption_key_with_salt("INKLOG_TEST_KEY_WITH_SALT_B64", b"ignored-salt!")
.unwrap();
assert_eq!(*key, key_bytes, "Base64 branch must ignore salt");
unsafe {
std::env::remove_var("INKLOG_TEST_KEY_WITH_SALT_B64");
}
}
#[test]
#[serial]
fn test_get_encryption_key_with_salt_ignores_salt_for_raw_32() {
let raw = "abcdefghijklmnopqrstuvwxyz123456"; unsafe {
std::env::set_var("INKLOG_TEST_KEY_WITH_SALT_RAW", raw);
}
let key =
get_encryption_key_with_salt("INKLOG_TEST_KEY_WITH_SALT_RAW", b"ignored!").unwrap();
assert_eq!(&*key, raw.as_bytes());
unsafe {
std::env::remove_var("INKLOG_TEST_KEY_WITH_SALT_RAW");
}
}
#[test]
fn test_get_encryption_key_with_salt_missing_env() {
unsafe {
std::env::remove_var("INKLOG_TEST_KEY_WITH_SALT_MISSING");
}
let result = get_encryption_key_with_salt("INKLOG_TEST_KEY_WITH_SALT_MISSING", b"salt");
assert!(result.is_err());
}
#[test]
#[serial]
fn test_env_key_is_password_classification() {
unsafe {
std::env::set_var("INKLOG_TEST_CLASSIFY_PWD", "plain-password-01");
std::env::set_var(
"INKLOG_TEST_CLASSIFY_B64",
general_purpose::STANDARD.encode([0x5Au8; 32]).as_str(),
);
std::env::set_var(
"INKLOG_TEST_CLASSIFY_RAW32",
"abcdefghijklmnopqrstuvwxyz123456",
);
std::env::set_var(
"INKLOG_TEST_CLASSIFY_B64_SHORT",
general_purpose::STANDARD.encode([0u8; 16]).as_str(),
);
}
assert!(env_key_is_password("INKLOG_TEST_CLASSIFY_PWD"));
assert!(!env_key_is_password("INKLOG_TEST_CLASSIFY_B64"));
assert!(!env_key_is_password("INKLOG_TEST_CLASSIFY_RAW32"));
assert!(!env_key_is_password("INKLOG_TEST_CLASSIFY_B64_SHORT"));
assert!(!env_key_is_password("INKLOG_TEST_CLASSIFY_MISSING"));
unsafe {
for var in [
"INKLOG_TEST_CLASSIFY_PWD",
"INKLOG_TEST_CLASSIFY_B64",
"INKLOG_TEST_CLASSIFY_RAW32",
"INKLOG_TEST_CLASSIFY_B64_SHORT",
] {
std::env::remove_var(var);
}
}
}
}