use aes::Aes256;
use aes::cipher::block_padding::Pkcs7;
use aes::cipher::{BlockModeDecrypt, BlockModeEncrypt, KeyIvInit};
use base64::Engine;
use base64::engine::general_purpose::STANDARD as BASE64;
use secp256k1::{Parity, ecdh};
use thiserror::Error;
use zeroize::Zeroize;
use crate::key::{PublicKey, SecretKey};
use crate::util::rng::{self, RngError};
const KEY_BYTES: usize = 32;
const IV_BYTES: usize = 16;
const SEPARATOR: &str = "?iv=";
type Aes256CbcEnc = cbc::Encryptor<Aes256>;
type Aes256CbcDec = cbc::Decryptor<Aes256>;
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum Nip04Error {
#[error("payload is missing the `?iv=` separator")]
MissingIvSeparator,
#[error("base64 decode failed: {0}")]
Base64(#[from] base64::DecodeError),
#[error("IV must be {expected} bytes, got {actual}")]
InvalidIvLength {
expected: usize,
actual: usize,
},
#[error("AES-256-CBC unpadding failed (wrong key or tampered payload)")]
Unpad,
#[error("decrypted plaintext is not valid UTF-8")]
InvalidUtf8,
#[error(transparent)]
Rng(#[from] RngError),
}
fn shared_secret_x(secret: &SecretKey, peer: &PublicKey) -> [u8; KEY_BYTES] {
let normalized = secp256k1::PublicKey::from_x_only_public_key(*peer.as_inner(), Parity::Even);
let ssp = ecdh::shared_secret_point(&normalized, secret.as_inner());
let mut x = [0_u8; KEY_BYTES];
x.copy_from_slice(&ssp[..KEY_BYTES]);
x
}
pub fn encrypt(
secret: &SecretKey,
peer: &PublicKey,
plaintext: &str,
) -> Result<String, Nip04Error> {
let mut key = shared_secret_x(secret, peer);
let iv = rng::random_bytes::<IV_BYTES>()?;
let ciphertext = Aes256CbcEnc::new((&key).into(), (&iv).into())
.encrypt_padded_vec::<Pkcs7>(plaintext.as_bytes());
key.zeroize();
Ok(format!(
"{}{SEPARATOR}{}",
BASE64.encode(&ciphertext),
BASE64.encode(iv),
))
}
pub fn decrypt(secret: &SecretKey, peer: &PublicKey, payload: &str) -> Result<String, Nip04Error> {
let (ciphertext_b64, iv_b64) = payload
.split_once(SEPARATOR)
.ok_or(Nip04Error::MissingIvSeparator)?;
let mut ciphertext = BASE64.decode(ciphertext_b64)?;
let iv_bytes = BASE64.decode(iv_b64)?;
let iv: [u8; IV_BYTES] =
iv_bytes
.as_slice()
.try_into()
.map_err(|_| Nip04Error::InvalidIvLength {
expected: IV_BYTES,
actual: iv_bytes.len(),
})?;
let mut key = shared_secret_x(secret, peer);
let plaintext = Aes256CbcDec::new((&key).into(), (&iv).into())
.decrypt_padded::<Pkcs7>(&mut ciphertext)
.map_err(|_| Nip04Error::Unpad)?;
let result = std::str::from_utf8(plaintext)
.map_err(|_| Nip04Error::InvalidUtf8)?
.to_owned();
key.zeroize();
Ok(result)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::key::Keys;
fn keys_alice() -> Keys {
Keys::parse("000000000000000000000000000000000000000000000000000000000000a1ce").unwrap()
}
fn keys_bob() -> Keys {
Keys::parse("00000000000000000000000000000000000000000000000000000000000000b0").unwrap()
}
#[test]
fn round_trip_short_message() {
let alice = keys_alice();
let bob = keys_bob();
let ciphertext = encrypt(alice.secret_key(), bob.public_key(), "hello").unwrap();
let recovered = decrypt(bob.secret_key(), alice.public_key(), &ciphertext).unwrap();
assert_eq!(recovered, "hello");
}
#[test]
fn round_trip_unicode_payload() {
let alice = keys_alice();
let bob = keys_bob();
let msg = "你好,nostr 🦀";
let ciphertext = encrypt(alice.secret_key(), bob.public_key(), msg).unwrap();
let recovered = decrypt(bob.secret_key(), alice.public_key(), &ciphertext).unwrap();
assert_eq!(recovered, msg);
}
#[test]
fn fresh_iv_per_call_yields_distinct_ciphertexts() {
let alice = keys_alice();
let bob = keys_bob();
let a = encrypt(alice.secret_key(), bob.public_key(), "same").unwrap();
let b = encrypt(alice.secret_key(), bob.public_key(), "same").unwrap();
assert_ne!(a, b, "two encryptions of the same plaintext must differ");
}
#[test]
fn empty_plaintext_round_trip() {
let alice = keys_alice();
let bob = keys_bob();
let ciphertext = encrypt(alice.secret_key(), bob.public_key(), "").unwrap();
let recovered = decrypt(bob.secret_key(), alice.public_key(), &ciphertext).unwrap();
assert_eq!(recovered, "");
}
#[test]
fn missing_separator_is_rejected() {
let alice = keys_alice();
let bob = keys_bob();
let err = decrypt(alice.secret_key(), bob.public_key(), "no-separator-here").unwrap_err();
assert!(matches!(err, Nip04Error::MissingIvSeparator));
}
#[test]
fn malformed_base64_is_rejected() {
let alice = keys_alice();
let bob = keys_bob();
let err = decrypt(
alice.secret_key(),
bob.public_key(),
"!!not-base64!!?iv=!!neither!!",
)
.unwrap_err();
assert!(matches!(err, Nip04Error::Base64(_)));
}
#[test]
fn wrong_iv_length_is_rejected() {
let alice = keys_alice();
let bob = keys_bob();
let payload = format!(
"{}{SEPARATOR}{}",
BASE64.encode([0_u8; 32]),
BASE64.encode([0_u8; 8]),
);
let err = decrypt(alice.secret_key(), bob.public_key(), &payload).unwrap_err();
assert!(matches!(
err,
Nip04Error::InvalidIvLength {
expected: 16,
actual: 8,
}
));
}
#[test]
fn wrong_peer_yields_unpad_error() {
let alice = keys_alice();
let bob = keys_bob();
let mallory =
Keys::parse("00000000000000000000000000000000000000000000000000000000000ca800")
.unwrap();
let ciphertext = encrypt(alice.secret_key(), bob.public_key(), "for bob").unwrap();
let err = decrypt(mallory.secret_key(), alice.public_key(), &ciphertext).unwrap_err();
assert!(matches!(err, Nip04Error::Unpad | Nip04Error::InvalidUtf8));
}
}