use dashmap::DashMap;
use crate::adapter::net::identity::EntityId;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CallerIdentityError {
Unavailable,
OriginMismatch,
}
impl std::fmt::Display for CallerIdentityError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Unavailable => {
write!(f, "no authenticated entity for the inbound session peer")
}
Self::OriginMismatch => write!(
f,
"claimed origin does not match the direct session peer \
(relayed or forged; direct-session-only in v1)"
),
}
}
}
impl std::error::Error for CallerIdentityError {}
pub fn resolve_direct_caller(
peer_entity_ids: &DashMap<u64, EntityId>,
from_node: u64,
claimed_origin_hash: u64,
) -> Result<EntityId, CallerIdentityError> {
if from_node == 0 {
return Err(CallerIdentityError::Unavailable);
}
let caller = peer_entity_ids
.get(&from_node)
.map(|e| e.clone())
.ok_or(CallerIdentityError::Unavailable)?;
if caller.origin_hash() != claimed_origin_hash {
return Err(CallerIdentityError::OriginMismatch);
}
Ok(caller)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::adapter::net::identity::EntityKeypair;
fn entity(seed: u8) -> EntityId {
EntityKeypair::from_bytes([seed; 32]).entity_id().clone()
}
#[test]
fn direct_session_caller_resolves_when_origin_matches() {
let caller = entity(0x11);
let map: DashMap<u64, EntityId> = DashMap::new();
let node_id = 0xABCD;
map.insert(node_id, caller.clone());
let resolved =
resolve_direct_caller(&map, node_id, caller.origin_hash()).expect("direct caller");
assert_eq!(resolved, caller);
}
#[test]
fn unpinned_peer_and_sentinel_are_unavailable() {
let map: DashMap<u64, EntityId> = DashMap::new();
assert_eq!(
resolve_direct_caller(&map, 0x1234, 999),
Err(CallerIdentityError::Unavailable)
);
map.insert(0, entity(0x22));
assert_eq!(
resolve_direct_caller(&map, 0, 999),
Err(CallerIdentityError::Unavailable)
);
}
#[test]
fn forged_origin_is_refused() {
let peer = entity(0x33);
let map: DashMap<u64, EntityId> = DashMap::new();
let node_id = 0x77;
map.insert(node_id, peer.clone());
let forged = peer.origin_hash().wrapping_add(1);
assert_eq!(
resolve_direct_caller(&map, node_id, forged),
Err(CallerIdentityError::OriginMismatch)
);
}
#[test]
fn relayed_request_is_not_mistaken_for_the_caller() {
let relay = entity(0x44);
let caller = entity(0x55);
assert_ne!(relay.origin_hash(), caller.origin_hash());
let map: DashMap<u64, EntityId> = DashMap::new();
let relay_node = 0x99;
map.insert(relay_node, relay.clone());
assert_eq!(
resolve_direct_caller(&map, relay_node, caller.origin_hash()),
Err(CallerIdentityError::OriginMismatch)
);
assert_eq!(
resolve_direct_caller(&map, relay_node, relay.origin_hash()).expect("relay entity"),
relay
);
}
}