use std::collections::HashMap;
use std::sync::Arc;
use aes_gcm::{
Aes256Gcm, KeyInit, Nonce,
aead::{Aead, Payload},
};
use axess_identity::{TenantId, UserId};
use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use crate::delegated::stored::credential::{DelegatedCredentialStore, StoredDelegation};
use axess_factors::ZeroizedString;
use axess_rng::SecureRng;
const ENVELOPE_VERSION: &str = "v1";
const NONCE_LEN: usize = 12;
const FIELD_ACCESS: &[u8] = b"access_token";
const FIELD_REFRESH: &[u8] = b"refresh_token";
const AAD_SEP: u8 = 0x1f;
pub struct EncryptionKey([u8; 32]);
impl EncryptionKey {
pub fn from_bytes(bytes: [u8; 32]) -> Self {
Self(bytes)
}
fn as_array(&self) -> &[u8; 32] {
&self.0
}
}
impl Drop for EncryptionKey {
fn drop(&mut self) {
zeroize::Zeroize::zeroize(&mut self.0);
}
}
impl core::fmt::Debug for EncryptionKey {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str("EncryptionKey(***)")
}
}
pub trait KeyProvider: Send + Sync + 'static {
fn current(&self) -> Result<CurrentKey, KeyProviderError>;
fn resolve(&self, key_id: &str) -> Result<Option<Arc<EncryptionKey>>, KeyProviderError>;
}
#[derive(Clone)]
pub struct CurrentKey {
pub key_id: Arc<str>,
pub key: Arc<EncryptionKey>,
}
#[derive(Debug, thiserror::Error)]
pub enum KeyProviderError {
#[error("key provider failed: {0}")]
Failed(String),
#[error("key id contains '.', which is reserved by the envelope format: {0:?}")]
InvalidKeyId(String),
}
#[derive(Clone, Debug)]
pub struct MemoryKeyProvider {
current_id: Arc<str>,
current_key: Arc<EncryptionKey>,
historical: HashMap<String, Arc<EncryptionKey>>,
}
impl MemoryKeyProvider {
pub fn new(key_id: impl Into<String>, key: [u8; 32]) -> Result<Self, KeyProviderError> {
let id = key_id.into();
validate_key_id(&id)?;
Ok(Self {
current_id: Arc::from(id),
current_key: Arc::new(EncryptionKey::from_bytes(key)),
historical: HashMap::new(),
})
}
pub fn with_historical(
mut self,
key_id: impl Into<String>,
key: [u8; 32],
) -> Result<Self, KeyProviderError> {
let id = key_id.into();
validate_key_id(&id)?;
self.historical
.insert(id, Arc::new(EncryptionKey::from_bytes(key)));
Ok(self)
}
}
impl KeyProvider for MemoryKeyProvider {
fn current(&self) -> Result<CurrentKey, KeyProviderError> {
Ok(CurrentKey {
key_id: self.current_id.clone(),
key: self.current_key.clone(),
})
}
fn resolve(&self, key_id: &str) -> Result<Option<Arc<EncryptionKey>>, KeyProviderError> {
if key_id == &*self.current_id {
return Ok(Some(self.current_key.clone()));
}
Ok(self.historical.get(key_id).cloned())
}
}
fn validate_key_id(id: &str) -> Result<(), KeyProviderError> {
if id.is_empty() || id.contains('.') {
return Err(KeyProviderError::InvalidKeyId(id.to_string()));
}
Ok(())
}
pub struct EncryptedDelegatedCredentialStore<S, K> {
inner: S,
keys: K,
}
impl<S, K> EncryptedDelegatedCredentialStore<S, K>
where
S: DelegatedCredentialStore,
K: KeyProvider,
{
pub fn new(inner: S, keys: K) -> Self {
Self { inner, keys }
}
pub fn inner(&self) -> &S {
&self.inner
}
}
impl<S, K> DelegatedCredentialStore for EncryptedDelegatedCredentialStore<S, K>
where
S: DelegatedCredentialStore,
K: KeyProvider,
{
async fn load(
&self,
tenant: &TenantId,
user: &UserId,
provider: &str,
) -> Result<Option<StoredDelegation>, String> {
let Some(cred) = self.inner.load(tenant, user, provider).await? else {
return Ok(None);
};
let aad_access = build_aad(provider, tenant, user, FIELD_ACCESS);
let access_plain = decrypt_envelope(&self.keys, &cred.access_token, &aad_access)
.map_err(|e| format!("decrypt access_token: {e}"))?;
let refresh_plain = match cred.refresh_token.as_deref() {
Some(rt) => {
let aad_refresh = build_aad(provider, tenant, user, FIELD_REFRESH);
let plain = decrypt_envelope(&self.keys, rt, &aad_refresh)
.map_err(|e| format!("decrypt refresh_token: {e}"))?;
Some(ZeroizedString::from(plain))
}
None => None,
};
Ok(Some(StoredDelegation {
provider: cred.provider,
access_token: ZeroizedString::from(access_plain),
refresh_token: refresh_plain,
expires_at: cred.expires_at,
scopes: cred.scopes,
token_type: cred.token_type,
}))
}
async fn save(
&self,
tenant: &TenantId,
user: &UserId,
credential: StoredDelegation,
) -> Result<(), String> {
let StoredDelegation {
provider,
access_token,
refresh_token,
expires_at,
scopes,
token_type,
} = credential;
let current = self.keys.current().map_err(|e| e.to_string())?;
let aad_access = build_aad(&provider, tenant, user, FIELD_ACCESS);
let access_env = encrypt_envelope(¤t, &access_token, &aad_access)
.map_err(|e| format!("encrypt access_token: {e}"))?;
let refresh_env = match refresh_token.as_deref() {
Some(rt) => {
let aad_refresh = build_aad(&provider, tenant, user, FIELD_REFRESH);
Some(ZeroizedString::from(
encrypt_envelope(¤t, rt, &aad_refresh)
.map_err(|e| format!("encrypt refresh_token: {e}"))?,
))
}
None => None,
};
let wrapped = StoredDelegation {
provider,
access_token: ZeroizedString::from(access_env),
refresh_token: refresh_env,
expires_at,
scopes,
token_type,
};
self.inner.save(tenant, user, wrapped).await
}
async fn revoke(&self, tenant: &TenantId, user: &UserId, provider: &str) -> Result<(), String> {
self.inner.revoke(tenant, user, provider).await
}
}
fn build_aad(provider: &str, tenant: &TenantId, user: &UserId, field: &[u8]) -> Vec<u8> {
let provider_bytes = provider.as_bytes();
let tenant_bytes = tenant.as_bytes();
let user_bytes = user.as_bytes();
let mut buf = Vec::with_capacity(provider_bytes.len() + 16 + 16 + field.len() + 3);
buf.extend_from_slice(provider_bytes);
buf.push(AAD_SEP);
buf.extend_from_slice(tenant_bytes);
buf.push(AAD_SEP);
buf.extend_from_slice(user_bytes);
buf.push(AAD_SEP);
buf.extend_from_slice(field);
buf
}
#[derive(Debug, thiserror::Error)]
enum EnvelopeError {
#[error("malformed envelope: {0}")]
Malformed(&'static str),
#[error("unknown envelope version {0:?}; store written by a newer axess?")]
UnknownVersion(String),
#[error("unknown key id {0:?}")]
UnknownKeyId(String),
#[error("key provider error: {0}")]
KeyProvider(#[from] KeyProviderError),
#[error("decryption failed (wrong key, corrupted ciphertext, or AAD mismatch)")]
Decrypt,
#[error("base64 decode failed")]
Base64,
#[error("encryption failed")]
Encrypt,
}
fn encrypt_envelope(
current: &CurrentKey,
plaintext: &str,
aad: &[u8],
) -> Result<String, EnvelopeError> {
let cipher =
Aes256Gcm::new_from_slice(current.key.as_array()).map_err(|_| EnvelopeError::Encrypt)?;
let mut nonce_bytes = [0u8; NONCE_LEN];
axess_rng::SystemRng.fill_bytes(&mut nonce_bytes);
let nonce = Nonce::try_from(&nonce_bytes[..]).map_err(|_| EnvelopeError::Encrypt)?;
let ct = cipher
.encrypt(
&nonce,
Payload {
msg: plaintext.as_bytes(),
aad,
},
)
.map_err(|_| EnvelopeError::Encrypt)?;
let mut body = Vec::with_capacity(NONCE_LEN + ct.len());
body.extend_from_slice(&nonce_bytes);
body.extend_from_slice(&ct);
let b64 = URL_SAFE_NO_PAD.encode(&body);
Ok(format!("{ENVELOPE_VERSION}.{}.{}", &*current.key_id, b64))
}
fn decrypt_envelope<K: KeyProvider>(
keys: &K,
envelope: &str,
aad: &[u8],
) -> Result<String, EnvelopeError> {
let mut parts = envelope.splitn(3, '.');
let version = parts.next().ok_or(EnvelopeError::Malformed("no version"))?;
let key_id = parts.next().ok_or(EnvelopeError::Malformed("no key id"))?;
let body_b64 = parts.next().ok_or(EnvelopeError::Malformed("no body"))?;
if version != ENVELOPE_VERSION {
return Err(EnvelopeError::UnknownVersion(version.to_string()));
}
if key_id.is_empty() {
return Err(EnvelopeError::Malformed("empty key id"));
}
let body = URL_SAFE_NO_PAD
.decode(body_b64)
.map_err(|_| EnvelopeError::Base64)?;
if body.len() <= NONCE_LEN {
return Err(EnvelopeError::Malformed("body shorter than nonce"));
}
let (nonce_bytes, ciphertext) = body.split_at(NONCE_LEN);
let key = keys
.resolve(key_id)?
.ok_or_else(|| EnvelopeError::UnknownKeyId(key_id.to_string()))?;
let cipher = Aes256Gcm::new_from_slice(key.as_array()).map_err(|_| EnvelopeError::Decrypt)?;
let nonce = Nonce::try_from(nonce_bytes).map_err(|_| EnvelopeError::Decrypt)?;
let plaintext = cipher
.decrypt(
&nonce,
Payload {
msg: ciphertext,
aad,
},
)
.map_err(|_| EnvelopeError::Decrypt)?;
String::from_utf8(plaintext).map_err(|_| EnvelopeError::Malformed("plaintext not utf-8"))
}
#[cfg(test)]
mod tests;