use std::sync::OnceLock;
use base64::Engine as _;
use crypto_box::{
SalsaBox,
aead::{Aead, AeadCore},
};
use rand_core::OsRng;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
const REDACTED: &str = "••••••";
const B64: base64::engine::general_purpose::GeneralPurpose =
base64::engine::general_purpose::STANDARD;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MaskError {
NoKeyring,
NoPrivateKey,
BadKey(String),
Malformed,
Decrypt,
}
impl std::fmt::Display for MaskError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MaskError::NoKeyring => f.write_str(
"no mask keyring configured (set UMBRAL_MASK_PUBLIC_KEY or call set_mask_keyring)",
),
MaskError::NoPrivateKey => f.write_str(
"mask keyring has no private key; cannot reveal (set UMBRAL_MASK_PRIVATE_KEY)",
),
MaskError::BadKey(why) => write!(f, "invalid mask key: {why}"),
MaskError::Malformed => f.write_str("masked ciphertext is malformed"),
MaskError::Decrypt => f.write_str("masked ciphertext failed to decrypt"),
}
}
}
impl std::error::Error for MaskError {}
#[derive(Clone)]
pub struct MaskKeyring {
public: crypto_box::PublicKey,
secret: Option<crypto_box::SecretKey>,
}
impl MaskKeyring {
pub fn from_base64(public_b64: &str, secret_b64: Option<&str>) -> Result<Self, MaskError> {
let public = decode_key(public_b64)?;
let public = crypto_box::PublicKey::from(public);
let secret = match secret_b64 {
Some(s) if !s.is_empty() => Some(crypto_box::SecretKey::from(decode_key(s)?)),
_ => None,
};
Ok(Self { public, secret })
}
pub fn from_env() -> Result<Self, MaskError> {
let public = std::env::var("UMBRAL_MASK_PUBLIC_KEY").map_err(|_| MaskError::NoKeyring)?;
let secret = std::env::var("UMBRAL_MASK_PRIVATE_KEY").ok();
Self::from_base64(&public, secret.as_deref())
}
pub fn generate() -> (String, String) {
let secret = crypto_box::SecretKey::generate(&mut OsRng);
let public = secret.public_key();
(B64.encode(public.as_bytes()), B64.encode(secret.to_bytes()))
}
pub fn seal(&self, plaintext: &[u8]) -> String {
let eph_secret = crypto_box::SecretKey::generate(&mut OsRng);
let eph_public = eph_secret.public_key();
let salsa = SalsaBox::new(&self.public, &eph_secret);
let nonce = SalsaBox::generate_nonce(&mut OsRng);
let ciphertext = salsa
.encrypt(&nonce, plaintext)
.expect("XSalsa20-Poly1305 encryption is infallible for in-memory plaintext");
let mut out = Vec::with_capacity(32 + nonce.len() + ciphertext.len());
out.extend_from_slice(eph_public.as_bytes());
out.extend_from_slice(nonce.as_slice());
out.extend_from_slice(&ciphertext);
B64.encode(out)
}
pub fn open(&self, b64_ciphertext: &str) -> Result<String, MaskError> {
let secret = self.secret.as_ref().ok_or(MaskError::NoPrivateKey)?;
let sealed = B64
.decode(b64_ciphertext)
.map_err(|_| MaskError::Malformed)?;
if sealed.len() < 32 + 24 {
return Err(MaskError::Malformed);
}
let eph_public: [u8; 32] = sealed[..32].try_into().map_err(|_| MaskError::Malformed)?;
let eph_public = crypto_box::PublicKey::from(eph_public);
let nonce = crypto_box::Nonce::from_slice(&sealed[32..56]);
let ciphertext = &sealed[56..];
let salsa = SalsaBox::new(&eph_public, secret);
let plaintext = salsa
.decrypt(nonce, ciphertext)
.map_err(|_| MaskError::Decrypt)?;
String::from_utf8(plaintext).map_err(|_| MaskError::Decrypt)
}
pub fn can_reveal(&self) -> bool {
self.secret.is_some()
}
}
fn decode_key(b64: &str) -> Result<[u8; 32], MaskError> {
let bytes = B64
.decode(b64.trim())
.map_err(|e| MaskError::BadKey(e.to_string()))?;
bytes
.try_into()
.map_err(|_| MaskError::BadKey("key is not 32 bytes".to_string()))
}
static KEYRING: OnceLock<Result<Option<MaskKeyring>, MaskError>> = OnceLock::new();
pub fn set_mask_keyring(keyring: MaskKeyring) -> bool {
KEYRING.set(Ok(Some(keyring))).is_ok()
}
fn keyring() -> Result<Option<&'static MaskKeyring>, &'static MaskError> {
KEYRING
.get_or_init(|| match MaskKeyring::from_env() {
Ok(k) => Ok(Some(k)),
Err(MaskError::NoKeyring) => Ok(None),
Err(e) => {
tracing::error!(
"UMBRAL_MASK_PUBLIC_KEY/UMBRAL_MASK_PRIVATE_KEY is set but could not be \
parsed ({e}); all Masked<T> seal/reveal calls will fail with BadKey. \
Fix the key or unset the variable."
);
Err(e)
}
})
.as_ref()
.map(|opt| opt.as_ref())
.map_err(|e| e)
}
pub(crate) fn ambient_seal(plaintext: &str) -> Result<String, MaskError> {
match keyring() {
Ok(Some(k)) => Ok(k.seal(plaintext.as_bytes())),
Ok(None) => Err(MaskError::NoKeyring),
Err(e) => Err(e.clone()),
}
}
fn ambient_open(ciphertext: &str) -> Result<String, MaskError> {
match keyring() {
Ok(Some(k)) => k.open(ciphertext),
Ok(None) => Err(MaskError::NoKeyring),
Err(e) => Err(e.clone()),
}
}
#[derive(Clone)]
pub struct Masked<T = String> {
inner: MaskInner,
_marker: std::marker::PhantomData<T>,
}
#[derive(Clone)]
enum MaskInner {
Plain(String),
Sealed(String),
}
impl<T> Masked<T> {
pub fn new(plaintext: impl Into<String>) -> Self {
Self {
inner: MaskInner::Plain(plaintext.into()),
_marker: std::marker::PhantomData,
}
}
pub fn reveal(&self) -> Result<String, MaskError> {
match &self.inner {
MaskInner::Plain(p) => Ok(p.clone()),
MaskInner::Sealed(c) => ambient_open(c),
}
}
pub fn is_revealable(&self) -> bool {
match &self.inner {
MaskInner::Plain(_) => true,
MaskInner::Sealed(_) => keyring()
.ok()
.and_then(|opt| opt)
.map(MaskKeyring::can_reveal)
.unwrap_or(false),
}
}
fn to_stored(&self) -> Result<String, MaskError> {
match &self.inner {
MaskInner::Plain(p) => ambient_seal(p),
MaskInner::Sealed(c) => Ok(c.clone()),
}
}
}
impl<T> Default for Masked<T> {
fn default() -> Self {
Masked::new(String::new())
}
}
impl<T> std::fmt::Debug for Masked<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("Masked(••••••)")
}
}
impl<T> std::fmt::Display for Masked<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(REDACTED)
}
}
impl<T> From<String> for Masked<T> {
fn from(plaintext: String) -> Self {
Masked::new(plaintext)
}
}
impl<T> From<&str> for Masked<T> {
fn from(plaintext: &str) -> Self {
Masked::new(plaintext)
}
}
impl<T> Serialize for Masked<T> {
fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
let stored = self.to_stored().map_err(serde::ser::Error::custom)?;
s.serialize_str(&stored)
}
}
impl<'de, T> Deserialize<'de> for Masked<T> {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
let s = String::deserialize(d)?;
if s == REDACTED {
Ok(Masked::new(String::new()))
} else {
Ok(Masked::new(s))
}
}
}
macro_rules! impl_masked_sqlx {
($db:ty, $valueref:ty, $argbuf:ty) => {
impl<T> sqlx::Type<$db> for Masked<T> {
fn type_info() -> <$db as sqlx::Database>::TypeInfo {
<String as sqlx::Type<$db>>::type_info()
}
fn compatible(ty: &<$db as sqlx::Database>::TypeInfo) -> bool {
<String as sqlx::Type<$db>>::compatible(ty)
}
}
impl<'r, T> sqlx::Decode<'r, $db> for Masked<T> {
fn decode(value: $valueref) -> Result<Self, Box<dyn std::error::Error + Send + Sync>> {
let ciphertext = <String as sqlx::Decode<$db>>::decode(value)?;
Ok(Masked {
inner: MaskInner::Sealed(ciphertext),
_marker: std::marker::PhantomData,
})
}
}
impl<'q, T> sqlx::Encode<'q, $db> for Masked<T> {
fn encode_by_ref(
&self,
buf: &mut $argbuf,
) -> Result<sqlx::encode::IsNull, Box<dyn std::error::Error + Send + Sync>> {
let stored = self.to_stored()?;
<String as sqlx::Encode<'q, $db>>::encode_by_ref(&stored, buf)
}
}
};
}
impl_masked_sqlx!(
sqlx::Sqlite,
sqlx::sqlite::SqliteValueRef<'r>,
<sqlx::Sqlite as sqlx::Database>::ArgumentBuffer<'q>
);
impl_masked_sqlx!(
sqlx::Postgres,
sqlx::postgres::PgValueRef<'r>,
<sqlx::Postgres as sqlx::Database>::ArgumentBuffer<'q>
);
#[cfg(test)]
mod tests {
use super::*;
fn test_keyring() -> MaskKeyring {
let (public, secret) = MaskKeyring::generate();
MaskKeyring::from_base64(&public, Some(&secret)).unwrap()
}
#[test]
fn seal_open_round_trips() {
let kr = test_keyring();
let sealed = kr.seal(b"+254712345678");
assert_ne!(sealed, "+254712345678", "stored form is not plaintext");
assert_eq!(kr.open(&sealed).unwrap(), "+254712345678");
}
#[test]
fn each_seal_is_distinct_ciphertext() {
let kr = test_keyring();
let a = kr.seal(b"secret");
let b = kr.seal(b"secret");
assert_ne!(a, b, "ephemeral keypair makes ciphertext non-deterministic");
assert_eq!(kr.open(&a).unwrap(), "secret");
assert_eq!(kr.open(&b).unwrap(), "secret");
}
#[test]
fn public_key_only_cannot_open() {
let (public, secret) = MaskKeyring::generate();
let write_only = MaskKeyring::from_base64(&public, None).unwrap();
let sealed = write_only.seal(b"pii");
assert_eq!(write_only.open(&sealed), Err(MaskError::NoPrivateKey));
let full = MaskKeyring::from_base64(&public, Some(&secret)).unwrap();
assert_eq!(full.open(&sealed).unwrap(), "pii");
}
#[test]
fn wrong_key_fails_to_decrypt() {
let a = test_keyring();
let b = test_keyring();
let sealed = a.seal(b"private");
assert_eq!(b.open(&sealed), Err(MaskError::Decrypt));
}
#[test]
fn masked_redacts_in_debug_and_display() {
let m: Masked = Masked::new("0712-secret");
assert_eq!(m.to_string(), REDACTED, "Display is redacted");
assert!(format!("{m:?}").contains("••••••"), "Debug is redacted");
}
#[test]
fn serialize_emits_ciphertext_not_plaintext() {
let (public, secret) = MaskKeyring::generate();
set_mask_keyring(MaskKeyring::from_base64(&public, Some(&secret)).unwrap());
let m: Masked = Masked::new("0712-secret");
let json = serde_json::to_string(&m).unwrap();
assert!(
!json.contains("0712-secret"),
"serialized form must not be the plaintext"
);
assert_ne!(
json,
format!("\"{REDACTED}\""),
"serialized form is ciphertext, not the redaction marker"
);
}
#[test]
fn in_memory_plaintext_reveals_without_keyring() {
let m: Masked = Masked::new("hello");
assert_eq!(m.reveal().unwrap(), "hello");
}
}