use std::io;
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncRead, AsyncReadExt};
use crate::caps::PEX_MAX_FRAME;
use crate::entry::PeerEntry;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum PexMessage {
PexHandshake {
version: u32,
network_id: String,
interval: u32,
#[serde(default)]
flags: Vec<String>,
},
PexSnapshot {
peers: Vec<PeerEntry>,
},
PexDelta {
added: Vec<PeerEntry>,
dropped: Vec<String>,
},
PexError {
code: u16,
message: String,
},
}
impl PexMessage {
#[must_use]
pub fn encode(&self) -> Vec<u8> {
let body = self.to_json_bytes();
let mut out = Vec::with_capacity(4 + body.len());
out.extend_from_slice(&(body.len() as u32).to_be_bytes());
out.extend_from_slice(&body);
out
}
pub async fn decode<R: AsyncRead + Unpin>(r: &mut R) -> io::Result<Self> {
let mut len_buf = [0u8; 4];
r.read_exact(&mut len_buf).await?;
let len = u32::from_be_bytes(len_buf) as usize;
if len > PEX_MAX_FRAME {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"pex frame too large",
));
}
let mut body = vec![0u8; len];
r.read_exact(&mut body).await?;
serde_json::from_slice(&body).map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))
}
#[must_use]
pub fn to_json(&self) -> String {
serde_json::to_string(self).expect("pex message serializes")
}
#[must_use]
pub fn to_json_bytes(&self) -> Vec<u8> {
serde_json::to_vec(self).expect("pex message serializes")
}
pub fn from_json(s: &str) -> Result<Self, serde_json::Error> {
serde_json::from_str(s)
}
#[must_use]
pub fn type_tag(&self) -> &'static str {
match self {
PexMessage::PexHandshake { .. } => "pex_handshake",
PexMessage::PexSnapshot { .. } => "pex_snapshot",
PexMessage::PexDelta { .. } => "pex_delta",
PexMessage::PexError { .. } => "pex_error",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::caps::PEX_VERSION;
use crate::entry::{Address, PeerEntry, Provenance};
use std::io::Cursor;
fn hex(b: u8) -> String {
format!("{b:02x}").repeat(32)
}
fn sample_entry() -> PeerEntry {
PeerEntry::new(hex(0x07), "mainnet", 1_719_763_200, Provenance::Direct)
.with_address(Address::direct("203.0.113.7", 9444))
.with_flag("storage")
}
#[test]
fn handshake_frozen_shape() {
let m = PexMessage::PexHandshake {
version: PEX_VERSION,
network_id: "mainnet".into(),
interval: 60,
flags: vec!["storage".into(), "holepunch".into()],
};
let s = m.to_json();
assert!(s.contains("\"type\":\"pex_handshake\""));
assert!(s.contains("\"version\":1"));
assert!(s.contains("\"network_id\":\"mainnet\""));
assert!(s.contains("\"interval\":60"));
assert!(s.contains("\"flags\":[\"storage\",\"holepunch\"]"));
}
#[test]
fn snapshot_frozen_shape() {
let m = PexMessage::PexSnapshot {
peers: vec![sample_entry()],
};
let s = m.to_json();
assert!(s.contains("\"type\":\"pex_snapshot\""));
assert!(s.contains("\"peers\":["));
}
#[test]
fn delta_frozen_shape() {
let m = PexMessage::PexDelta {
added: vec![sample_entry()],
dropped: vec![hex(0x09)],
};
let s = m.to_json();
assert!(s.contains("\"type\":\"pex_delta\""));
assert!(s.contains("\"added\":["));
assert!(s.contains("\"dropped\":["));
}
#[test]
fn error_frozen_shape() {
let m = PexMessage::PexError {
code: 3,
message: "rate violation".into(),
};
let s = m.to_json();
assert_eq!(
s,
r#"{"type":"pex_error","code":3,"message":"rate violation"}"#
);
}
#[test]
fn all_messages_round_trip_through_json() {
for m in [
PexMessage::PexHandshake {
version: 1,
network_id: "mainnet".into(),
interval: 60,
flags: vec!["storage".into()],
},
PexMessage::PexSnapshot {
peers: vec![sample_entry()],
},
PexMessage::PexDelta {
added: vec![sample_entry()],
dropped: vec![hex(0x09)],
},
PexMessage::PexError {
code: 6,
message: "protocol violation".into(),
},
] {
let back = PexMessage::from_json(&m.to_json()).unwrap();
assert_eq!(m, back);
}
}
#[tokio::test]
async fn framed_round_trip() {
let m = PexMessage::PexSnapshot {
peers: vec![sample_entry()],
};
let bytes = m.encode();
let declared = u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]) as usize;
assert_eq!(declared, bytes.len() - 4);
let mut cur = Cursor::new(bytes);
let back = PexMessage::decode(&mut cur).await.unwrap();
assert_eq!(m, back);
}
#[tokio::test]
async fn oversize_length_prefix_rejected_before_body() {
let mut buf = ((PEX_MAX_FRAME + 1) as u32).to_be_bytes().to_vec();
buf.extend_from_slice(b"{}");
let mut cur = Cursor::new(buf);
let err = PexMessage::decode(&mut cur).await.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[tokio::test]
async fn truncated_frame_errors() {
let mut buf = 100u32.to_be_bytes().to_vec();
buf.extend_from_slice(b"{}");
let mut cur = Cursor::new(buf);
assert!(PexMessage::decode(&mut cur).await.is_err());
}
#[test]
fn unknown_fields_ignored_on_receive() {
let s = r#"{"type":"pex_handshake","version":1,"network_id":"mainnet","interval":60,"future":42}"#;
let m = PexMessage::from_json(s).unwrap();
assert_eq!(m.type_tag(), "pex_handshake");
}
#[test]
fn unknown_type_tag_is_error() {
assert!(PexMessage::from_json(r#"{"type":"pex_bogus"}"#).is_err());
}
#[test]
fn missing_required_field_is_error() {
assert!(PexMessage::from_json(r#"{"type":"pex_snapshot"}"#).is_err());
}
}