use std::fmt;
use std::net::SocketAddr;
use dashmap::DashMap;
use dashmap::mapref::entry::Entry;
use ed25519_dalek::{Signature as Ed25519Signature, Signer, SigningKey, VerifyingKey};
use serde::de::{self, Visitor};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use crate::{Error, Message, MessageId, Payload, Result};
const SIGNING_DOMAIN: &[u8] = b"grapevine.message.v1";
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct PeerId(pub [u8; 32]);
impl PeerId {
pub const UNSIGNED: Self = Self([0u8; 32]);
pub const fn from_bytes(bytes: [u8; 32]) -> Self {
Self(bytes)
}
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
}
impl fmt::Display for PeerId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for byte in &self.0[..8] {
write!(f, "{byte:02x}")?;
}
f.write_str("..")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Signature([u8; 64]);
impl Signature {
pub const UNSIGNED: Self = Self([0u8; 64]);
pub fn as_bytes(&self) -> &[u8; 64] {
&self.0
}
}
impl Serialize for Signature {
fn serialize<S: Serializer>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error> {
serializer.serialize_bytes(&self.0)
}
}
impl<'de> Deserialize<'de> for Signature {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> std::result::Result<Self, D::Error> {
struct SignatureVisitor;
impl<'de> Visitor<'de> for SignatureVisitor {
type Value = Signature;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a 64-byte Ed25519 signature")
}
fn visit_bytes<E: de::Error>(self, value: &[u8]) -> std::result::Result<Signature, E> {
<[u8; 64]>::try_from(value)
.map(Signature)
.map_err(|_| E::invalid_length(value.len(), &self))
}
fn visit_byte_buf<E: de::Error>(
self,
value: Vec<u8>,
) -> std::result::Result<Signature, E> {
self.visit_bytes(&value)
}
fn visit_seq<A: de::SeqAccess<'de>>(
self,
mut seq: A,
) -> std::result::Result<Signature, A::Error> {
let mut bytes = [0u8; 64];
for (index, slot) in bytes.iter_mut().enumerate() {
*slot = seq
.next_element()?
.ok_or_else(|| de::Error::invalid_length(index, &self))?;
}
Ok(Signature(bytes))
}
}
deserializer.deserialize_bytes(SignatureVisitor)
}
}
pub struct Identity {
signing_key: SigningKey,
peer_id: PeerId,
}
impl fmt::Debug for Identity {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Identity")
.field("peer_id", &self.peer_id)
.finish_non_exhaustive()
}
}
impl Identity {
pub fn generate() -> Self {
let seed: [u8; 32] = rand::random();
let signing_key = SigningKey::from_bytes(&seed);
let peer_id = PeerId(signing_key.verifying_key().to_bytes());
Self {
signing_key,
peer_id,
}
}
pub fn peer_id(&self) -> PeerId {
self.peer_id
}
pub fn author(&self, origin: SocketAddr, sequence: u64, payload: Payload) -> Result<Message> {
self.author_with_ttl(origin, sequence, payload, Message::DEFAULT_TTL)
}
pub fn author_with_ttl(
&self,
origin: SocketAddr,
sequence: u64,
payload: Payload,
ttl: u8,
) -> Result<Message> {
let preimage = preimage_bytes(origin, sequence, &payload)?;
let signature = Signature(self.signing_key.sign(&preimage).to_bytes());
Ok(Message {
id: MessageId::new(origin, sequence),
ttl,
payload,
origin_key: self.peer_id,
signature,
})
}
}
pub fn authenticate(message: &Message, pins: &DashMap<SocketAddr, PeerId>) -> Result<()> {
verify_message(message)?;
let origin = message.id.origin;
match pins.entry(origin) {
Entry::Occupied(pinned) if *pinned.get() != message.origin_key => {
Err(Error::OriginKeyMismatch(origin))
}
Entry::Occupied(_) => Ok(()),
Entry::Vacant(slot) => {
slot.insert(message.origin_key);
Ok(())
}
}
}
pub fn verify_message(message: &Message) -> Result<()> {
let origin = message.id.origin;
if message.origin_key == PeerId::UNSIGNED {
return Err(Error::InvalidSignature(origin));
}
let verifying_key = VerifyingKey::from_bytes(&message.origin_key.0)
.map_err(|_| Error::InvalidSignature(origin))?;
let preimage = preimage_bytes(origin, message.id.sequence, &message.payload)?;
let signature = Ed25519Signature::from_bytes(&message.signature.0);
verifying_key
.verify_strict(&preimage, &signature)
.map_err(|_| Error::InvalidSignature(origin))
}
fn preimage_bytes(origin: SocketAddr, sequence: u64, payload: &Payload) -> Result<Vec<u8>> {
#[derive(Serialize)]
struct Preimage<'a> {
domain: &'static [u8],
origin: SocketAddr,
sequence: u64,
payload: &'a Payload,
}
let preimage = Preimage {
domain: SIGNING_DOMAIN,
origin,
sequence,
payload,
};
Ok(bincode::serde::encode_to_vec(
&preimage,
bincode::config::standard(),
)?)
}
#[cfg(test)]
mod tests {
use bytes::Bytes;
use super::*;
fn addr(port: u16) -> SocketAddr {
SocketAddr::from(([127, 0, 0, 1], port))
}
fn app(origin: SocketAddr, sequence: u64, body: &str) -> Message {
Identity::generate()
.author(
origin,
sequence,
Payload::Application(Bytes::from(body.to_owned())),
)
.expect("authoring a well-formed message succeeds")
}
#[test]
fn peer_id_is_the_verifying_key() {
let identity = Identity::generate();
let message = identity
.author(addr(8000), 0, Payload::PeerListRequest)
.unwrap();
assert_eq!(message.origin_key, identity.peer_id());
}
#[test]
fn peer_id_serde() {
let id = Identity::generate().peer_id();
let encoded = bincode::serde::encode_to_vec(id, bincode::config::standard()).unwrap();
let (decoded, _): (PeerId, _) =
bincode::serde::decode_from_slice(&encoded, bincode::config::standard()).unwrap();
assert_eq!(id, decoded);
}
#[test]
fn sign_then_verify() {
let identity = Identity::generate();
let message = identity
.author(
addr(8000),
7,
Payload::Application(Bytes::from_static(b"hi")),
)
.unwrap();
assert!(verify_message(&message).is_ok());
}
#[test]
fn signature_survives_ttl_decrement() {
let identity = Identity::generate();
let mut message = identity
.author(
addr(8000),
1,
Payload::Application(Bytes::from_static(b"x")),
)
.unwrap();
message.decrement_ttl();
message.decrement_ttl();
assert!(
verify_message(&message).is_ok(),
"ttl is excluded from the signature so forwarding cannot break it"
);
}
#[test]
fn tampered_payload_fails_verification() {
let mut message = app(addr(8000), 1, "original");
message.payload = Payload::Application(Bytes::from_static(b"tampered"));
assert!(matches!(
verify_message(&message),
Err(Error::InvalidSignature(_))
));
}
#[test]
fn tampered_origin_fails_verification() {
let mut message = app(addr(8000), 1, "body");
message.id.origin = addr(9999);
assert!(matches!(
verify_message(&message),
Err(Error::InvalidSignature(_))
));
}
#[test]
fn tampered_sequence_fails_verification() {
let mut message = app(addr(8000), 1, "body");
message.id.sequence = 2;
assert!(matches!(
verify_message(&message),
Err(Error::InvalidSignature(_))
));
}
#[test]
fn unsigned_message_is_rejected() {
let message = Message::new(addr(8000), 0, Payload::PeerListRequest);
assert!(matches!(
verify_message(&message),
Err(Error::InvalidSignature(_))
));
}
#[test]
fn swapped_key_fails_verification() {
let mut message = app(addr(8000), 1, "body");
message.origin_key = Identity::generate().peer_id();
assert!(matches!(
verify_message(&message),
Err(Error::InvalidSignature(_))
));
}
#[test]
fn authenticate_pins_first_key_and_rejects_later_changes() {
let pins: DashMap<SocketAddr, PeerId> = DashMap::new();
let origin = addr(8000);
let honest = Identity::generate();
let first = honest
.author(origin, 0, Payload::Application(Bytes::from_static(b"one")))
.unwrap();
assert!(authenticate(&first, &pins).is_ok());
assert_eq!(pins.get(&origin).map(|k| *k), Some(honest.peer_id()));
let second = honest
.author(origin, 1, Payload::Application(Bytes::from_static(b"two")))
.unwrap();
assert!(authenticate(&second, &pins).is_ok());
let forger = Identity::generate();
let forged = forger
.author(
origin,
2,
Payload::Application(Bytes::from_static(b"forged")),
)
.unwrap();
assert!(verify_message(&forged).is_ok());
assert!(matches!(
authenticate(&forged, &pins),
Err(Error::OriginKeyMismatch(o)) if o == origin
));
}
}