use std::io;
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncRead, AsyncReadExt};
use crate::record::ProviderRecord;
use crate::routing::Contact;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum DhtRequest {
FindNode {
target: String,
},
FindProviders {
content_key: String,
},
AddProvider {
record: ProviderRecord,
},
Ping {
nonce: u64,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum DhtResponse {
Nodes {
nodes: Vec<Contact>,
},
Providers {
providers: Vec<ProviderRecord>,
closer: Vec<Contact>,
},
AddProviderOk,
Pong {
nonce: u64,
},
Error {
code: u32,
message: String,
},
}
pub const MAX_FRAMED_BODY: usize = 256 * 1024;
impl DhtRequest {
pub fn encode(&self) -> Vec<u8> {
encode_framed(self)
}
pub async fn decode<R: AsyncRead + Unpin>(r: &mut R) -> io::Result<Self> {
decode_framed(r).await
}
}
impl DhtResponse {
pub fn encode(&self) -> Vec<u8> {
encode_framed(self)
}
pub async fn decode<R: AsyncRead + Unpin>(r: &mut R) -> io::Result<Self> {
decode_framed(r).await
}
}
fn encode_framed<T: Serialize>(value: &T) -> Vec<u8> {
let body = serde_json::to_vec(value).expect("dht message serializes");
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
}
async fn decode_framed<T: for<'de> Deserialize<'de>, R: AsyncRead + Unpin>(
r: &mut R,
) -> io::Result<T> {
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 > MAX_FRAMED_BODY {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"dht message 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))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::record::{AddressKind, CandidateAddr};
use std::io::Cursor;
#[tokio::test]
async fn find_node_round_trips_framed() {
let req = DhtRequest::FindNode {
target: "ab".repeat(32),
};
let bytes = req.encode();
let mut cur = Cursor::new(bytes);
let back = DhtRequest::decode(&mut cur).await.unwrap();
assert_eq!(req, back);
}
#[tokio::test]
async fn providers_response_round_trips() {
let resp = DhtResponse::Providers {
providers: vec![ProviderRecord {
content_key: "cd".repeat(32),
provider_peer_id: "ef".repeat(32),
addresses: vec![CandidateAddr::direct("203.0.113.7", 9444)],
expires_at: 1_719_763_200,
}],
closer: vec![Contact {
peer_id: "12".repeat(32),
addresses: vec![CandidateAddr {
host: "h".into(),
port: 1,
kind: AddressKind::Mapped,
}],
}],
};
let mut cur = Cursor::new(resp.encode());
let back = DhtResponse::decode(&mut cur).await.unwrap();
assert_eq!(resp, back);
}
#[test]
fn request_type_tags_are_snake_case() {
let s = serde_json::to_string(&DhtRequest::Ping { nonce: 7 }).unwrap();
assert!(s.contains("\"type\":\"ping\""));
assert!(s.contains("\"nonce\":7"));
let fp = serde_json::to_string(&DhtRequest::FindProviders {
content_key: "00".repeat(32),
})
.unwrap();
assert!(fp.contains("\"type\":\"find_providers\""));
}
#[test]
fn response_type_tags_are_snake_case() {
assert!(serde_json::to_string(&DhtResponse::AddProviderOk)
.unwrap()
.contains("\"type\":\"add_provider_ok\""));
assert!(serde_json::to_string(&DhtResponse::Pong { nonce: 9 })
.unwrap()
.contains("\"type\":\"pong\""));
}
#[tokio::test]
async fn oversize_length_prefix_is_rejected() {
let mut buf = ((MAX_FRAMED_BODY + 1) as u32).to_be_bytes().to_vec();
buf.extend_from_slice(b"{}");
let mut cur = Cursor::new(buf);
let err = DhtRequest::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!(DhtRequest::decode(&mut cur).await.is_err());
}
}