#![allow(dead_code)]
use std::sync::OnceLock;
use aes_gcm::aead::{Aead, Nonce, Payload};
use aes_gcm::{Aes256Gcm, KeyInit};
use super::{CodecError, FieldCodec, MissingKeyKind};
use crate::migrate::OnlineSafetyClassification;
#[doc(hidden)]
pub const ENV_VAR: &str = "DJOGI_FIELD_CODEC_KEY_0";
static CODEC_RING: OnceLock<Vec<[u8; 32]>> = OnceLock::new();
#[cfg(all(test, feature = "aes-codec"))]
pub(crate) static TEST_CODEC_ENV_MUTEX: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[doc(hidden)]
pub fn load_ring() -> Result<(), CodecError> {
if CODEC_RING.get().is_some() {
return Ok(());
}
let mut ring: Vec<[u8; 32]> = Vec::new();
for index in 0u8..=31 {
let name = format!("DJOGI_FIELD_CODEC_KEY_{index}");
match std::env::var(&name) {
Ok(raw) => {
if ring.len() != index as usize {
return Err(CodecError::MissingKey {
index: ring.len() as u8,
kind: MissingKeyKind::Gap,
});
}
let is_lower_hex = raw.len() == 64
&& raw
.bytes()
.all(|b| b.is_ascii_digit() || (b'a'..=b'f').contains(&b));
if !is_lower_hex {
return Err(CodecError::MissingKey {
index,
kind: MissingKeyKind::Malformed,
});
}
let bytes = hex::decode(&raw).map_err(|_| CodecError::MissingKey {
index,
kind: MissingKeyKind::Malformed,
})?;
let arr: [u8; 32] = bytes.try_into().map_err(|_| CodecError::MissingKey {
index,
kind: MissingKeyKind::Malformed,
})?;
ring.push(arr);
}
Err(_) => continue,
}
}
if ring.is_empty() {
return Err(CodecError::RingEmpty);
}
let _ = CODEC_RING.set(ring);
Ok(())
}
fn derive_subkey(ikm: &[u8; 32], model: &str, field: &str) -> [u8; 32] {
let hk = hkdf::Hkdf::<sha2::Sha256>::new(None, ikm); let mut info = Vec::with_capacity(20 + model.len() + 1 + field.len());
info.extend_from_slice(b"djogi:aes256_gcm_v1\x00");
info.extend_from_slice(model.as_bytes());
info.push(0);
info.extend_from_slice(field.as_bytes());
let mut okm = [0u8; 32];
hk.expand(&info, &mut okm)
.expect("32 bytes is a valid HKDF-SHA256 output length");
okm
}
fn generate_nonce() -> Result<[u8; 12], CodecError> {
let mut nonce = [0u8; 12];
getrandom::fill(&mut nonce)?; Ok(nonce)
}
pub struct Aes256GcmV1;
impl FieldCodec for Aes256GcmV1 {
const ID: &'static str = "aes256_gcm_v1";
type Decoded = String;
type Encoded = Vec<u8>;
type Error = CodecError;
fn encode(
model: &'static str,
field: &'static str,
value: &String,
) -> Result<Vec<u8>, CodecError> {
load_ring()?; let ring = CODEC_RING.get().ok_or(CodecError::RingEmpty)?;
let active_index = ring.len() - 1; let subkey = derive_subkey(&ring[active_index], model, field);
let nonce_bytes = generate_nonce()?;
let nonce = Nonce::<Aes256Gcm>::from_slice(&nonce_bytes);
let aad = format!("{model}\x00{field}");
let aead =
Aes256Gcm::new_from_slice(&subkey).expect("32-byte slice is a valid AES-256 key");
let ciphertext = aead.encrypt(
nonce,
Payload {
msg: value.as_bytes(),
aad: aad.as_bytes(),
},
)?; let mut out = Vec::with_capacity(2 + 12 + ciphertext.len());
out.push(0x01); out.push(active_index as u8); out.extend_from_slice(&nonce_bytes);
out.extend_from_slice(&ciphertext);
Ok(out)
}
fn decode(
model: &'static str,
field: &'static str,
stored: &Vec<u8>,
) -> Result<String, CodecError> {
if stored.len() < 30 {
return Err(CodecError::CiphertextTooShort);
}
let version = stored[0];
if version != 0x01 {
return Err(CodecError::UnknownVersion(version));
}
load_ring()?; let ring = CODEC_RING.get().ok_or(CodecError::RingEmpty)?;
let key_index = stored[1] as usize;
if key_index >= ring.len() {
return Err(CodecError::UnknownKeyIndex {
index: stored[1],
ring_len: ring.len() as u8,
});
}
let nonce = Nonce::<Aes256Gcm>::from_slice(&stored[2..14]);
let ciphertext = &stored[14..];
let subkey = derive_subkey(&ring[key_index], model, field);
let aad = format!("{model}\x00{field}");
let aead =
Aes256Gcm::new_from_slice(&subkey).expect("32-byte slice is a valid AES-256 key");
let plaintext = aead.decrypt(
nonce,
Payload {
msg: ciphertext,
aad: aad.as_bytes(),
},
)?; Ok(String::from_utf8(plaintext)?) }
fn classify_transition<Other: FieldCodec>() -> OnlineSafetyClassification {
if Other::ID == Self::ID {
OnlineSafetyClassification::OnlineSafe
} else {
OnlineSafetyClassification::ExpandContract
}
}
}
#[cfg(all(test, feature = "aes-codec"))]
pub(crate) fn test_with_codec_ring(keys_hex: &[&str], f: impl FnOnce()) {
let _lock = TEST_CODEC_ENV_MUTEX.lock().unwrap();
let prev: Vec<Option<String>> = (0..32)
.map(|i| std::env::var(format!("DJOGI_FIELD_CODEC_KEY_{i}")).ok())
.collect();
for i in 0..32 {
let name = format!("DJOGI_FIELD_CODEC_KEY_{i}");
match keys_hex.get(i) {
Some(k) => unsafe { std::env::set_var(&name, k) },
None => unsafe { std::env::remove_var(&name) },
}
}
let _ = load_ring();
f();
for (i, p) in prev.into_iter().enumerate() {
let name = format!("DJOGI_FIELD_CODEC_KEY_{i}");
match p {
Some(v) => unsafe { std::env::set_var(&name, &v) },
None => unsafe { std::env::remove_var(&name) },
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::error::Error;
const TEST_KEY_HEX: &str = "0000000000000000000000000000000000000000000000000000000000000001";
const TEST_KEY_HEX_B: &str = "0000000000000000000000000000000000000000000000000000000000000002";
#[test]
fn encode_decode_round_trip() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let original = String::from("djogi encrypted-at-rest round-trip test");
let encoded = Aes256GcmV1::encode("TestModel", "secret_field", &original)
.expect("encode should succeed with valid ring");
assert_eq!(encoded.len(), original.len() + 30);
let decoded = Aes256GcmV1::decode("TestModel", "secret_field", &encoded)
.expect("decode should succeed with matching ring and AAD");
assert_eq!(decoded, original);
});
}
#[test]
fn encode_produces_different_ciphertext_each_time() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let original = String::from("same plaintext");
let encoded1 = Aes256GcmV1::encode("TestModel", "field", &original).unwrap();
let encoded2 = Aes256GcmV1::encode("TestModel", "field", &original).unwrap();
assert_ne!(
encoded1, encoded2,
"two encodes of same plaintext should differ (random nonce)"
);
assert_eq!(
Aes256GcmV1::decode("TestModel", "field", &encoded1).unwrap(),
original
);
assert_eq!(
Aes256GcmV1::decode("TestModel", "field", &encoded2).unwrap(),
original
);
});
}
#[test]
fn aad_mismatch_returns_aead_error() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let original = String::from("aad binding test");
let encoded = Aes256GcmV1::encode("ModelA", "field_a", &original).unwrap();
let err = Aes256GcmV1::decode("ModelB", "field_a", &encoded).unwrap_err();
assert!(matches!(err, CodecError::AeadError(_)));
let err = Aes256GcmV1::decode("ModelA", "field_b", &encoded).unwrap_err();
assert!(matches!(err, CodecError::AeadError(_)));
assert_eq!(
Aes256GcmV1::decode("ModelA", "field_a", &encoded).unwrap(),
original
);
});
}
#[test]
fn ciphertext_too_short_returns_error() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let short = vec![0u8; 10]; let err = Aes256GcmV1::decode("Model", "field", &short).unwrap_err();
assert!(matches!(err, CodecError::CiphertextTooShort));
let msg = err.to_string();
assert!(
msg.contains("30"),
"message should mention the 30-byte minimum: {msg}"
);
});
}
#[test]
fn unknown_version_decode() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let mut blob = vec![0u8; 30];
blob[0] = 0x02;
let err = Aes256GcmV1::decode("Model", "field", &blob).unwrap_err();
assert!(
matches!(err, CodecError::UnknownVersion(0x02)),
"got: {err:?}"
);
});
}
#[test]
fn unknown_key_index_decode() {
test_with_codec_ring(&[TEST_KEY_HEX, TEST_KEY_HEX_B], || {
let original = String::from("rotate me");
let mut encoded = Aes256GcmV1::encode("Model", "field", &original).unwrap();
encoded[1] = 5; let err = Aes256GcmV1::decode("Model", "field", &encoded).unwrap_err();
match err {
CodecError::UnknownKeyIndex { index, ring_len } => {
assert_eq!(index, 5);
assert_eq!(
err.to_string(),
format!("key index 5 not in ring of length {ring_len}")
);
}
other => panic!("expected UnknownKeyIndex, got: {other:?}"),
}
});
}
#[test]
fn empty_string_round_trip() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let original = String::new();
let encoded = Aes256GcmV1::encode("Model", "field", &original).unwrap();
assert_eq!(encoded.len(), 30);
let decoded = Aes256GcmV1::decode("Model", "field", &encoded).unwrap();
assert_eq!(decoded, original);
});
}
#[test]
fn hkdf_domain_separation() {
test_with_codec_ring(&[TEST_KEY_HEX], || {
let original = String::from("domain separation");
let blob_a = Aes256GcmV1::encode("ModelA", "f", &original).unwrap();
let err = Aes256GcmV1::decode("ModelB", "f", &blob_a).unwrap_err();
assert!(matches!(err, CodecError::AeadError(_)), "got: {err:?}");
});
}
#[test]
fn aead_error_display_safety() {
let aead_err = CodecError::AeadError(aes_gcm::aead::Error);
let msg = aead_err.to_string();
assert_eq!(msg, "generic AEAD error");
assert!(!msg.contains("0x"));
assert!(!msg.contains("key"));
assert!(!msg.contains("nonce"));
}
#[test]
fn codec_error_source_chaining() {
let utf8_err = String::from_utf8(vec![0xff]).unwrap_err();
assert!(CodecError::Utf8Error(utf8_err).source().is_some());
assert!(
CodecError::AeadError(aes_gcm::aead::Error)
.source()
.is_none()
);
assert!(
CodecError::MissingKey {
index: 0,
kind: MissingKeyKind::Gap
}
.source()
.is_none()
);
assert!(CodecError::RingEmpty.source().is_none());
assert!(CodecError::CiphertextTooShort.source().is_none());
assert!(CodecError::UnknownVersion(2).source().is_none());
assert!(
CodecError::UnknownKeyIndex {
index: 5,
ring_len: 2
}
.source()
.is_none()
);
}
#[test]
fn unknown_key_index_display_exact() {
let err = CodecError::UnknownKeyIndex {
index: 5,
ring_len: 2,
};
assert_eq!(err.to_string(), "key index 5 not in ring of length 2");
}
#[test]
fn classify_transition_same_codec_is_online_safe() {
let classification = <Aes256GcmV1 as FieldCodec>::classify_transition::<Aes256GcmV1>();
assert_eq!(classification, OnlineSafetyClassification::OnlineSafe);
}
#[test]
fn classify_transition_different_codec_is_expand_contract() {
struct DummyCodec;
impl FieldCodec for DummyCodec {
const ID: &'static str = "_djogi_test_dummy_codec";
type Decoded = String;
type Encoded = Vec<u8>;
type Error = CodecError;
fn encode(
_model: &'static str,
_field: &'static str,
value: &Self::Decoded,
) -> Result<Self::Encoded, Self::Error> {
Ok(value.as_bytes().to_vec())
}
fn decode(
_model: &'static str,
_field: &'static str,
stored: &Self::Encoded,
) -> Result<Self::Decoded, Self::Error> {
Ok(String::from_utf8(stored.clone())?)
}
fn classify_transition<Other: FieldCodec>() -> OnlineSafetyClassification {
if Self::ID == Other::ID {
OnlineSafetyClassification::OnlineSafe
} else {
OnlineSafetyClassification::ExpandContract
}
}
}
let classification = <Aes256GcmV1 as FieldCodec>::classify_transition::<DummyCodec>();
assert_eq!(classification, OnlineSafetyClassification::ExpandContract);
}
#[test]
fn codec_startup_error_display_format() {
use super::super::CodecStartupError;
let err = CodecStartupError {
codec_id: "aes256_gcm_v1",
env_var: ENV_VAR,
error: "no field codec key configured".to_owned(),
};
let msg = err.to_string();
assert!(msg.contains("aes256_gcm_v1"));
assert!(msg.contains(ENV_VAR));
assert!(msg.contains("no field codec key configured"));
}
}
#[cfg(all(test, feature = "aes-codec"))]
mod ring_isolation_tests {
use super::*;
const KEY_0: &str = "0000000000000000000000000000000000000000000000000000000000000001";
const KEY_1: &str = "0000000000000000000000000000000000000000000000000000000000000002";
#[serial_test::serial]
#[test]
fn ring_loading_one_key_active_index_zero() {
test_with_codec_ring(&[KEY_0], || {
if CODEC_RING.get().map(|r| r.len()) == Some(1) {
let blob = Aes256GcmV1::encode("M", "f", &"x".to_string()).unwrap();
assert_eq!(blob[1], 0, "single-key ring active index must be 0");
}
});
}
#[serial_test::serial]
#[test]
fn ring_loading_two_keys_active_index_one() {
test_with_codec_ring(&[KEY_0, KEY_1], || {
if CODEC_RING.get().map(|r| r.len()) == Some(2) {
let blob = Aes256GcmV1::encode("M", "f", &"x".to_string()).unwrap();
assert_eq!(blob[1], 1, "two-key ring active index must be 1");
}
});
}
#[serial_test::serial]
#[test]
fn cross_ring_entry_non_interference() {
test_with_codec_ring(&[KEY_0, KEY_1], || {
if CODEC_RING.get().map(|r| r.len()) == Some(2) {
let mut blob = Aes256GcmV1::encode("M", "f", &"x".to_string()).unwrap();
assert_eq!(blob[1], 1);
blob[1] = 0; let err = Aes256GcmV1::decode("M", "f", &blob).unwrap_err();
assert!(
matches!(err, CodecError::AeadError(_)),
"forging the key_index must fail authentication: {err:?}"
);
}
});
}
#[serial_test::serial]
#[test]
fn key_rotation_round_trip() {
test_with_codec_ring(&[KEY_0], || {
let len = CODEC_RING.get().map(|r| r.len());
if len == Some(1) {
let blob = Aes256GcmV1::encode("M", "f", &"v".to_string()).unwrap();
assert_eq!(blob[1], 0, "first ring entry encodes under index 0");
assert_eq!(Aes256GcmV1::decode("M", "f", &blob).unwrap(), "v");
}
});
}
#[serial_test::serial]
#[test]
fn ring_immutability() {
test_with_codec_ring(&[KEY_0], || {
if CODEC_RING.get().is_some() {
unsafe { std::env::set_var("DJOGI_FIELD_CODEC_KEY_0", "not-hex-garbage") };
let blob = Aes256GcmV1::encode("M", "f", &"v".to_string()).unwrap();
assert_eq!(Aes256GcmV1::decode("M", "f", &blob).unwrap(), "v");
}
});
}
#[serial_test::serial]
#[test]
fn missing_ring_startup_failure() {
let _lock = TEST_CODEC_ENV_MUTEX.lock().unwrap();
let prev: Vec<Option<String>> = (0..32)
.map(|i| std::env::var(format!("DJOGI_FIELD_CODEC_KEY_{i}")).ok())
.collect();
for i in 0..32 {
unsafe { std::env::remove_var(format!("DJOGI_FIELD_CODEC_KEY_{i}")) };
}
if CODEC_RING.get().is_none() {
let err = load_ring().unwrap_err();
assert!(matches!(err, CodecError::RingEmpty), "got: {err:?}");
}
for (i, p) in prev.into_iter().enumerate() {
let name = format!("DJOGI_FIELD_CODEC_KEY_{i}");
match p {
Some(v) => unsafe { std::env::set_var(&name, &v) },
None => unsafe { std::env::remove_var(&name) },
}
}
}
#[serial_test::serial]
#[test]
fn ring_gap_detection() {
let _lock = TEST_CODEC_ENV_MUTEX.lock().unwrap();
let prev: Vec<Option<String>> = (0..32)
.map(|i| std::env::var(format!("DJOGI_FIELD_CODEC_KEY_{i}")).ok())
.collect();
for i in 0..32 {
unsafe { std::env::remove_var(format!("DJOGI_FIELD_CODEC_KEY_{i}")) };
}
unsafe {
std::env::set_var("DJOGI_FIELD_CODEC_KEY_0", KEY_0);
std::env::set_var("DJOGI_FIELD_CODEC_KEY_2", KEY_1);
}
if CODEC_RING.get().is_none() {
let err = load_ring().unwrap_err();
match err {
CodecError::MissingKey {
index,
kind: MissingKeyKind::Gap,
} => {
assert_eq!(index, 1);
assert!(err.to_string().contains("DJOGI_FIELD_CODEC_KEY_1"));
}
other => panic!("expected MissingKey gap at index 1, got: {other:?}"),
}
}
for (i, p) in prev.into_iter().enumerate() {
let name = format!("DJOGI_FIELD_CODEC_KEY_{i}");
match p {
Some(v) => unsafe { std::env::set_var(&name, &v) },
None => unsafe { std::env::remove_var(&name) },
}
}
}
#[serial_test::serial]
#[test]
fn malformed_entry_detection() {
let _lock = TEST_CODEC_ENV_MUTEX.lock().unwrap();
let prev: Vec<Option<String>> = (0..32)
.map(|i| std::env::var(format!("DJOGI_FIELD_CODEC_KEY_{i}")).ok())
.collect();
for i in 0..32 {
unsafe { std::env::remove_var(format!("DJOGI_FIELD_CODEC_KEY_{i}")) };
}
unsafe { std::env::set_var("DJOGI_FIELD_CODEC_KEY_0", "not-64-lowercase-hex") };
if CODEC_RING.get().is_none() {
let err = load_ring().unwrap_err();
match err {
CodecError::MissingKey {
index,
kind: MissingKeyKind::Malformed,
} => {
assert_eq!(index, 0);
assert!(err.to_string().contains("DJOGI_FIELD_CODEC_KEY_0"));
}
other => panic!("expected MissingKey malformed at index 0, got: {other:?}"),
}
}
for (i, p) in prev.into_iter().enumerate() {
let name = format!("DJOGI_FIELD_CODEC_KEY_{i}");
match p {
Some(v) => unsafe { std::env::set_var(&name, &v) },
None => unsafe { std::env::remove_var(&name) },
}
}
}
#[serial_test::serial]
#[test]
fn base_key_only_missing_distinction() {
let _lock = TEST_CODEC_ENV_MUTEX.lock().unwrap();
let prev: Vec<Option<String>> = (0..32)
.map(|i| std::env::var(format!("DJOGI_FIELD_CODEC_KEY_{i}")).ok())
.collect();
for i in 0..32 {
unsafe { std::env::remove_var(format!("DJOGI_FIELD_CODEC_KEY_{i}")) };
}
unsafe { std::env::set_var("DJOGI_FIELD_CODEC_KEY_1", KEY_1) };
if CODEC_RING.get().is_none() {
let err = load_ring().unwrap_err();
match err {
CodecError::MissingKey {
index: 0,
kind: MissingKeyKind::Gap,
} => {
assert!(err.to_string().contains("DJOGI_FIELD_CODEC_KEY_0"));
}
other => panic!("expected MissingKey gap at index 0, got: {other:?}"),
}
}
for (i, p) in prev.into_iter().enumerate() {
let name = format!("DJOGI_FIELD_CODEC_KEY_{i}");
match p {
Some(v) => unsafe { std::env::set_var(&name, &v) },
None => unsafe { std::env::remove_var(&name) },
}
}
}
}