use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{SystemTime, UNIX_EPOCH};
use cryptovault::{CryptoError, CryptoVault};
use rusqlite::{params, Connection, OptionalExtension};
use zeroize::Zeroizing;
use super::{MaskedDek, VaultError};
const CREATE_VAULT_TABLE_SQL: &str = "CREATE TABLE IF NOT EXISTS vault (
name TEXT PRIMARY KEY,
value_blob BLOB NOT NULL,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
)";
const SET_SQL: &str = "INSERT INTO vault (name, value_blob, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4)
ON CONFLICT(name) DO UPDATE SET value_blob = excluded.value_blob, updated_at = excluded.updated_at";
const GET_SQL: &str = "SELECT value_blob FROM vault WHERE name = ?1";
const REMOVE_SQL: &str = "DELETE FROM vault WHERE name = ?1";
const LIST_SQL: &str =
"SELECT name, datetime(created_at, 'unixepoch'), datetime(updated_at, 'unixepoch')
FROM vault ORDER BY name ASC";
const CONTAINS_SQL: &str = "SELECT EXISTS(SELECT 1 FROM vault WHERE name = ?1)";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SecretEntry {
pub name: String,
pub created_at: String,
pub updated_at: String,
}
pub trait SecretStore: Send {
fn set(&mut self, name: &str, value: &str) -> Result<(), VaultError>;
fn get(&mut self, name: &str) -> Result<Zeroizing<String>, VaultError>;
fn remove(&mut self, name: &str) -> Result<(), VaultError>;
fn list(&mut self) -> Result<Vec<SecretEntry>, VaultError>;
fn contains(&mut self, name: &str) -> Result<bool, VaultError>;
}
pub struct VaultStore {
conn: Arc<Mutex<Connection>>,
vault: CryptoVault,
dek: MaskedDek,
}
pub fn wire(conn: Arc<Mutex<Connection>>, dek: MaskedDek) -> Result<VaultStore, VaultError> {
{
let guard = conn.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
guard
.execute(CREATE_VAULT_TABLE_SQL, [])
.map_err(|e| VaultError::Storage(e.to_string()))?;
}
Ok(VaultStore {
conn,
vault: CryptoVault::default(),
dek,
})
}
fn map_store_crypto_err(e: CryptoError) -> VaultError {
VaultError::Crypto(e.to_string())
}
fn raw_column_bytes(value: rusqlite::types::ValueRef<'_>) -> Vec<u8> {
match value {
rusqlite::types::ValueRef::Blob(b) => b.to_vec(),
rusqlite::types::ValueRef::Text(t) => t.to_vec(),
_ => Vec::new(),
}
}
fn validate_name(name: &str) -> Result<(), VaultError> {
if name.trim().is_empty() {
return Err(VaultError::SecretNotFound(name.to_string()));
}
Ok(())
}
fn now_epoch_secs() -> i64 {
let secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
i64::try_from(secs).unwrap_or(i64::MAX)
}
impl VaultStore {
fn locked_conn(&self) -> MutexGuard<'_, Connection> {
self.conn
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
}
impl SecretStore for VaultStore {
fn set(&mut self, name: &str, value: &str) -> Result<(), VaultError> {
validate_name(name)?;
let vault = &self.vault;
let sealed = self
.dek
.with_dek(|k| vault.encrypt_with_key(k, value))
.map_err(map_store_crypto_err)?;
let now = now_epoch_secs();
let conn = self.locked_conn();
conn.execute(SET_SQL, params![name, sealed.into_bytes(), now, now])
.map_err(|e| VaultError::Storage(e.to_string()))?;
Ok(())
}
fn get(&mut self, name: &str) -> Result<Zeroizing<String>, VaultError> {
validate_name(name)?;
let raw: Vec<u8> = {
let conn = self.locked_conn();
conn.query_row(GET_SQL, params![name], |row| {
Ok(raw_column_bytes(row.get_ref(0)?))
})
.optional()
.map_err(|e| VaultError::Storage(e.to_string()))?
.ok_or_else(|| VaultError::SecretNotFound(name.to_string()))?
};
let blob = String::from_utf8(raw)
.map_err(|_| VaultError::Crypto("value blob is not valid UTF-8".to_string()))?;
let vault = &self.vault;
self.dek
.with_dek(|k| vault.decrypt_with_key(k, &blob))
.map_err(map_store_crypto_err)
}
fn remove(&mut self, name: &str) -> Result<(), VaultError> {
validate_name(name)?;
let conn = self.locked_conn();
let affected = conn
.execute(REMOVE_SQL, params![name])
.map_err(|e| VaultError::Storage(e.to_string()))?;
if affected == 0 {
return Err(VaultError::SecretNotFound(name.to_string()));
}
Ok(())
}
fn list(&mut self) -> Result<Vec<SecretEntry>, VaultError> {
let conn = self.locked_conn();
let mut stmt = conn
.prepare(LIST_SQL)
.map_err(|e| VaultError::Storage(e.to_string()))?;
let rows = stmt
.query_map([], |row| {
Ok(SecretEntry {
name: row.get(0)?,
created_at: row.get(1)?,
updated_at: row.get(2)?,
})
})
.map_err(|e| VaultError::Storage(e.to_string()))?;
let mut entries = Vec::new();
for row in rows {
entries.push(row.map_err(|e| VaultError::Storage(e.to_string()))?);
}
Ok(entries)
}
fn contains(&mut self, name: &str) -> Result<bool, VaultError> {
validate_name(name)?;
let conn = self.locked_conn();
conn.query_row(CONTAINS_SQL, params![name], |row| row.get(0))
.map_err(|e| VaultError::Storage(e.to_string()))
}
}
#[doc(hidden)]
pub fn fuzz_value_roundtrip_entrypoint(data: &[u8]) {
let Ok(conn) = Connection::open_in_memory() else {
return;
};
let Ok(dek) = MaskedDek::new(Zeroizing::new(vec![7u8; 32])) else {
return;
};
let Ok(mut store) = wire(Arc::new(Mutex::new(conn)), dek) else {
return;
};
let value = String::from_utf8_lossy(data);
if store.set("k", &value).is_ok() {
let _ = store.get("k");
}
}
#[cfg(test)]
impl VaultStore {
pub(crate) fn debug_conn(&self) -> Arc<Mutex<Connection>> {
self.conn.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vault::MaskedDek;
use std::sync::{Arc, Mutex};
use zeroize::Zeroizing;
fn fixture() -> VaultStore {
let conn = rusqlite::Connection::open_in_memory().expect("mem db");
let dek = MaskedDek::new(Zeroizing::new(vec![7u8; 32])).expect("32B");
crate::vault::wire(Arc::new(Mutex::new(conn)), dek).expect("wire")
}
#[test]
fn test_set_then_get_roundtrips_the_exact_value() {
let mut s = fixture();
s.set("OPENAI_API_KEY", "sk-test-123\nline2").expect("set");
let v = s.get("OPENAI_API_KEY").expect("get");
assert_eq!(v.as_str(), "sk-test-123\nline2"); }
#[test]
fn test_set_overwrites_existing_and_bumps_updated_at_only() {
let mut s = fixture();
s.set("K", "v1").expect("set1");
let e1 = &s.list().expect("ls")[0];
let (c1, _) = (e1.created_at.clone(), e1.updated_at.clone());
std::thread::sleep(std::time::Duration::from_millis(1100)); s.set("K", "v2").expect("set2");
let e2 = &s.list().expect("ls")[0];
assert_eq!(e2.created_at, c1); assert_ne!(e2.updated_at, e2.created_at); assert_eq!(s.get("K").expect("get").as_str(), "v2");
}
#[test]
fn test_remove_deletes_and_missing_names_are_typed_errors() {
let mut s = fixture();
s.set("K", "v").expect("set");
s.remove("K").expect("rm");
assert!(matches!(s.get("K"), Err(VaultError::SecretNotFound(_))));
assert!(matches!(s.remove("K"), Err(VaultError::SecretNotFound(_))));
}
#[test]
fn test_empty_name_is_rejected() {
let mut s = fixture();
assert!(matches!(s.set("", "v"), Err(VaultError::SecretNotFound(_))));
assert!(matches!(
s.set(" ", "v"),
Err(VaultError::SecretNotFound(_))
));
assert!(matches!(s.get(""), Err(VaultError::SecretNotFound(_))));
assert!(matches!(s.remove(""), Err(VaultError::SecretNotFound(_))));
}
#[test]
fn test_contains_reports_presence_and_rejects_empty_name() {
let mut s = fixture();
assert!(!s.contains("K").expect("absent")); s.set("K", "v").expect("set");
assert!(s.contains("K").expect("present")); assert!(matches!(s.contains(""), Err(VaultError::SecretNotFound(_)))); }
#[test]
fn test_list_yields_names_and_dates_never_values() {
let mut s = fixture();
s.set("A", "secret-value-A").expect("set");
let rows = s.list().expect("ls");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].name, "A");
assert!(!format!(
"{:?} {} {}",
rows[0].name, rows[0].created_at, rows[0].updated_at
)
.contains("secret-value")); }
#[test]
fn test_tampered_blob_fails_typed_never_returns_wrong_value() {
let mut s = fixture();
s.set("K", "v").expect("set");
{
let conn = s.debug_conn();
let guard = conn.lock().expect("lock");
let original: Vec<u8> = guard
.query_row("SELECT value_blob FROM vault WHERE name='K'", [], |r| {
Ok(raw_column_bytes(r.get_ref(0)?))
})
.expect("read original blob");
let tampered: Vec<u8> = original.iter().map(|b| b ^ 0xFF).collect();
guard
.execute(
"UPDATE vault SET value_blob = ?1 WHERE name = 'K'",
rusqlite::params![tampered],
)
.expect("tamper");
}
assert!(matches!(s.get("K"), Err(VaultError::Crypto(_))));
}
#[test]
fn test_vault_error_messages_never_contain_secret_values() {
let probe = "sk-super-secret-PROBE-9f3a";
let mut s = fixture();
let errs: Vec<String> = vec![
VaultError::SecretNotFound("A".into()).to_string(),
VaultError::Crypto("tag mismatch".into()).to_string(),
VaultError::WrongPassphrase.to_string(),
VaultError::VaultMetaCorrupt.to_string(),
];
s.set("A", probe).expect("set");
for e in errs {
assert!(!e.contains(probe), "VaultError leaked a value: {e}");
}
}
#[test]
fn test_value_is_not_plaintext_at_rest() {
let mut s = fixture();
s.set("K", "super-secret-readable").expect("set");
let blob: Vec<u8> = {
let conn = s.debug_conn();
let guard = conn.lock().expect("lock");
guard
.query_row("SELECT value_blob FROM vault WHERE name='K'", [], |r| {
Ok(raw_column_bytes(r.get_ref(0)?))
})
.expect("row")
};
assert_ne!(blob.as_slice(), b"super-secret-readable");
assert!(!String::from_utf8_lossy(&blob).contains("super-secret-readable"));
assert_eq!(
s.get("K").expect("roundtrip").as_str(),
"super-secret-readable"
);
}
#[test]
fn test_fuzz_value_roundtrip_entrypoint_never_panics_on_arbitrary_input() {
for data in [
&b""[..],
&b"\x00"[..],
&b"\xff\xfe\x80"[..], &[0x41u8; 100_000], ] {
super::fuzz_value_roundtrip_entrypoint(data);
}
}
}