use std::collections::BTreeMap;
use nostro2_traits::NostrKeypair;
use nostro2_traits::hex::Hexable;
use zeroize::Zeroize;
use super::{Nip104Crypto, Nip104Error};
type Result<T> = std::result::Result<T, Nip104Error>;
pub const SENDER_KEY_MAX_SKIP: u32 = 10_000;
pub const SENDER_KEY_MAX_STORED_SKIPPED_KEYS: usize = 2_000;
const SENDER_KEY_KDF_SALT: &[u8] = b"ndr-sender-key-v1";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SenderKeyState {
pub key_id: u32,
chain_key: String,
iteration: u32,
skipped_message_keys: BTreeMap<u32, String>,
}
impl Drop for SenderKeyState {
fn drop(&mut self) {
self.chain_key.zeroize();
for key in self.skipped_message_keys.values_mut() {
key.zeroize();
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SenderKeyEncryptPlan {
pub next_state: SenderKeyState,
pub key_id: u32,
pub message_number: u32,
pub ciphertext: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SenderKeyDecryptPlan {
pub next_state: SenderKeyState,
pub plaintext: Vec<u8>,
}
impl SenderKeyState {
#[must_use]
pub fn new(key_id: u32, chain_key: &[u8; 32], iteration: u32) -> Self {
Self {
key_id,
chain_key: chain_key.to_hex(),
iteration,
skipped_message_keys: BTreeMap::new(),
}
}
#[must_use]
pub const fn key_id(&self) -> u32 {
self.key_id
}
#[must_use]
pub const fn iteration(&self) -> u32 {
self.iteration
}
#[must_use]
pub fn chain_key_hex(&self) -> String {
self.chain_key.clone()
}
#[must_use]
pub fn skipped_len(&self) -> usize {
self.skipped_message_keys.len()
}
#[must_use]
pub const fn skipped_keys(&self) -> &BTreeMap<u32, String> {
&self.skipped_message_keys
}
#[must_use]
pub const fn from_parts(
key_id: u32,
chain_key_hex: String,
iteration: u32,
skipped_message_keys: BTreeMap<u32, String>,
) -> Self {
Self {
key_id,
chain_key: chain_key_hex,
iteration,
skipped_message_keys,
}
}
pub fn plan_encrypt<K: NostrKeypair>(&self, plaintext: &[u8]) -> Result<SenderKeyEncryptPlan> {
let mut next_state = self.clone();
let message_number = next_state.iteration;
let (next_chain_key, message_key) =
Self::derive_message_key::<K>(&K::decode_hex_32(&next_state.chain_key)?);
next_state.chain_key = next_chain_key.to_hex();
next_state.iteration = next_state
.iteration
.checked_add(1)
.ok_or(Nip104Error::SessionNotReady)?;
let ciphertext = K::encrypt_with_message_key(&message_key, plaintext)?;
Ok(SenderKeyEncryptPlan {
next_state,
key_id: self.key_id,
message_number,
ciphertext,
})
}
pub fn apply_encrypt(&mut self, plan: SenderKeyEncryptPlan) {
*self = plan.next_state;
}
pub fn encrypt<K: NostrKeypair>(&mut self, plaintext: &[u8]) -> Result<(u32, String)> {
let plan = self.plan_encrypt::<K>(plaintext)?;
let out = (plan.message_number, plan.ciphertext.clone());
self.apply_encrypt(plan);
Ok(out)
}
pub fn plan_decrypt<K: NostrKeypair>(
&self,
key_id: u32,
message_number: u32,
ciphertext_b64: &str,
) -> Result<SenderKeyDecryptPlan> {
if key_id != self.key_id {
return Err(Nip104Error::InvalidHeader);
}
let mut next_state = self.clone();
let plaintext = next_state.decrypt_in_place::<K>(message_number, ciphertext_b64)?;
Ok(SenderKeyDecryptPlan {
next_state,
plaintext,
})
}
pub fn apply_decrypt(&mut self, plan: SenderKeyDecryptPlan) -> Vec<u8> {
*self = plan.next_state;
plan.plaintext
}
pub fn decrypt<K: NostrKeypair>(
&mut self,
message_number: u32,
ciphertext_b64: &str,
) -> Result<Vec<u8>> {
let plan = self.plan_decrypt::<K>(self.key_id, message_number, ciphertext_b64)?;
Ok(self.apply_decrypt(plan))
}
fn decrypt_in_place<K: NostrKeypair>(
&mut self,
message_number: u32,
ciphertext_b64: &str,
) -> Result<Vec<u8>> {
if message_number < self.iteration {
let key = self
.skipped_message_keys
.remove(&message_number)
.ok_or(Nip104Error::InvalidHeader)?;
return K::decrypt_with_message_key(&K::decode_hex_32(&key)?, ciphertext_b64);
}
let delta = message_number - self.iteration;
if delta > SENDER_KEY_MAX_SKIP {
return Err(Nip104Error::TooManySkippedMessages);
}
while self.iteration < message_number {
let (next_chain_key, message_key) =
Self::derive_message_key::<K>(&K::decode_hex_32(&self.chain_key)?);
self.chain_key = next_chain_key.to_hex();
self.skipped_message_keys
.insert(self.iteration, message_key.to_hex());
self.iteration = self
.iteration
.checked_add(1)
.ok_or(Nip104Error::SessionNotReady)?;
}
let (next_chain_key, message_key) =
Self::derive_message_key::<K>(&K::decode_hex_32(&self.chain_key)?);
self.chain_key = next_chain_key.to_hex();
self.iteration = self
.iteration
.checked_add(1)
.ok_or(Nip104Error::SessionNotReady)?;
Self::prune_skipped(&mut self.skipped_message_keys);
K::decrypt_with_message_key(&message_key, ciphertext_b64)
}
fn derive_message_key<K: NostrKeypair>(chain_key: &[u8; 32]) -> ([u8; 32], [u8; 32]) {
let outs = K::kdf(chain_key, SENDER_KEY_KDF_SALT, 2);
(outs[0], outs[1])
}
fn prune_skipped(map: &mut BTreeMap<u32, String>) {
while map.len() > SENDER_KEY_MAX_STORED_SKIPPED_KEYS {
let Some(first) = map.keys().next().copied() else {
break;
};
map.remove(&first);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
type K = crate::tests::NipTester;
#[test]
fn roundtrip_single_message() {
let ck = [7_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
let (n, ct) = sender.encrypt::<K>(b"hello").unwrap();
assert_eq!(n, 0);
assert_eq!(receiver.decrypt::<K>(n, &ct).unwrap(), b"hello");
assert_eq!(sender.iteration(), receiver.iteration());
}
#[test]
fn decrypt_out_of_order() {
let ck = [9_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
let (n0, c0) = sender.encrypt::<K>(b"m0").unwrap();
let (n1, c1) = sender.encrypt::<K>(b"m1").unwrap();
assert_eq!(receiver.decrypt::<K>(n1, &c1).unwrap(), b"m1");
assert_eq!(receiver.decrypt::<K>(n0, &c0).unwrap(), b"m0");
}
#[test]
fn rejects_duplicate_message() {
let ck = [11_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
let (n, c) = sender.encrypt::<K>(b"once").unwrap();
assert_eq!(receiver.decrypt::<K>(n, &c).unwrap(), b"once");
assert!(receiver.decrypt::<K>(n, &c).is_err());
}
#[test]
fn wrong_key_id_does_not_mutate_receiver() {
let ck = [13_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let receiver = SenderKeyState::new(1, &ck, 0);
let (n, c) = sender.encrypt::<K>(b"x").unwrap();
let before = receiver.clone();
assert!(receiver.plan_decrypt::<K>(2, n, &c).is_err());
assert_eq!(receiver, before);
}
#[test]
fn plan_encrypt_is_pure_until_apply() {
let ck = [4_u8; 32];
let mut sender = SenderKeyState::new(7, &ck, 0);
let before = sender.clone();
let plan = sender.plan_encrypt::<K>(b"deferred").unwrap();
assert_eq!(sender, before);
assert_eq!(plan.key_id, 7);
assert_eq!(plan.message_number, 0);
sender.apply_encrypt(plan);
assert_eq!(sender.iteration(), 1);
assert_ne!(sender, before);
}
#[test]
fn skip_ahead_then_backfill() {
let ck = [19_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
let (n0, c0) = sender.encrypt::<K>(b"m0").unwrap();
let (n1, c1) = sender.encrypt::<K>(b"m1").unwrap();
let (n2, c2) = sender.encrypt::<K>(b"m2").unwrap();
assert_eq!(receiver.decrypt::<K>(n2, &c2).unwrap(), b"m2");
assert_eq!(receiver.skipped_len(), 2);
assert_eq!(receiver.decrypt::<K>(n0, &c0).unwrap(), b"m0");
assert_eq!(receiver.decrypt::<K>(n1, &c1).unwrap(), b"m1");
assert_eq!(receiver.skipped_len(), 0);
}
#[test]
fn rejects_too_many_skipped() {
let ck = [3_u8; 32];
let receiver = SenderKeyState::new(1, &ck, 0);
let err = receiver.plan_decrypt::<K>(1, SENDER_KEY_MAX_SKIP + 1, "AA");
assert!(matches!(err, Err(Nip104Error::TooManySkippedMessages)));
}
#[test]
#[allow(clippy::cast_possible_truncation)] fn skip_ahead_prunes_oldest_skipped_keys() {
const AHEAD: u32 = SENDER_KEY_MAX_STORED_SKIPPED_KEYS as u32 + 500;
let ck = [21_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
let mut first = None;
let mut recent = None;
for i in 0..=AHEAD {
let (n, c) = sender.encrypt::<K>(format!("m{i}").as_bytes()).unwrap();
if i == 0 {
first = Some((n, c.clone()));
}
if i == AHEAD - 1 {
recent = Some((n, c.clone()));
}
if i == AHEAD {
assert_eq!(
receiver.decrypt::<K>(n, &c).unwrap(),
format!("m{i}").as_bytes()
);
}
}
assert_eq!(
receiver.skipped_len(),
SENDER_KEY_MAX_STORED_SKIPPED_KEYS,
"stored skipped keys must be capped"
);
let (rn, rc) = recent.unwrap();
assert_eq!(
receiver.decrypt::<K>(rn, &rc).unwrap(),
format!("m{}", AHEAD - 1).as_bytes()
);
let (fn0, fc0) = first.unwrap();
assert!(matches!(
receiver.plan_decrypt::<K>(1, fn0, &fc0),
Err(Nip104Error::InvalidHeader)
));
}
#[test]
fn skip_ceiling_is_inclusive() {
let ck = [23_u8; 32];
let receiver = SenderKeyState::new(1, &ck, 0);
let at_ceiling = receiver.plan_decrypt::<K>(1, SENDER_KEY_MAX_SKIP, "not-base64!!");
assert!(
!matches!(at_ceiling, Err(Nip104Error::TooManySkippedMessages)),
"delta == MAX_SKIP must pass the skip gate"
);
let over = receiver.plan_decrypt::<K>(1, SENDER_KEY_MAX_SKIP + 1, "not-base64!!");
assert!(matches!(over, Err(Nip104Error::TooManySkippedMessages)));
}
#[test]
fn iteration_overflow_on_encrypt_is_rejected() {
let ck = [25_u8; 32];
let sender = SenderKeyState::new(1, &ck, u32::MAX);
assert!(matches!(
sender.plan_encrypt::<K>(b"boom"),
Err(Nip104Error::SessionNotReady)
));
}
#[test]
fn payload_size_bounds_match_nip44() {
let ck = [27_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
let (n0, c0) = sender.encrypt::<K>(b"x").unwrap();
assert_eq!(receiver.decrypt::<K>(n0, &c0).unwrap(), b"x");
let max = vec![0x5A_u8; 65_535];
let (n1, c1) = sender.encrypt::<K>(&max).unwrap();
assert_eq!(receiver.decrypt::<K>(n1, &c1).unwrap(), max);
assert!(sender.plan_encrypt::<K>(b"").is_err());
let over = vec![0_u8; 65_536];
assert!(sender.plan_encrypt::<K>(&over).is_err());
assert_eq!(sender.iteration(), 2);
}
#[test]
fn tampered_ciphertext_fails_without_advancing() {
use base64::engine::{Engine as _, general_purpose};
let ck = [29_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let receiver = SenderKeyState::new(1, &ck, 0);
let (n, c) = sender.encrypt::<K>(b"authentic").unwrap();
let mut raw = general_purpose::STANDARD.decode(&c).unwrap();
let mid = raw.len() / 2;
raw[mid] ^= 0xFF;
let tampered = general_purpose::STANDARD.encode(&raw);
let before = receiver.clone();
assert!(receiver.plan_decrypt::<K>(1, n, &tampered).is_err());
assert_eq!(receiver, before, "failed decrypt must not mutate state");
let mut receiver = receiver;
assert_eq!(receiver.decrypt::<K>(n, &c).unwrap(), b"authentic");
}
#[test]
fn sustained_in_order_volume() {
const N: u32 = 5_000;
let ck = [31_u8; 32];
let mut sender = SenderKeyState::new(1, &ck, 0);
let mut receiver = SenderKeyState::new(1, &ck, 0);
for i in 0..N {
let (n, c) = sender.encrypt::<K>(format!("msg-{i}").as_bytes()).unwrap();
assert_eq!(n, i);
assert_eq!(
receiver.decrypt::<K>(n, &c).unwrap(),
format!("msg-{i}").as_bytes()
);
}
assert_eq!(sender.iteration(), N);
assert_eq!(receiver.iteration(), N);
assert_eq!(receiver.skipped_len(), 0);
}
}