use std::collections::BTreeMap;
use super::{Invite, MESSAGE_EVENT_KIND, Nip104Error, Session};
use nostro2_traits::NostrKeypair;
type Result<T> = std::result::Result<T, Nip104Error>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReceivedMessage {
pub peer: String,
pub device_id: String,
pub plaintext: Vec<u8>,
}
#[derive(Debug, Clone)]
struct PeerRecord<K: NostrKeypair> {
devices: BTreeMap<String, Session<K>>,
}
impl<K: NostrKeypair> Default for PeerRecord<K> {
fn default() -> Self {
Self {
devices: BTreeMap::new(),
}
}
}
type SessionKey = (String, String);
#[derive(Debug, Clone)]
pub struct SessionManager<K: NostrKeypair> {
identity: K,
our_pubkey: String,
peers: BTreeMap<String, PeerRecord<K>>,
sender_index: BTreeMap<String, SessionKey>,
}
impl<K: NostrKeypair> SessionManager<K> {
#[must_use]
pub fn new(identity: K) -> Self {
let our_pubkey = identity.public_key();
Self {
identity,
our_pubkey,
peers: BTreeMap::new(),
sender_index: BTreeMap::new(),
}
}
#[must_use]
pub fn our_pubkey(&self) -> &str {
&self.our_pubkey
}
#[must_use]
pub fn has_session(&self, peer: &str) -> bool {
self.peers.get(peer).is_some_and(|p| !p.devices.is_empty())
}
pub fn peers(&self) -> impl Iterator<Item = &String> {
self.peers.keys()
}
#[must_use]
pub fn devices(&self, peer: &str) -> Vec<String> {
self.peers
.get(peer)
.map(|p| p.devices.keys().cloned().collect())
.unwrap_or_default()
}
#[must_use]
pub fn session_count(&self) -> usize {
self.peers.values().map(|p| p.devices.len()).sum()
}
pub fn install_session(&mut self, peer: &str, device_id: &str, session: Session<K>) {
let slot = (peer.to_owned(), device_id.to_owned());
self.forget_in_index(&slot);
for sender in session.accepted_senders() {
self.sender_index.insert(sender, slot.clone());
}
self.peers
.entry(peer.to_owned())
.or_default()
.devices
.insert(device_id.to_owned(), session);
}
pub fn sessions(&self) -> impl Iterator<Item = ((&str, &str), &super::SessionState)> {
self.peers.iter().flat_map(|(peer, record)| {
record
.devices
.iter()
.map(move |(device, session)| ((peer.as_str(), device.as_str()), &session.state))
})
}
fn forget_in_index(&mut self, slot: &SessionKey) {
self.sender_index.retain(|_, v| v != slot);
}
fn reindex(&mut self, slot: &SessionKey) {
self.forget_in_index(slot);
if let Some(session) = self.peers.get(&slot.0).and_then(|p| p.devices.get(&slot.1)) {
for sender in session.accepted_senders() {
self.sender_index.insert(sender, slot.clone());
}
}
}
pub fn accept_invite(
&mut self,
invite: &Invite,
owner_pubkey: Option<&str>,
created_at: i64,
) -> Result<nostro2::NostrNote> {
let (session, response) = invite.accept::<K>(&self.identity, owner_pubkey, created_at)?;
let device_id = invite
.device_id
.clone()
.unwrap_or_else(|| invite.inviter.clone());
self.install_session(&invite.inviter, &device_id, session);
Ok(response)
}
pub fn receive_invite_response(
&mut self,
invite: &Invite,
event: &nostro2::NostrNote,
) -> Result<String> {
let (session, recovered) = invite.receive::<K>(event, &self.identity)?;
let peer = recovered
.owner_public_key
.clone()
.unwrap_or_else(|| recovered.invitee_identity.clone());
self.install_session(&peer, &recovered.invitee_identity, session);
Ok(peer)
}
pub fn process_event(&mut self, event: &nostro2::NostrNote) -> Option<ReceivedMessage> {
if event.kind != MESSAGE_EVENT_KIND {
return None;
}
let slot = self.sender_index.get(&event.pubkey)?.clone();
let session = self.peers.get_mut(&slot.0)?.devices.get_mut(&slot.1)?;
let (next, plaintext) = session.plan_receive_event(event).ok()?;
session.apply(next);
self.reindex(&slot);
Some(ReceivedMessage {
peer: slot.0,
device_id: slot.1,
plaintext,
})
}
pub fn send(
&mut self,
peer: &str,
payload: &[u8],
created_at: i64,
) -> Result<Vec<nostro2::NostrNote>> {
let record = self
.peers
.get_mut(peer)
.filter(|p| !p.devices.is_empty())
.ok_or_else(|| Nip104Error::UnknownPeer(peer.to_owned()))?;
let mut events = Vec::with_capacity(record.devices.len());
for session in record.devices.values_mut() {
if !session.can_send() {
continue;
}
let (next, event) = session.plan_send_event(payload, created_at)?;
session.apply(next);
events.push(event);
}
Ok(events)
}
}
#[cfg(test)]
mod tests {
use super::*;
use nostro2_traits::NostrSigner as _;
type K = crate::tests::NipTester;
fn ident(seed: u8) -> K {
K::from_secret_bytes(&[seed; 32]).unwrap()
}
const NOW: i64 = 1_700_000_000;
#[test]
fn two_managers_handshake_and_chat() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let invite = Invite::create_new::<K>(alice.our_pubkey(), None).unwrap();
let response = bob.accept_invite(&invite, None, NOW).unwrap();
assert!(bob.has_session(alice.our_pubkey()));
let peer = alice.receive_invite_response(&invite, &response).unwrap();
assert_eq!(peer, bob.our_pubkey());
assert!(alice.has_session(bob.our_pubkey()));
let outbound = bob.send(alice.our_pubkey(), b"hello alice", NOW).unwrap();
assert_eq!(outbound.len(), 1);
let got = alice.process_event(&outbound[0]).expect("alice decrypts");
assert_eq!(got.peer, bob.our_pubkey());
assert_eq!(got.plaintext, b"hello alice");
let reply = alice.send(bob.our_pubkey(), b"hi bob", NOW).unwrap();
assert_eq!(reply.len(), 1);
let got = bob.process_event(&reply[0]).expect("bob decrypts");
assert_eq!(got.peer, alice.our_pubkey());
assert_eq!(got.plaintext, b"hi bob");
}
#[test]
fn send_fans_out_to_every_device() {
let mut alice = SessionManager::new(ident(0x10));
let bob_owner = ident(0x20).public_key();
let mut bob_dev1 = SessionManager::new(ident(0x21));
let mut bob_dev2 = SessionManager::new(ident(0x22));
let invite = Invite::create_new::<K>(alice.our_pubkey(), None).unwrap();
let r1 = bob_dev1
.accept_invite(&invite, Some(&bob_owner), NOW)
.unwrap();
let r2 = bob_dev2
.accept_invite(&invite, Some(&bob_owner), NOW)
.unwrap();
let p1 = alice.receive_invite_response(&invite, &r1).unwrap();
let p2 = alice.receive_invite_response(&invite, &r2).unwrap();
assert_eq!(p1, bob_owner);
assert_eq!(p2, bob_owner);
assert_eq!(alice.devices(&bob_owner).len(), 2);
let m1 = bob_dev1.send(alice.our_pubkey(), b"d1 up", NOW).unwrap();
let m2 = bob_dev2.send(alice.our_pubkey(), b"d2 up", NOW).unwrap();
assert_eq!(alice.process_event(&m1[0]).unwrap().plaintext, b"d1 up");
assert_eq!(alice.process_event(&m2[0]).unwrap().plaintext, b"d2 up");
let fanned = alice.send(&bob_owner, b"broadcast", NOW).unwrap();
assert_eq!(fanned.len(), 2);
let to_dev1 = fanned
.iter()
.filter_map(|e| bob_dev1.process_event(e))
.collect::<Vec<_>>();
let to_dev2 = fanned
.iter()
.filter_map(|e| bob_dev2.process_event(e))
.collect::<Vec<_>>();
assert_eq!(to_dev1.len(), 1);
assert_eq!(to_dev2.len(), 1);
assert_eq!(to_dev1[0].plaintext, b"broadcast");
assert_eq!(to_dev2[0].plaintext, b"broadcast");
}
#[test]
fn sessions_snapshot_round_trip() {
let mut alice = SessionManager::new(ident(0x51));
let mut bob = SessionManager::new(ident(0x52));
let invite = Invite::create_new::<K>(alice.our_pubkey(), None).unwrap();
let response = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &response).unwrap();
let up = bob.send(alice.our_pubkey(), b"hello", NOW).unwrap();
assert_eq!(alice.process_event(&up[0]).unwrap().plaintext, b"hello");
let snaps: Vec<((String, String), super::super::SessionState)> = alice
.sessions()
.map(|((p, d), st)| ((p.to_owned(), d.to_owned()), st.clone()))
.collect();
assert_eq!(snaps.len(), 1);
let mut alice2 = SessionManager::new(ident(0x51));
for ((peer, device), state) in snaps {
alice2.install_session(&peer, &device, Session::from_state(state));
}
let up2 = bob.send(alice.our_pubkey(), b"again", NOW).unwrap();
assert_eq!(alice2.process_event(&up2[0]).unwrap().plaintext, b"again");
let reply = alice2.send(bob.our_pubkey(), b"hi back", NOW).unwrap();
assert_eq!(bob.process_event(&reply[0]).unwrap().plaintext, b"hi back");
}
#[test]
fn send_to_unknown_peer_errors() {
let mut alice = SessionManager::new(ident(0x30));
let err = alice.send("deadbeef", b"hi", NOW).unwrap_err();
assert!(matches!(err, Nip104Error::UnknownPeer(_)));
}
#[test]
fn process_ignores_foreign_and_non_message_events() {
let mut alice = SessionManager::new(ident(0x40));
let mut bob = SessionManager::new(ident(0x41));
let invite = Invite::create_new::<K>(alice.our_pubkey(), None).unwrap();
let response = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &response).unwrap();
let outbound = bob.send(alice.our_pubkey(), b"hi", NOW).unwrap();
let mut stranger = SessionManager::new(ident(0x42));
assert!(stranger.process_event(&outbound[0]).is_none());
assert!(alice.process_event(&response).is_none());
}
#[test]
fn fan_out_to_many_devices() {
const DEVICES: u8 = 24;
let mut alice = SessionManager::new(ident(0x01));
let owner = ident(0x02).public_key();
let invite = Invite::create_new::<K>(alice.our_pubkey(), None).unwrap();
let mut devices: Vec<SessionManager<K>> = Vec::new();
for d in 0..DEVICES {
let mut dev = SessionManager::new(ident(0x10 + d));
let resp = dev.accept_invite(&invite, Some(&owner), NOW).unwrap();
alice.receive_invite_response(&invite, &resp).unwrap();
let up = dev.send(alice.our_pubkey(), b"up", NOW).unwrap();
assert_eq!(alice.process_event(&up[0]).unwrap().plaintext, b"up");
devices.push(dev);
}
assert_eq!(alice.devices(&owner).len(), DEVICES as usize);
let fanned = alice.send(&owner, b"broadcast", NOW).unwrap();
assert_eq!(fanned.len(), DEVICES as usize);
for dev in &mut devices {
let hits: Vec<_> = fanned.iter().filter_map(|e| dev.process_event(e)).collect();
assert_eq!(hits.len(), 1, "each device takes exactly one copy");
assert_eq!(hits[0].plaintext, b"broadcast");
}
}
#[test]
fn sustained_bidirectional_conversation() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let apk = alice.our_pubkey().to_owned();
let bpk = bob.our_pubkey().to_owned();
let invite = Invite::create_new::<K>(&apk, None).unwrap();
let resp = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &resp).unwrap();
let first = bob.send(&apk, b"hi", NOW).unwrap();
assert_eq!(alice.process_event(&first[0]).unwrap().plaintext, b"hi");
for i in 0..100 {
let a_body = format!("a{i}");
let ev = alice.send(&bpk, a_body.as_bytes(), NOW).unwrap();
assert_eq!(
bob.process_event(&ev[0]).unwrap().plaintext,
a_body.as_bytes()
);
let b_body = format!("b{i}");
let ev = bob.send(&apk, b_body.as_bytes(), NOW).unwrap();
assert_eq!(
alice.process_event(&ev[0]).unwrap().plaintext,
b_body.as_bytes()
);
}
}
#[test]
fn replayed_message_event_ignored() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let apk = alice.our_pubkey().to_owned();
let invite = Invite::create_new::<K>(&apk, None).unwrap();
let resp = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &resp).unwrap();
let ev = bob.send(&apk, b"only once", NOW).unwrap();
assert_eq!(alice.process_event(&ev[0]).unwrap().plaintext, b"only once");
assert!(alice.process_event(&ev[0]).is_none());
}
#[test]
fn message_does_not_decrypt_under_foreign_session() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let mut mallory = SessionManager::new(ident(0x03));
let apk = alice.our_pubkey().to_owned();
let mpk_owner = mallory.our_pubkey().to_owned();
let _ = mpk_owner;
let inv_b = Invite::create_new::<K>(&apk, None).unwrap();
let rb = bob.accept_invite(&inv_b, None, NOW).unwrap();
alice.receive_invite_response(&inv_b, &rb).unwrap();
let inv_m = Invite::create_new::<K>(&apk, None).unwrap();
let rm = mallory.accept_invite(&inv_m, None, NOW).unwrap();
alice.receive_invite_response(&inv_m, &rm).unwrap();
let ev = bob.send(&apk, b"for alice only", NOW).unwrap();
assert!(mallory.process_event(&ev[0]).is_none());
assert_eq!(
alice.process_event(&ev[0]).unwrap().plaintext,
b"for alice only"
);
}
#[test]
fn out_of_order_events_route_and_backfill() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let apk = alice.our_pubkey().to_owned();
let bpk = bob.our_pubkey().to_owned();
let invite = Invite::create_new::<K>(&apk, None).unwrap();
let resp = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &resp).unwrap();
let first = bob.send(&apk, b"open", NOW).unwrap();
alice.process_event(&first[0]).unwrap();
let e1 = alice.send(&bpk, b"m1", NOW).unwrap().pop().unwrap();
let e2 = alice.send(&bpk, b"m2", NOW).unwrap().pop().unwrap();
let e3 = alice.send(&bpk, b"m3", NOW).unwrap().pop().unwrap();
assert_eq!(bob.process_event(&e1).unwrap().plaintext, b"m1");
assert_eq!(bob.process_event(&e3).unwrap().plaintext, b"m3");
assert_eq!(bob.process_event(&e2).unwrap().plaintext, b"m2");
}
#[test]
fn routing_index_tracks_ratchet_turns() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let apk = alice.our_pubkey().to_owned();
let bpk = bob.our_pubkey().to_owned();
let invite = Invite::create_new::<K>(&apk, None).unwrap();
let resp = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &resp).unwrap();
let first = bob.send(&apk, b"open", NOW).unwrap();
alice.process_event(&first[0]).unwrap();
for i in 0..40 {
let a = format!("a{i}");
let ea = alice.send(&bpk, a.as_bytes(), NOW).unwrap();
assert_eq!(bob.process_event(&ea[0]).unwrap().plaintext, a.as_bytes());
let b = format!("b{i}");
let eb = bob.send(&apk, b.as_bytes(), NOW).unwrap();
assert_eq!(alice.process_event(&eb[0]).unwrap().plaintext, b.as_bytes());
}
}
#[test]
fn routing_index_keeps_old_chain_reachable() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let apk = alice.our_pubkey().to_owned();
let bpk = bob.our_pubkey().to_owned();
let invite = Invite::create_new::<K>(&apk, None).unwrap();
let resp = bob.accept_invite(&invite, None, NOW).unwrap();
alice.receive_invite_response(&invite, &resp).unwrap();
let first = bob.send(&apk, b"open", NOW).unwrap();
alice.process_event(&first[0]).unwrap();
let e1 = alice.send(&bpk, b"old-1", NOW).unwrap().pop().unwrap();
let e2 = alice.send(&bpk, b"old-2", NOW).unwrap().pop().unwrap();
assert_eq!(bob.process_event(&e2).unwrap().plaintext, b"old-2");
let reply = bob.send(&apk, b"hi back", NOW).unwrap();
alice.process_event(&reply[0]).unwrap();
assert_eq!(bob.process_event(&e1).unwrap().plaintext, b"old-1");
}
#[test]
fn reinstalling_session_clears_stale_index_rows() {
let mut alice = SessionManager::new(ident(0x01));
let mut bob = SessionManager::new(ident(0x02));
let apk = alice.our_pubkey().to_owned();
let invite = Invite::create_new::<K>(&apk, None).unwrap();
let resp = bob.accept_invite(&invite, None, NOW).unwrap();
let peer = alice.receive_invite_response(&invite, &resp).unwrap();
let ev = bob.send(&apk, b"hello", NOW).unwrap();
let mut carol = SessionManager::new(ident(0x03));
let inv2 = Invite::create_new::<K>(&apk, None).unwrap();
let resp2 = carol.accept_invite(&inv2, None, NOW).unwrap();
let (replacement, _rec) = inv2.receive::<K>(&resp2, &ident(0x01)).unwrap();
let device = alice.devices(&peer).pop().unwrap();
alice.install_session(&peer, &device, replacement);
assert!(alice.process_event(&ev[0]).is_none());
}
#[test]
fn install_and_introspection() {
let mut alice = SessionManager::new(ident(0x50));
let mut bob = SessionManager::new(ident(0x51));
assert_eq!(alice.session_count(), 0);
assert!(!alice.has_session(bob.our_pubkey()));
let invite = Invite::create_new::<K>(alice.our_pubkey(), None).unwrap();
let response = bob.accept_invite(&invite, None, NOW).unwrap();
let peer = alice.receive_invite_response(&invite, &response).unwrap();
assert_eq!(alice.session_count(), 1);
assert!(alice.has_session(&peer));
assert_eq!(alice.peers().count(), 1);
assert_eq!(alice.devices(&peer), vec![bob.our_pubkey().to_owned()]);
}
}