Implement fast userspace multihop mesh routing

Replace the one-intermediate-peer relay with protocol-scoped connectivity
graphs and precomputed shortest-path/ECMP forwarding snapshots. Independent
transport readers forward opaque transit frames without a plugin or TUN
round trip. Carry source, destination, a bounded hop limit and stable flow
tags; preserve end-to-end WireGuard links across topology changes.

Classify IP flows before encryption and preserve their tags through the
WireGuard pending queue. Add offline four-agent path-change coverage, loop
and isolation tests, and an opt-in release forwarding microbenchmark. Bump
control/data ALPNs while preserving persistent network identities and state.

Also include the pending Windows Mainline idle-timeout fix and its regression
test, using a reproducible vendored dependency patch.

Validation: fmt, workspace Clippy with warnings denied, and 307 release tests
passed. Two pre-existing Windows SQLite wipe failures were excluded; public
DHT and the manual benchmark remain ignored by default. The forwarding
microbenchmark measured 103 ns (64 B) and 202 ns (1280 B) per transit packet,
excluding encryption and socket I/O.
This commit is contained in:
ab
2026-09-22 18:08:26 +03:00
parent d1a0eca723
commit b4f3e57c8d
50 changed files with 11073 additions and 857 deletions
+147
View File
@@ -0,0 +1,147 @@
# THIS FILE IS AUTOMATICALLY GENERATED BY CARGO
#
# When uploading crates to the registry Cargo will automatically
# "normalize" Cargo.toml files for maximal compatibility
# with all versions of Cargo and also rewrite `path` dependencies
# to registry (e.g., crates.io) dependencies.
#
# If you are reading this file be aware that the original Cargo.toml
# will likely look very different (and much more reasonable).
# See Cargo.toml.orig for the original contents.
[package]
edition = "2021"
rust-version = "1.85"
name = "mainline"
version = "8.0.0"
authors = [
"nuh.dev",
"SeverinAlexB <severin@synonym.to>",
"SHAcollision <shacollision@synonym.to>",
"dzdidi <denys@synonym.to>",
"Kevin Karsopawiro <kevin@synonym.to>",
]
build = false
exclude = [
"/docs/*",
"/examples/*",
]
autolib = false
autobins = false
autoexamples = false
autotests = false
autobenches = false
description = "Simple, robust, BitTorrent's Mainline DHT implementation"
homepage = "https://github.com/pubky/mainline"
readme = "README.md"
keywords = [
"bittorrent",
"torrent",
"dht",
"kademlia",
"mainline",
]
categories = ["network-programming"]
license = "MIT"
repository = "https://github.com/pubky/mainline"
[package.metadata.docs.rs]
all-features = true
[features]
async = [
"node",
"flume/async",
"dep:futures-lite",
]
default = ["full"]
full = ["async"]
node = ["dep:flume"]
[lib]
name = "mainline"
path = "src/lib.rs"
[dependencies.crc]
version = "3.4.0"
default-features = false
[dependencies.document-features]
version = "0.2.12"
[dependencies.dyn-clone]
version = "1.0.20"
default-features = false
[dependencies.ed25519-dalek]
version = "3.0.0-pre.1"
default-features = false
[dependencies.flume]
version = "0.12.0"
optional = true
default-features = false
[dependencies.futures-lite]
version = "2.6.1"
optional = true
default-features = false
[dependencies.getrandom]
version = "0.4"
default-features = false
[dependencies.lru]
version = "0.16.2"
default-features = false
[dependencies.serde]
version = "1.0.228"
features = ["derive"]
[dependencies.serde_bencode]
version = "0.2.4"
default-features = false
[dependencies.serde_bytes]
version = "0.11.19"
[dependencies.sha1_smol]
version = "1.0.1"
default-features = false
[dependencies.thiserror]
version = "2.0.18"
default-features = false
[dependencies.tracing]
version = "0.1.44"
[dev-dependencies.clap]
version = "4.5.57"
features = ["derive"]
[dev-dependencies.colored]
version = "3.1.1"
[dev-dependencies.ctrlc]
version = "3.5.1"
[dev-dependencies.dashmap]
version = "6.1"
[dev-dependencies.flume]
version = "0.12.0"
default-features = false
[dev-dependencies.futures]
version = "0.3.31"
[dev-dependencies.histo]
version = "1.0.0"
[dev-dependencies.rayon]
version = "1.11.0"
[dev-dependencies.tracing-subscriber]
version = "0.3"
+64
View File
@@ -0,0 +1,64 @@
[package]
name = "mainline"
version = "8.0.0"
authors = [
"nuh.dev",
"SeverinAlexB <severin@synonym.to>",
"SHAcollision <shacollision@synonym.to>",
"dzdidi <denys@synonym.to>",
"Kevin Karsopawiro <kevin@synonym.to>"
]
edition = "2021"
rust-version = "1.85"
description = "Simple, robust, BitTorrent's Mainline DHT implementation"
homepage = "https://github.com/pubky/mainline"
license = "MIT"
keywords = ["bittorrent", "torrent", "dht", "kademlia", "mainline"]
categories = ["network-programming"]
repository = "https://github.com/pubky/mainline"
exclude = ["/docs/*", "/examples/*"]
[dependencies]
getrandom = { version = "0.4", default-features = false }
serde_bencode = { version = "0.2.4", default-features = false }
serde = { version = "1.0.228", features = ["derive"] }
serde_bytes = "0.11.19"
thiserror = { version = "2.0.18", default-features = false }
crc = { version = "3.4.0", default-features = false }
sha1_smol = { version = "1.0.1", default-features = false }
ed25519-dalek = { version = "3.0.0-pre.1", default-features = false }
tracing = "0.1.44"
lru = { version = "0.16.2", default-features = false }
dyn-clone = { version = "1.0.20", default-features = false }
document-features = "0.2.12"
# `node` dependencies
flume = { version = "0.12.0", default-features = false, optional = true }
# `async` dependencies
futures-lite = { version = "2.6.1", default-features = false, optional = true }
[dev-dependencies]
clap = { version = "4.5.57", features = ["derive"] }
futures = "0.3.31"
tracing-subscriber = "0.3"
ctrlc = "3.5.1"
histo = "1.0.0"
rayon = "1.11.0"
dashmap = "6.1"
flume = { version = "0.12.0", default-features = false }
colored = "3.1.1"
[features]
## Include [Dht] node.
node = ["dep:flume"]
## Enable [Dht::as_async()] to use [async_dht::AsyncDht].
async = ["node", "flume/async", "dep:futures-lite"]
full = ["async"]
default = ["full"]
[package.metadata.docs.rs]
all-features = true
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2021 raptorswing
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+36
View File
@@ -0,0 +1,36 @@
# Mainline 8.0.0: Windows receive timeout correction
This directory contains the published `mainline` 8.0.0 sources, manifest and
MIT license. The workspace selects this copy through `[patch.crates-io]`.
- Upstream: https://github.com/pubky/mainline
- Published source commit: `b0cabe684f310004c6dcfe8099b91f0d239b11e3`
- Crates.io archive SHA-256:
`d32eaee3dcba6e0bbbefe8bd896a8bd6039d5e74b199c0fe248e9feb547c2a26`
- The upstream manifest is retained, including its development dependencies;
the crate is excluded from the application workspace.
The only source change is in `src/rpc/socket.rs`: `KrpcSocket::recv_from`
treats `ErrorKind::TimedOut` like `WouldBlock`. Its 50 ms socket read timeout
is a normal idle tick, not a failed DHT operation. Windows reports this as
WSAETIMEDOUT (10060), producing approximately 16 warnings per second with
the unmodified crate. Other receive errors still emit the original warning.
Request deadlines, DHT routing, publication and lookup behavior are unchanged.
Rust documents the platform difference here:
https://doc.rust-lang.org/std/net/struct.UdpSocket.html#method.set_read_timeout
Regression coverage is part of the regular workspace suite:
```sh
cargo test --locked -p tsunagi --test mainline_socket
```
It drives the actual DHT actor through idle receive ticks on loopback, checks
that no socket warnings were emitted, then exchanges a real KRPC ping. The
existing Mainline Testnet tests also cover publication, lookup and recovery.
No public DHT, privileged interface or changes to host settings are needed.
When upstream releases this correction, upgrade the dependency and remove
this directory, the patch entry and its workspace exclusion. Keep the
regression test. Do not patch Cargo's global registry cache.
+69
View File
@@ -0,0 +1,69 @@
# Mainline
Simple, robust, BitTorrent's [Mainline](https://en.wikipedia.org/wiki/Mainline_DHT) DHT implementation.
This library is focused on being the best and simplest Rust client for Mainline, especially focused on reliable and fast time-to-first-response.
It should work as a routing / storing node (server mode) as well, and has been running in production for many months without an issue.
However if you are concerned about spam or DoS, you should consider implementing [rate limiting](#rate-limiting).
**[API Docs](https://docs.rs/mainline/latest/mainline/)**
## Getting started
Check the [Examples](https://github.com/Pubky/mainline/tree/main/examples).
## Features
### Client
Running as a client, means you can store and query for values on the DHT, but not accept any incoming requests.
```rust
use mainline::Dht;
let dht = Dht::client().unwrap();
```
Supported BEPs:
- [x] [BEP_0005 DHT Protocol](https://www.bittorrent.org/beps/bep_0005.html)
- [x] [BEP_0042 DHT Security extension](https://www.bittorrent.org/beps/bep_0042.html)
- [x] [BEP_0043 Read-only DHT Nodes](https://www.bittorrent.org/beps/bep_0043.html)
- [x] [BEP_0044 Storing arbitrary data in the DHT](https://www.bittorrent.org/beps/bep_0044.html)
This implementation also includes [measures against Vertical Sybil Attacks](./docs/sybil-resistance.md).
### Server
Running as a server is the same as a client, but you also respond to incoming requests and serve as a routing and storing node, supporting the general routing of the DHT, and contributing to the storage capacity of the DHT.
```rust
use mainline::Dht;
let dht = Dht::server().unwrap(); // or `Dht::builder::server_mode().build();`
```
Supported BEPs:
- [x] [BEP_0005 DHT Protocol](https://www.bittorrent.org/beps/bep_0005.html)
- [x] [BEP_0042 DHT Security extension](https://www.bittorrent.org/beps/bep_0042.html)
- [x] [BEP_0043 Read-only DHT Nodes](https://www.bittorrent.org/beps/bep_0043.html)
- [x] [BEP_0044 Storing arbitrary data in the DHT](https://www.bittorrent.org/beps/bep_0044.html)
#### Rate limiting
The server implementation has no rate-limiting, you can run your own [request filter](./examples/request_filter.rs) and apply your custom rate-limiting.
However, that limit/block will only apply _after_ parsing incoming messages, and it won't affect handling incoming responses.
### Adaptive mode
The default Adaptive mode will start the node in client mode, and after 15 minutes of running with a publicly accessible address,
it will switch to server mode. This way nodes that can serve as routing nodes (accessible and less likely to churn), serve as such.
If you want to explicitly start in Server mode, because you know you are not running behind firewall,
you can call `Dht::builder().server_mode().build()`, and you can optionally add your known public ip so the node doesn't have to depend on,
votes from responding nodes: `Dht::builder().server_mode().public_ip().build()`.
## Acknowledgment
This implementation was possible thanks to [Webtorrent's Bittorrent-dht](https://github.com/webtorrent/bittorrent-dht) as a reference,
and [Rustydht-lib](https://github.com/raptorswing/rustydht-lib) that saved me a lot of time, especially at the serialization and deserialization of Bencode messages.
+1031
View File
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
//! Miscellaneous common structs used throughout the library.
mod id;
mod immutable;
pub mod messages;
mod mutable;
mod node;
mod routing_table;
pub use id::*;
pub use immutable::*;
pub use messages::*;
pub(crate) use mutable::most_recent_mutable_item;
pub use mutable::*;
pub use node::*;
pub use routing_table::*;
+354
View File
@@ -0,0 +1,354 @@
//! Kademlia node Id or a lookup target
use crc::{Crc, CRC_32_ISCSI};
use serde::{Deserialize, Serialize};
use std::convert::TryInto;
use std::{
fmt::{self, Debug, Display, Formatter},
net::{IpAddr, Ipv4Addr, SocketAddr},
str::FromStr,
};
/// The size of node IDs in bits.
pub const ID_SIZE: usize = 20;
pub const MAX_DISTANCE: u8 = ID_SIZE as u8 * 8;
const IPV4_MASK: u32 = 0x030f3fff;
const CASTAGNOLI: Crc<u32> = Crc::<u32>::new(&CRC_32_ISCSI);
#[derive(Clone, Copy, PartialEq, Ord, PartialOrd, Eq, Hash, Serialize, Deserialize)]
/// Kademlia node Id or a lookup target
pub struct Id([u8; ID_SIZE]);
impl Id {
/// Generate a random Id
pub fn random() -> Id {
let mut bytes: [u8; 20] = [0; 20];
getrandom::fill(&mut bytes).expect("getrandom");
Id(bytes)
}
/// Create a new Id from some bytes. Returns Err if the input is not 20 bytes long.
pub fn from_bytes<T: AsRef<[u8]>>(bytes: T) -> Result<Id, InvalidIdSize> {
let bytes = bytes.as_ref();
if bytes.len() != ID_SIZE {
return Err(InvalidIdSize(bytes.len()));
}
let mut tmp: [u8; ID_SIZE] = [0; ID_SIZE];
tmp[..ID_SIZE].clone_from_slice(&bytes[..ID_SIZE]);
Ok(Id(tmp))
}
/// Simplified XOR distance between this Id and a target Id.
///
/// The distance is the number of trailing non zero bits in the XOR result.
///
/// Distance to self is 0
/// Distance to the furthest Id is 160
/// Distance to an Id with 5 leading matching bits is 155
pub fn distance(&self, other: &Id) -> u8 {
MAX_DISTANCE - self.xor(other).leading_zeros()
}
/// Returns the number of leading zeros in the binary representation of `self`.
pub fn leading_zeros(&self) -> u8 {
for (i, byte) in self.0.iter().enumerate() {
if *byte != 0 {
// leading zeros so far + laedinge zeros of this byte
return (i as u32 * 8 + byte.leading_zeros()) as u8;
}
}
160
}
/// Performs bitwise XOR between two Ids
pub fn xor(&self, other: &Id) -> Id {
let mut result = [0_u8; 20];
for (i, (a, b)) in self.0.iter().zip(other.0).enumerate() {
result[i] = a ^ b;
}
result.into()
}
/// Returns a byte slice of this Id.
pub fn as_bytes(&self) -> &[u8; 20] {
&self.0
}
/// Create a new Id according to [BEP_0042](http://bittorrent.org/beps/bep_0042.html).
pub fn from_addr(addr: &SocketAddr) -> Id {
let ip = addr.ip();
Id::from_ip(ip)
}
/// Create a new Id from an Ipv4 address according to [BEP_0042](http://bittorrent.org/beps/bep_0042.html).
pub fn from_ip(ip: IpAddr) -> Id {
match ip {
IpAddr::V4(addr) => Id::from_ipv4(addr),
IpAddr::V6(_addr) => unimplemented!("Ipv6 is not supported"),
}
}
/// Create a new Id from an Ipv4 address according to [BEP_0042](http://bittorrent.org/beps/bep_0042.html).
pub fn from_ipv4(ipv4: Ipv4Addr) -> Id {
let mut bytes = [0_u8; 21];
getrandom::fill(&mut bytes).expect("getrandom");
from_ipv4_and_r(bytes[1..].try_into().expect("infallible"), ipv4, bytes[0])
}
/// Validate that this Id is valid with respect to [BEP_0042](http://bittorrent.org/beps/bep_0042.html).
pub fn is_valid_for_ip(&self, ipv4: Ipv4Addr) -> bool {
if ipv4.is_private() || ipv4.is_link_local() || ipv4.is_loopback() {
return true;
}
let expected = first_21_bits(&id_prefix_ipv4(ipv4, self.0[ID_SIZE - 1]));
self.first_21_bits() == expected
}
pub(crate) fn first_21_bits(&self) -> [u8; 3] {
first_21_bits(&self.0)
}
}
fn first_21_bits(bytes: &[u8]) -> [u8; 3] {
[bytes[0], bytes[1], bytes[2] & 0xf8]
}
fn from_ipv4_and_r(bytes: [u8; 20], ip: Ipv4Addr, r: u8) -> Id {
let mut bytes = bytes;
let prefix = id_prefix_ipv4(ip, r);
// Set first 21 bits to the prefix
bytes[0] = prefix[0];
bytes[1] = prefix[1];
// set the first 5 bits of the 3rd byte to the remaining 5 bits of the prefix
bytes[2] = (prefix[2] & 0xf8) | (bytes[2] & 0x7);
// Set the last byte to the random r
bytes[ID_SIZE - 1] = r;
Id(bytes)
}
fn id_prefix_ipv4(ip: Ipv4Addr, r: u8) -> [u8; 3] {
let r32: u32 = r.into();
let ip_int: u32 = u32::from_be_bytes(ip.octets());
let masked_ip: u32 = (ip_int & IPV4_MASK) | (r32 << 29);
let mut digest = CASTAGNOLI.digest();
digest.update(&masked_ip.to_be_bytes());
let crc = digest.finalize();
crc.to_be_bytes()[..3]
.try_into()
.expect("Failed to convert bytes 0-2 of the crc into a 3-byte array")
}
impl Display for Id {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
#[allow(clippy::format_collect)]
let hex_chars: String = self.0.iter().map(|byte| format!("{:02x}", byte)).collect();
write!(f, "{}", hex_chars)
}
}
impl From<[u8; ID_SIZE]> for Id {
fn from(bytes: [u8; ID_SIZE]) -> Id {
Id(bytes)
}
}
impl From<&[u8; ID_SIZE]> for Id {
fn from(bytes: &[u8; ID_SIZE]) -> Id {
Id(*bytes)
}
}
impl From<Id> for [u8; ID_SIZE] {
fn from(value: Id) -> Self {
value.0
}
}
impl FromStr for Id {
type Err = DecodeIdError;
fn from_str(s: &str) -> Result<Id, DecodeIdError> {
if s.len() % 2 != 0 {
return Err(DecodeIdError::OddNumberOfCharacters);
}
let mut bytes = Vec::with_capacity(s.len() / 2);
for i in 0..s.len() / 2 {
let byte_str = &s[i * 2..(i * 2) + 2];
if let Ok(byte) = u8::from_str_radix(byte_str, 16) {
bytes.push(byte);
} else {
return Err(DecodeIdError::InvalidHexCharacter(byte_str.into()));
}
}
Ok(Id::from_bytes(bytes)?)
}
}
impl Debug for Id {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "Id({})", self)
}
}
#[derive(Debug)]
pub struct InvalidIdSize(usize);
impl std::error::Error for InvalidIdSize {}
impl std::fmt::Display for InvalidIdSize {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "Invalid Id size, expected 20, got {0}", self.0)
}
}
#[derive(thiserror::Error, Debug)]
/// Mainline crate error enum.
pub enum DecodeIdError {
/// Id is expected to by 20 bytes.
#[error(transparent)]
InvalidIdSize(#[from] InvalidIdSize),
#[error("Hex encoding should contain an even number of hex characters")]
/// Hex encoding should contain an even number of hex characters
OddNumberOfCharacters,
/// Invalid hex character
#[error("Invalid Id encoding: {0}")]
InvalidHexCharacter(String),
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn distance_to_self() {
let id = Id::random();
let distance = id.distance(&id);
assert_eq!(distance, 0)
}
#[test]
fn distance_to_id() {
let id = Id::from_str("0639A1E24FBB8AB277DF033476AB0DE10FAB3BDC").unwrap();
let target = Id::from_str("035b1aeb9737ade1a80933594f405d3f772aa08e").unwrap();
let distance = id.distance(&target);
assert_eq!(distance, 155)
}
#[test]
fn distance_to_random_id() {
let id = Id::random();
let target = Id::random();
let distance = id.distance(&target);
assert_ne!(distance, 0)
}
#[test]
fn distance_to_furthest() {
let id = Id::random();
let mut opposite = [0_u8; 20];
for (i, &value) in id.as_bytes().iter().enumerate() {
opposite[i] = value ^ 0xff;
}
let target = Id::from_bytes(opposite).unwrap();
let distance = id.distance(&target);
assert_eq!(distance, MAX_DISTANCE)
}
#[test]
fn from_u8_20() {
let bytes = [8; 20];
let id: Id = bytes.into();
assert_eq!(*id.as_bytes(), bytes);
}
#[test]
fn from_ipv4() {
let vectors = vec![
(Ipv4Addr::new(124, 31, 75, 21), 1, [0x5f, 0xbf, 0xbf]),
(Ipv4Addr::new(21, 75, 31, 124), 86, [0x5a, 0x3c, 0xe9]),
(Ipv4Addr::new(65, 23, 51, 170), 22, [0xa5, 0xd4, 0x32]),
(Ipv4Addr::new(84, 124, 73, 14), 65, [0x1b, 0x03, 0x21]),
(Ipv4Addr::new(43, 213, 53, 83), 90, [0xe5, 0x6f, 0x6c]),
];
for vector in vectors {
test(vector.0, vector.1, vector.2);
}
fn test(ip: Ipv4Addr, r: u8, expected_prefix: [u8; 3]) {
let id = Id::random();
let result = from_ipv4_and_r(*id.as_bytes(), ip, r);
let prefix = first_21_bits(result.as_bytes());
assert_eq!(prefix, first_21_bits(&expected_prefix));
assert_eq!(result.as_bytes()[ID_SIZE - 1], r);
}
}
#[test]
fn is_valid_for_ipv4() {
let valid_vectors = vec![
(
Ipv4Addr::new(124, 31, 75, 21),
"5fbfbff10c5d6a4ec8a88e4c6ab4c28b95eee401",
),
(
Ipv4Addr::new(21, 75, 31, 124),
"5a3ce9c14e7a08645677bbd1cfe7d8f956d53256",
),
(
Ipv4Addr::new(65, 23, 51, 170),
"a5d43220bc8f112a3d426c84764f8c2a1150e616",
),
(
Ipv4Addr::new(84, 124, 73, 14),
"1b0321dd1bb1fe518101ceef99462b947a01ff41",
),
(
Ipv4Addr::new(43, 213, 53, 83),
"e56f6cbf5b7c4be0237986d5243b87aa6d51305a",
),
];
for vector in valid_vectors {
test(vector.0, vector.1);
}
fn test(ip: Ipv4Addr, hex: &str) {
let id = Id::from_str(hex).unwrap();
assert!(id.is_valid_for_ip(ip));
}
}
}
+52
View File
@@ -0,0 +1,52 @@
//! Helper functions for immutable items.
use sha1_smol::Sha1;
use super::ID_SIZE;
use crate::Id;
pub fn validate_immutable(v: &[u8], target: Id) -> bool {
hash_immutable(v) == *target.as_bytes()
}
pub fn hash_immutable(v: &[u8]) -> [u8; ID_SIZE] {
let mut encoded = Vec::with_capacity(v.len() + 3);
encoded.extend(format!("{}:", v.len()).bytes());
encoded.extend_from_slice(v);
let mut hasher = Sha1::new();
hasher.update(&encoded);
hasher.digest().bytes()
}
#[cfg(test)]
mod test {
use super::*;
use std::str::FromStr;
#[test]
fn test_validate_immutable() {
let v = vec![
171, 118, 111, 111, 174, 109, 195, 32, 138, 140, 113, 176, 76, 135, 116, 132, 156, 126,
75, 173,
];
let target = Id::from_bytes([
2, 23, 113, 43, 67, 11, 185, 26, 26, 30, 204, 238, 204, 1, 13, 84, 52, 40, 86, 231,
])
.unwrap();
assert!(validate_immutable(&v, target));
assert!(!validate_immutable(&v[1..], target));
}
#[test]
fn test_hash_immutable() {
let v = b"From the river to the sea, Palestine will be free";
let target = Id::from_str("4238af8aff56cf6e0007d9d2003bf23d33eea7c3").unwrap();
assert_eq!(hash_immutable(v), *target.as_bytes());
}
}
File diff suppressed because it is too large Load Diff
+333
View File
@@ -0,0 +1,333 @@
use serde::{Deserialize, Serialize};
use serde_bytes::ByteBuf;
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTMessage {
#[serde(rename = "t", with = "serde_bytes")]
// Only few messages received seems to not use exactly 2 bytes,
// and they don't seem to have a version.
pub transaction_id: [u8; 4],
#[serde(default)]
#[serde(rename = "v", with = "serde_bytes")]
pub version: Option<[u8; 4]>,
#[serde(flatten)]
pub variant: DHTMessageVariant,
#[serde(default)]
#[serde(with = "serde_bytes")]
// Ipv6 is not supported anyways.
pub ip: Option<[u8; 6]>,
#[serde(default)]
#[serde(rename = "ro")]
pub read_only: Option<i32>,
}
impl DHTMessage {
pub fn from_bytes(bytes: &[u8]) -> Result<DHTMessage, serde_bencode::Error> {
let obj = serde_bencode::from_bytes(bytes)?;
Ok(obj)
}
pub fn to_bytes(&self) -> Result<Vec<u8>, serde_bencode::Error> {
serde_bencode::to_bytes(self)
}
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
#[serde(tag = "y")]
pub enum DHTMessageVariant {
#[serde(rename = "q")]
Request(DHTRequestSpecific),
#[serde(rename = "r")]
Response(DHTResponseSpecific),
#[serde(rename = "e")]
Error(DHTErrorSpecific),
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
#[serde(tag = "q")]
pub enum DHTRequestSpecific {
#[serde(rename = "ping")]
Ping {
#[serde(rename = "a")]
arguments: DHTPingRequestArguments,
},
#[serde(rename = "find_node")]
FindNode {
#[serde(rename = "a")]
arguments: DHTFindNodeRequestArguments,
},
#[serde(rename = "get_peers")]
GetPeers {
#[serde(rename = "a")]
arguments: DHTGetPeersRequestArguments,
},
#[serde(rename = "announce_peer")]
AnnouncePeer {
#[serde(rename = "a")]
arguments: DHTAnnouncePeerRequestArguments,
},
#[serde(rename = "get")]
GetValue {
#[serde(rename = "a")]
arguments: DHTGetValueRequestArguments,
},
#[serde(rename = "put")]
PutValue {
#[serde(rename = "a")]
arguments: DHTPutValueRequestArguments,
},
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
#[serde(untagged)] // This means order matters! Order these from most to least detailed
pub enum DHTResponseSpecific {
GetMutable {
#[serde(rename = "r")]
arguments: DHTGetMutableResponseArguments,
},
NoMoreRecentValue {
#[serde(rename = "r")]
arguments: DHTNoMoreRecentValueResponseArguments,
},
GetImmutable {
#[serde(rename = "r")]
arguments: DHTGetImmutableResponseArguments,
},
GetPeers {
#[serde(rename = "r")]
arguments: DHTGetPeersResponseArguments,
},
NoValues {
#[serde(rename = "r")]
arguments: DHTNoValuesResponseArguments,
},
FindNode {
#[serde(rename = "r")]
arguments: DHTFindNodeResponseArguments,
},
Ping {
#[serde(rename = "r")]
arguments: DHTPingResponseArguments,
},
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTErrorSpecific {
#[serde(rename = "e")]
pub error_info: (i32, String),
}
// === PING ===
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTPingRequestArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTPingResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
}
// === FIND NODE ===
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTFindNodeRequestArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub target: [u8; 20],
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTFindNodeResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub nodes: Box<[u8]>,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTNoValuesResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub nodes: Option<Box<[u8]>>,
}
// === Get Peers ===
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTGetPeersRequestArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub info_hash: [u8; 20],
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTGetPeersResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub nodes: Option<Box<[u8]>>,
// values are not optional, because if they are missing this missing
// we can just treat this as DHTNoValuesResponseArguments
pub values: Vec<ByteBuf>,
}
// === Announce Peer ===
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTAnnouncePeerRequestArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub info_hash: [u8; 20],
pub port: u16,
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(default)]
pub implied_port: Option<u8>,
}
// === Get Value ===
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTGetValueRequestArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub target: [u8; 20],
#[serde(default)]
pub seq: Option<i64>,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTGetImmutableResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub nodes: Option<Box<[u8]>>,
#[serde(with = "serde_bytes")]
pub v: Box<[u8]>,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTNoMoreRecentValueResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub nodes: Option<Box<[u8]>>,
pub seq: i64,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTGetMutableResponseArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub nodes: Option<Box<[u8]>>,
#[serde(with = "serde_bytes")]
pub v: Box<[u8]>,
#[serde(with = "serde_bytes")]
pub k: [u8; 32],
#[serde(with = "serde_bytes")]
pub sig: [u8; 64],
pub seq: i64,
}
// === Put Value ===
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
pub struct DHTPutValueRequestArguments {
#[serde(with = "serde_bytes")]
pub id: [u8; 20],
#[serde(with = "serde_bytes")]
pub target: [u8; 20],
#[serde(with = "serde_bytes")]
pub token: Box<[u8]>,
#[serde(with = "serde_bytes")]
pub v: Box<[u8]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub k: Option<[u8; 32]>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub sig: Option<[u8; 64]>,
#[serde(default)]
pub seq: Option<i64>,
#[serde(default)]
pub cas: Option<i64>,
#[serde(with = "serde_bytes")]
#[serde(default)]
pub salt: Option<Box<[u8]>>,
}
+286
View File
@@ -0,0 +1,286 @@
//! Helper functions and structs for mutable items.
use ed25519_dalek::{Signature, Signer, SigningKey, Verifier, VerifyingKey};
use serde::{Deserialize, Serialize};
use sha1_smol::Sha1;
use std::convert::TryFrom;
use crate::Id;
use super::PutMutableRequestArguments;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
/// [BEP_0044](https://www.bittorrent.org/beps/bep_0044.html)'s Mutable item.
pub struct MutableItem {
/// hash of the key and optional salt
target: Id,
/// ed25519 public key
key: [u8; 32],
/// sequence number
pub(crate) seq: i64,
/// mutable value
pub(crate) value: Box<[u8]>,
/// ed25519 signature
#[serde(with = "serde_bytes")]
signature: [u8; 64],
/// Optional salt
salt: Option<Box<[u8]>>,
}
impl MutableItem {
/// Create a new mutable item from a signing key, value, sequence number and optional salt.
pub fn new(signer: SigningKey, value: &[u8], seq: i64, salt: Option<&[u8]>) -> Self {
let signable = encode_signable(seq, value, salt);
let signature = signer.sign(&signable);
Self::new_signed_unchecked(
signer.verifying_key().to_bytes(),
signature.into(),
value,
seq,
salt,
)
}
/// Return the target of a [MutableItem] by hashing its `public_key` and an optional `salt`
pub fn target_from_key(public_key: &[u8; 32], salt: Option<&[u8]>) -> Id {
let mut encoded = vec![];
encoded.extend(public_key);
if let Some(salt) = salt {
encoded.extend(salt);
}
let mut hasher = Sha1::new();
hasher.update(&encoded);
let bytes = hasher.digest().bytes();
bytes.into()
}
/// Create a new mutable item from an already signed value.
pub fn new_signed_unchecked(
key: [u8; 32],
signature: [u8; 64],
value: &[u8],
seq: i64,
salt: Option<&[u8]>,
) -> Self {
Self {
target: MutableItem::target_from_key(&key, salt),
key,
value: value.into(),
seq,
signature,
salt: salt.map(|s| s.into()),
}
}
pub(crate) fn from_dht_message(
target: Id,
key: &[u8],
v: Box<[u8]>,
seq: i64,
signature: &[u8],
salt: Option<Box<[u8]>>,
) -> Result<Self, MutableError> {
let key = VerifyingKey::try_from(key).map_err(|_| MutableError::InvalidMutablePublicKey)?;
let signature =
Signature::from_slice(signature).map_err(|_| MutableError::InvalidMutableSignature)?;
key.verify(&encode_signable(seq, &v, salt.as_deref()), &signature)
.map_err(|_| MutableError::InvalidMutableSignature)?;
if Self::target_from_key(&key.to_bytes(), salt.as_deref()) != target {
return Err(MutableError::InvalidMutableTarget);
}
Ok(Self {
target,
key: key.to_bytes(),
value: v,
seq,
signature: signature.to_bytes(),
salt,
})
}
// === Getters ===
/// Returns the target (info hash) of this item.
pub fn target(&self) -> &Id {
&self.target
}
/// Returns a reference to the 32 bytes Ed25519 public key of this item.
pub fn key(&self) -> &[u8; 32] {
&self.key
}
/// Returns a byte slice of the value of this item.
pub fn value(&self) -> &[u8] {
&self.value
}
/// Returns the `seq` (sequence) number of this item.
pub fn seq(&self) -> i64 {
self.seq
}
/// Returns the signature over this item.
pub fn signature(&self) -> &[u8; 64] {
&self.signature
}
/// Returns the `Salt` value used for generating the
/// [Self::target] if any.
pub fn salt(&self) -> Option<&[u8]> {
self.salt.as_deref()
}
}
pub fn encode_signable(seq: i64, value: &[u8], salt: Option<&[u8]>) -> Box<[u8]> {
let mut signable = vec![];
if let Some(salt) = salt {
signable.extend(format!("4:salt{}:", salt.len()).into_bytes());
signable.extend(salt);
}
signable.extend(format!("3:seqi{}e1:v{}:", seq, value.len()).into_bytes());
signable.extend(value);
signable.into()
}
pub(crate) fn most_recent_mutable_item(
most_recent: Option<MutableItem>,
item: MutableItem,
) -> Option<MutableItem> {
match most_recent {
Some(mr)
if mr.seq() > item.seq() || (mr.seq() == item.seq() && mr.value() >= item.value()) =>
{
Some(mr)
}
_ => Some(item),
}
}
#[derive(thiserror::Error, Debug)]
/// Mainline crate error enum.
pub enum MutableError {
#[error("Invalid mutable item signature")]
/// Invalid mutable item signature
InvalidMutableSignature,
#[error("Invalid mutable item public key")]
/// Invalid mutable item public key
InvalidMutablePublicKey,
#[error("Mutable item target does not match its public key and salt")]
/// Mutable item target does not match its public key and salt
InvalidMutableTarget,
}
impl PutMutableRequestArguments {
/// Create a [PutMutableRequestArguments] from a [MutableItem],
/// and an optional CAS condition, which is usually the [MutableItem::seq]
/// of the most recent known [MutableItem]
pub fn from(item: MutableItem, cas: Option<i64>) -> Self {
Self {
target: item.target,
v: item.value,
k: item.key,
seq: item.seq,
sig: item.signature,
salt: item.salt,
cas,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use ed25519_dalek::SigningKey;
#[test]
fn signable_without_salt() {
let signable = encode_signable(4, b"Hello world!", None);
assert_eq!(&*signable, b"3:seqi4e1:v12:Hello world!");
}
#[test]
fn signable_with_salt() {
let signable = encode_signable(4, b"Hello world!", Some(b"foobar"));
assert_eq!(&*signable, b"4:salt6:foobar3:seqi4e1:v12:Hello world!");
}
#[test]
fn most_recent_mutable_item_selects_by_seq_then_value() {
let signer = SigningKey::from_bytes(&[
56, 171, 62, 85, 105, 58, 155, 209, 189, 8, 59, 109, 137, 84, 84, 201, 221, 115, 7,
228, 127, 70, 4, 204, 182, 64, 77, 98, 92, 215, 27, 103,
]);
let lower_seq = MutableItem::new(signer.clone(), b"lower-seq", 999, None);
let current = MutableItem::new(signer.clone(), b"current1", 1000, None);
let higher_seq = MutableItem::new(signer.clone(), b"higher-seq", 1001, None);
let same_seq_lower_value = MutableItem::new(signer.clone(), b"current0", 1000, None);
let same_seq_higher_value = MutableItem::new(signer, b"current2", 1000, None);
assert_eq!(
most_recent_mutable_item(None, current.clone()),
Some(current.clone())
);
assert_eq!(
most_recent_mutable_item(Some(current.clone()), higher_seq.clone()),
Some(higher_seq)
);
assert_eq!(
most_recent_mutable_item(Some(current.clone()), lower_seq),
Some(current.clone())
);
assert_eq!(
most_recent_mutable_item(Some(current.clone()), same_seq_higher_value.clone()),
Some(same_seq_higher_value)
);
assert_eq!(
most_recent_mutable_item(Some(current.clone()), same_seq_lower_value),
Some(current.clone())
);
assert_eq!(
most_recent_mutable_item(Some(current.clone()), current.clone()),
Some(current)
);
}
#[test]
fn from_dht_message_rejects_a_signed_item_for_a_different_target() {
let signer = SigningKey::from_bytes(&[42; 32]);
let item = MutableItem::new(signer, b"value", 1, Some(b"salt"));
let mut wrong_target = *item.target().as_bytes();
wrong_target[0] ^= 1;
let result = MutableItem::from_dht_message(
wrong_target.into(),
item.key(),
item.value().into(),
item.seq(),
item.signature(),
item.salt().map(Into::into),
);
assert!(matches!(result, Err(MutableError::InvalidMutableTarget)));
}
}
+142
View File
@@ -0,0 +1,142 @@
//! Struct and implementation of the Node entry in the Kademlia routing table
use std::{
fmt::{self, Debug, Formatter},
net::SocketAddrV4,
sync::Arc,
time::{Duration, Instant},
};
use crate::common::Id;
/// The age of a node's last_seen time before it is considered stale and removed from a full bucket
/// on inserting a new node.
pub const STALE_TIME: Duration = Duration::from_secs(15 * 60);
const MIN_PING_BACKOFF_INTERVAL: Duration = Duration::from_secs(10);
pub const TOKEN_ROTATE_INTERVAL: Duration = Duration::from_secs(60 * 5);
#[derive(PartialEq)]
pub(crate) struct NodeInner {
pub(crate) id: Id,
pub(crate) address: SocketAddrV4,
pub(crate) token: Option<Box<[u8]>>,
pub(crate) last_seen: Instant,
}
impl NodeInner {
pub fn random() -> Self {
Self {
id: Id::random(),
address: SocketAddrV4::new(0.into(), 0),
token: None,
last_seen: Instant::now(),
}
}
}
#[derive(Clone, PartialEq)]
/// Node entry in Kademlia routing table
pub struct Node(pub(crate) Arc<NodeInner>);
impl Debug for Node {
fn fmt(&self, fmt: &mut Formatter<'_>) -> fmt::Result {
fmt.debug_struct("Node")
.field("id", &self.0.id)
.field("address", &self.0.address)
.field("last_seen", &self.0.last_seen.elapsed().as_secs())
.finish()
}
}
impl Node {
/// Creates a new Node from an id and socket address.
pub fn new(id: Id, address: SocketAddrV4) -> Node {
Node(Arc::new(NodeInner {
id,
address,
token: None,
last_seen: Instant::now(),
}))
}
pub(crate) fn new_with_token(id: Id, address: SocketAddrV4, token: Box<[u8]>) -> Self {
Node(Arc::new(NodeInner {
id,
address,
token: Some(token),
last_seen: Instant::now(),
}))
}
/// Creates a node with random Id for testing purposes.
pub fn random() -> Node {
Node(Arc::new(NodeInner::random()))
}
/// Create a node that is unique per `i` as it has a random Id and sets IP and port to `i`
#[cfg(test)]
pub fn unique(i: usize) -> Node {
Node::new(Id::random(), SocketAddrV4::new((i as u32).into(), i as u16))
}
// === Getters ===
/// Returns the id of this node
pub fn id(&self) -> &Id {
&self.0.id
}
/// Returns the address of this node
pub fn address(&self) -> SocketAddrV4 {
self.0.address
}
/// Returns the token we received from this node if any.
pub fn token(&self) -> Option<Box<[u8]>> {
self.0.token.clone()
}
/// Node is last seen more than a threshold ago.
pub fn is_stale(&self) -> bool {
self.0.last_seen.elapsed() > STALE_TIME
}
/// Node's token was received 5 minutes ago or less
pub fn valid_token(&self) -> bool {
self.0.last_seen.elapsed() <= TOKEN_ROTATE_INTERVAL
}
pub(crate) fn should_ping(&self) -> bool {
self.0.last_seen.elapsed() > MIN_PING_BACKOFF_INTERVAL
}
/// Returns true if both nodes have the same ip and port
pub fn same_address(&self, other: &Self) -> bool {
self.0.address == other.0.address
}
/// Returns true if both nodes have the same ip
pub fn same_ip(&self, other: &Self) -> bool {
self.0.address.ip() == other.0.address.ip()
}
/// Node [Id] is valid for its IP address.
///
/// Check [BEP_0042](https://www.bittorrent.org/beps/bep_0042.html).
pub fn is_secure(&self) -> bool {
self.0.id.is_valid_for_ip(*self.0.address.ip())
}
/// Returns true if Any of the existing nodes:
/// - Have the same IP as this node, And:
/// = The existing nodes is Not secure.
/// = The existing nodes is secure And shares the same first 21 bits.
///
/// Effectively, allows only One non-secure node or Eight secure nodes from the same IP, in the routing table or ClosestNodes.
pub(crate) fn already_exists(&self, nodes: &[Self]) -> bool {
nodes.iter().any(|existing| {
self.same_ip(existing)
&& (!existing.is_secure()
|| self.id().first_21_bits() == existing.id().first_21_bits())
})
}
}
+649
View File
@@ -0,0 +1,649 @@
//! Simplified Kademlia routing table
use std::collections::BTreeMap;
use std::slice::Iter;
use crate::common::{Id, Node};
use crate::rpc::ClosestNodes;
/// K = the default maximum size of a k-bucket.
pub const MAX_BUCKET_SIZE_K: usize = 20;
#[derive(Debug, Clone)]
/// Simplified Kademlia routing table
pub struct RoutingTable {
id: Id,
buckets: BTreeMap<u8, KBucket>,
}
impl RoutingTable {
/// Create a new [RoutingTable] with a given id.
pub fn new(id: Id) -> Self {
let buckets = BTreeMap::new();
RoutingTable { id, buckets }
}
/// Returns the [Id] of this node, where the distance is measured from.
pub fn id(&self) -> &Id {
&self.id
}
/// Returns the map of distances and their [KBucket]
pub(crate) fn buckets(&self) -> &BTreeMap<u8, KBucket> {
&self.buckets
}
// === Public Methods ===
/// Attempts to add a node to this routing table, and return `true` if it did.
pub fn add(&mut self, node: Node) -> bool {
let distance = self.id.distance(node.id());
if distance == 0 {
// Do not add self to the routing_table
return false;
}
if self
.buckets()
.values()
.any(|bucket| node.already_exists(&bucket.nodes))
{
return false;
};
let bucket = self.buckets.entry(distance).or_default();
bucket.add(node)
}
/// Remove a node from this routing table.
pub fn remove(&mut self, node_id: &Id) {
let distance = self.id.distance(node_id);
if let Some(bucket) = self.buckets.get_mut(&distance) {
bucket.remove(node_id)
}
}
/// Return the closest nodes to the target while prioritizing secure nodes,
/// as defined in [BEP_0042](https://www.bittorrent.org/beps/bep_0042.html)
pub fn closest(&self, target: Id) -> Box<[Node]> {
let mut closest = ClosestNodes::new(target);
for bucket in self.buckets.values() {
for node in &bucket.nodes {
closest.add(node.clone());
}
}
closest.nodes()[..MAX_BUCKET_SIZE_K.min(closest.len())].into()
}
/// Secure version of [Self::closest] that tries to circumvent sybil attacks.
pub fn closest_secure(
&self,
target: Id,
dht_size_estimate: usize,
subnets: usize,
) -> Vec<Node> {
let mut closest = ClosestNodes::new(target);
for node in self.nodes() {
closest.add(node);
}
closest
.take_until_secure(dht_size_estimate, subnets)
.to_vec()
}
/// Returns `true` if this routing table is empty.
pub fn is_empty(&self) -> bool {
self.buckets.values().all(|bucket| bucket.is_empty())
}
/// Return the number of nodes in this routing table.
pub fn size(&self) -> usize {
self.buckets
.values()
.fold(0, |acc, bucket| acc + bucket.nodes.len())
}
/// Returns an iterator over the nodes in this routing table.
pub fn nodes(&self) -> RoutingTableIterator<'_> {
RoutingTableIterator {
bucket_index: 1,
node_index: 0,
table: self,
}
}
/// Export an owned vector of nodes from this routing table.
pub fn to_owned_nodes(&self) -> Vec<Node> {
self.nodes().collect()
}
/// Turn this routing table to a list of bootstrapping nodes.
pub fn to_bootstrap(&self) -> Vec<String> {
self.nodes()
.filter(|n| !n.is_stale())
.map(|n| n.address().to_string())
.collect()
}
// === Private Methods ===
#[cfg(test)]
fn contains(&self, node_id: &Id) -> bool {
let distance = self.id.distance(node_id);
if let Some(bucket) = self.buckets.get(&distance) {
if bucket.contains(node_id) {
return true;
}
}
false
}
}
pub struct RoutingTableIterator<'a> {
bucket_index: u8,
node_index: usize,
table: &'a RoutingTable,
}
impl Iterator for RoutingTableIterator<'_> {
type Item = Node;
fn next(&mut self) -> Option<Self::Item> {
while self.bucket_index <= 160 {
if let Some(current_bucket) = self.table.buckets.get(&self.bucket_index) {
if let Some(current_node) = current_bucket.nodes.get(self.node_index) {
self.node_index += 1;
if self.node_index == current_bucket.nodes.len() {
self.node_index = 0;
self.bucket_index += 1;
}
return Some(current_node.clone());
}
};
self.bucket_index += 1;
}
None
}
}
/// Kbuckets are similar to LRU caches that checks and evicts unresponsive nodes,
/// without dropping any responsive nodes in the process.
#[derive(Debug, Clone)]
pub struct KBucket {
/// Nodes in the k-bucket, sorted by the least recently seen.
nodes: Vec<Node>,
}
impl KBucket {
pub fn new() -> Self {
KBucket {
nodes: Vec::with_capacity(MAX_BUCKET_SIZE_K),
}
}
// === Getters ===
// === Public Methods ===
pub fn add(&mut self, incoming: Node) -> bool {
if let Some(index) = self.iter().position(|n| n.id() == incoming.id()) {
let existing = self.nodes[index].clone();
// If the incoming node is secure, then we trust its IP address for this Id,
// and even if it changed its port number, we should accept it.
//
// If neither nodes are secure for this Id, but the incoming is the same IP,
// then add the incoming one, effectively updating the node's
// `last_seen` and moving it to the end of the bucket.
// Possibly also updating the port, which is a good thing, instead of waiting
// for the old port to timeout (not responding to Pings).
//
// Using same ip instead of same address, allow
if incoming.is_secure() || (!existing.is_secure() && existing.same_ip(&incoming)) {
self.nodes.remove(index);
self.nodes.push(incoming);
true
} else {
false
}
} else if self.nodes.len() < MAX_BUCKET_SIZE_K {
self.nodes.push(incoming);
true
} else if self.nodes[0].is_stale() {
// Remove the least recently seen node and add the new one
self.nodes.remove(0);
self.nodes.push(incoming);
true
} else {
false
}
}
pub fn remove(&mut self, node_id: &Id) {
self.nodes.retain(|node| node.id() != node_id);
}
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
pub fn iter(&self) -> Iter<'_, Node> {
self.nodes.iter()
}
#[cfg(test)]
fn contains(&self, id: &Id) -> bool {
self.iter().any(|node| node.id() == id)
}
}
impl Default for KBucket {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod test {
use std::net::SocketAddrV4;
use std::str::FromStr;
use std::sync::Arc;
use std::time::Instant;
use crate::common::{Id, KBucket, Node, NodeInner, RoutingTable, MAX_BUCKET_SIZE_K};
#[test]
fn table_is_empty() {
let mut table = RoutingTable::new(Id::random());
assert!(table.is_empty());
table.add(Node::random());
assert!(!table.is_empty());
}
#[test]
fn to_vec() {
let mut table = RoutingTable::new(Id::random());
let mut expected_nodes: Vec<Node> = vec![];
for i in 0..MAX_BUCKET_SIZE_K {
expected_nodes.push(Node::unique(i));
}
for node in &expected_nodes {
table.add(node.clone());
}
let mut sorted_table = table.nodes().collect::<Vec<_>>();
sorted_table.sort_by(|a, b| a.id().cmp(b.id()));
let mut sorted_expected = expected_nodes.to_vec();
sorted_expected.sort_by(|a, b| a.id().cmp(b.id()));
assert_eq!(sorted_table, sorted_expected);
}
#[test]
fn contains() {
let mut table = RoutingTable::new(Id::random());
let node = Node::random();
assert!(!table.contains(node.id()));
table.add(node.clone());
assert!(table.contains(node.id()));
}
#[test]
fn remove() {
let mut table = RoutingTable::new(Id::random());
let node = Node::random();
table.add(node.clone());
assert!(table.contains(node.id()));
table.remove(node.id());
assert!(!table.contains(node.id()));
}
#[test]
fn buckets_are_sets() {
let mut table = RoutingTable::new(Id::random());
let node1 = Node::random();
let node2 = Node::new(*node1.id(), node1.address());
table.add(node1);
table.add(node2);
assert_eq!(table.size(), 1);
}
#[test]
fn should_not_add_self() {
let mut table = RoutingTable::new(Id::random());
let node = Node::new(*table.id(), SocketAddrV4::new(0.into(), 0));
table.add(node.clone());
assert!(!table.add(node));
assert!(table.is_empty())
}
#[test]
fn should_not_add_more_than_k() {
let mut bucket = KBucket::new();
for i in 0..MAX_BUCKET_SIZE_K {
let node = Node::random();
assert!(bucket.add(node), "Failed to add node {}", i);
}
let node = Node::random();
assert!(!bucket.add(node));
}
#[test]
fn should_update_existing_node() {
// Same address
{
let mut bucket = KBucket::new();
let node1 = Node::random();
let node2 = Node::new(*node1.id(), node1.address());
bucket.add(node1.clone());
bucket.add(Node::random());
assert_ne!(bucket.nodes[1].id(), node1.id());
bucket.add(node2);
assert_eq!(bucket.nodes.len(), 2);
assert_eq!(bucket.nodes[1].id(), node1.id());
}
// Different port
{
let mut bucket = KBucket::new();
let node1 = Node::random();
let node2 = Node::new(*node1.id(), SocketAddrV4::new(*node1.address().ip(), 1));
bucket.add(node1.clone());
bucket.add(Node::random());
assert_ne!(bucket.nodes[1].id(), node1.id());
bucket.add(node2.clone());
assert_eq!(bucket.nodes.len(), 2);
assert_eq!(bucket.nodes[1].id(), node1.id());
}
{
let mut bucket = KBucket::new();
let secure = Node(Arc::new(NodeInner {
id: Id::from_str("5a3ce9c14e7a08645677bbd1cfe7d8f956d53256").unwrap(),
address: SocketAddrV4::new([21, 75, 31, 124].into(), 0),
token: None,
last_seen: Instant::now(),
}));
let unsecure = Node::new(*secure.id(), SocketAddrV4::new([0, 0, 0, 0].into(), 1));
{
bucket.add(unsecure.clone());
bucket.add(secure.clone());
assert_eq!(bucket.nodes[0].address(), secure.address())
}
{
bucket.add(secure.clone());
bucket.add(unsecure.clone());
assert_eq!(bucket.nodes[0].address(), secure.address())
}
}
// Different ip
{
let mut bucket = KBucket::new();
let node1 = Node::random();
let node2 = Node::new(*node1.id(), SocketAddrV4::new([0, 0, 0, 1].into(), 1));
bucket.add(node1.clone());
bucket.add(Node::random());
assert_ne!(bucket.nodes[1].id(), node1.id());
bucket.add(node2.clone());
assert_eq!(bucket.nodes.len(), 2);
assert_ne!(bucket.nodes[1].id(), node1.id());
assert_ne!(bucket.nodes[1].address(), node2.address());
}
}
#[test]
fn closest() {
let ids = [
"fb449c17f6c34fadea26a5a83e1952e815e001ea",
"e63b72f95aacee40ad087f83afb475645739f669",
"58c65677e3833cb0f15733a6363cc4cb1352f90a",
"fd042ff1404b495720ad8345404ff5f25acd02a8",
"dbed34a2c8db568fe59c10adcca9e81825b3dcfd",
"079d40b746b5721f59972ebde423429739844914",
"094f1d2fb4b95ba2c3250b014a9f06d13cd9eb9a",
"98805a55523458c56d59339266bdcecc82370ecd",
"0a1d6cce47c60f2c7357e9fec2910192de6eb336",
"fb689ce0e18c2c22f316976d3ae524aed4137773",
"0d01c32b4cf386b0b784b718b999d0e9dac07876",
"9465e80d80f707b222c4ae6ee81c02b62f607629",
"6cdc012328cc7a3a9a5b967e93387686e19c9f75",
"99719dfc220b145e2aac71d6b3e276731d85be1c",
"94d2037bbc534a5f1d672ce3e3350576c2b78ed1",
"b48d0aeb94cd3766f23d2ac098bbccf01485dc20",
"3b6e1c05f199edd7dee87d3cc8422c8f0ed02358",
"d9b50c6ca730c89f8fc9f518136cef6139dd2252",
"15827c92e6efbc4f56e507e548409c4bc04360bf",
"3c8ff1e484c21132f8e6b8112a2feab984536f57",
"c9a8163fa3e85065d46567bfac39b5452cfb3ae8",
"ef79f77e9eed9ad51094ce2747e2c4fdc3a81326",
"81f038cabb8a845f39da0d40716bf0707da55187",
"907fdf0aa137200b395bc210763ed947b03dfc2e",
"b0bce9873042aee29cbc7ec395647f6cc7a482f8",
"e6b8d5567bc05d9b68f23d562645bc030729abc9",
"74667cb7c629fb7e63749134b16e27446984c517",
"cdc7f4d5825dc316de20d998bc0f1c5e91e36a5e",
"701e7b5af5fabcf0bc3de97cb05a7c00da3e53c6",
"36eb09b1db4af2b11312742faa2bb42621fce753",
"9e4923966754c02b036698e95f95cec8fc40a9d2",
"0e43d66e9da1bfc7e2581155dfd1b8f4be57d3f1",
"647679a0d8816d2f62200e7b6ef6171297756dd5",
"c03d9008add37f8414cb41549448bb2dcb5c6c9b",
"dff82b028a6ec033e00b387df8e386417b92a47c",
"42e8b38494b0ee11003592da11b5cbe43332190e",
"03161976385301ac9b965202e8f3922cef840790",
"7d598e5726fb58501d8cc65faf6b676bab7cb4bc",
"54ddde105d3f2c6ea7a5e7641ff24522eea2e784",
"3a75532b5916c772c1b7a18627bf170cf915aeb3",
"fa2b38321419e63cb890f8a8b5c53a1c4728a10a",
"a3ba598bee9da287092f4f2f3864322af38e1824",
"a94df01f21d870a006748b6ab3c04d31428c959d",
"396aabc66c603617f376409053d1e2cec3813101",
"a7b4becc2304da63792eb6c33f95677b2e7c9f8c",
"58b1623af15a9828ccf41b8cee47d123c5cfe8b6",
"3cb7eeac7be3a0195a9243537d452f790ccf1ca9",
"e0296cfc4726d91a1f7f041e24638a1276a08bed",
"aeb03edad3edc7c54a3c5f7916ecba981e65ce91",
"4a81a4596b7c4b8706fd8b5c88ddfde18ca72293",
"0d4e9ae7c486e5a0361bd4e3b918b6bdca89cfcb",
"81d394b44403315f9845c3da6f018b8daedd89ef",
"345630675ff0f319c8f2bb355edf59f9bd93072f",
"b61fbd992a13af05feba939f597b5f6ee61188e3",
"5ea45447e2e79a5f3b3d8c2f68aebdabf71c42f9",
"84325dd6fbd9a93f4ab61d091a9562a6c6111df4",
"e7c796aeecd47cfd01a2d62fd3fb1d41aafa2464",
"897457b33c4eb1ffcab08331877108cbf3fac6de",
"833843b1f33e720c17bccfb75647a49040861b4c",
"06b49c253d3fc9800cfd75605d26426f8ccb89af",
"5024212c42bed9f45e48c450147fecb3e934fc4e",
"5a9de8041b045a7a4f85b71a6dc6a794a7fcd4ea",
"70cad33774ddacb51ed1918adedeb67ff13a3b1e",
"840d201e3c213c01b4ab85983efaac44f0671552",
"aa7ffc7999a1b1bb79ce19b61c37f70331f492d6",
"e2ec0c07e15411564292b5fa75246e4c385f4411",
"38c1a0d14f548d4d81655920ec564b08e9fcf5e6",
"1d128b8343569c7e9a8985879fafd325d458d31c",
"4fcd30cbe02b74cece57babac93aded26ecdc893",
"57d8a6d782ee1df62ceebd5d10884805ed382336",
"54443ed3476d1d542f37bf069973bbd2b64c1b27",
"0e7ba6c5e4c29cf4fff25733892b63cf2a6efdfc",
"18824378226a6d33bcdbe39dd3bc9ee656ce20a2",
"93b0cb01befc90b65a0026acf85bea2fefec7d44",
"d65e378a1ec70cc79ae5b4469ae7f0e8939033fe",
"9230a2f8ac81e73f16c63dd60adb030328fbc983",
"302de797c9d73275ea184d7f6a8bf77364a8fd52",
"cea92f6e6612ef408d8c22ad5c1ed602bb2aedbf",
"353f2ff278f4ee038e7b217276a82d6ed0617130",
"e962e3a1946afa0d3ee97f3a0418cb3489a5f84c",
"a4e42b6cf98e957684aa4e7006940d31bcb76b1f",
"57af8f960b2450ffa0dc5bc7314fece53996d4d0",
"28e73f73084bc8e91fe9ec0a5581b583ef468d8c",
"9481589ddec9a6d9ad2cee7f73e8319aab3f1e95",
"edec09cc7476cd019560874def4af852bfeaffe3",
"6c3ae2cf5f9452d5176788e15635c5958581c931",
"f547b9717e84036c3d5eefec6d6ee3bfa5af89cb",
"87b51f4bf1ccd41cda3aa85c71da5de56aeeda33",
"e743092a576b92c8c05e04d5d2b23f2838825fd1",
"e713b84894b761e2a4e20fd0e5a81ae48a6b6f9d",
"6b1abca34099d2436bac8ab25aa17a57cbfe1564",
"93cb2977e536a680c043b158345254c14b946d52",
"8c2754fa9e93cbf1cccfd9241ebe0cc141199cfe",
"13e4abf95a8a9e6525419b4db7b1704ed0a2789d",
"8d53d453d7cfbb9bc386e128fa68aca388a5ddc6",
"caebf39e9c9b48d87277f2a13faa5931a24819a4",
"5025ca6cda98f31bc3ef321dd9a015b7f06b8bfa",
"531fbf18fdf3e513091614f20d65e920a505ca41",
"2f81e6159f7de0bc90c8a1db661b33bffbee85fd",
"85d4d9954f3a28228a2786b320ad58a46a13f37b",
];
let nodes: Vec<Node> = ids
.iter()
.enumerate()
.map(|(i, str)| {
let id = Id::from_str(str).unwrap();
Node(Arc::new(NodeInner {
id,
address: SocketAddrV4::new((i as u32).into(), i as u16),
token: None,
last_seen: Instant::now(),
}))
})
.collect();
let local_id = Id::from_str("ba3042eb2d373b19e7c411ce6826e31b37be0b2e").unwrap();
let mut table = RoutingTable::new(local_id);
for node in nodes {
table.add(node);
}
{
let expected_closest_ids: Vec<_> = [
"897457b33c4eb1ffcab08331877108cbf3fac6de",
"907fdf0aa137200b395bc210763ed947b03dfc2e",
"9230a2f8ac81e73f16c63dd60adb030328fbc983",
"93b0cb01befc90b65a0026acf85bea2fefec7d44",
"93cb2977e536a680c043b158345254c14b946d52",
"9465e80d80f707b222c4ae6ee81c02b62f607629",
"9481589ddec9a6d9ad2cee7f73e8319aab3f1e95",
"94d2037bbc534a5f1d672ce3e3350576c2b78ed1",
"98805a55523458c56d59339266bdcecc82370ecd",
"99719dfc220b145e2aac71d6b3e276731d85be1c",
"9e4923966754c02b036698e95f95cec8fc40a9d2",
"a3ba598bee9da287092f4f2f3864322af38e1824",
"a4e42b6cf98e957684aa4e7006940d31bcb76b1f",
"a7b4becc2304da63792eb6c33f95677b2e7c9f8c",
"a94df01f21d870a006748b6ab3c04d31428c959d",
"aa7ffc7999a1b1bb79ce19b61c37f70331f492d6",
"aeb03edad3edc7c54a3c5f7916ecba981e65ce91",
"b0bce9873042aee29cbc7ec395647f6cc7a482f8",
"b48d0aeb94cd3766f23d2ac098bbccf01485dc20",
"b61fbd992a13af05feba939f597b5f6ee61188e3",
]
.iter()
.map(|id| Id::from_str(id).unwrap())
.collect();
let target = local_id;
let closest = table.closest(target);
let mut closest_ids: Vec<Id> = closest.iter().map(|n| *n.id()).collect();
closest_ids.sort();
assert_eq!(closest_ids, expected_closest_ids);
}
{
let expected_closest_ids: Vec<_> = [
"c03d9008add37f8414cb41549448bb2dcb5c6c9b",
"c9a8163fa3e85065d46567bfac39b5452cfb3ae8",
"cdc7f4d5825dc316de20d998bc0f1c5e91e36a5e",
"cea92f6e6612ef408d8c22ad5c1ed602bb2aedbf",
"d65e378a1ec70cc79ae5b4469ae7f0e8939033fe",
"d9b50c6ca730c89f8fc9f518136cef6139dd2252",
"dbed34a2c8db568fe59c10adcca9e81825b3dcfd",
"dff82b028a6ec033e00b387df8e386417b92a47c",
"e0296cfc4726d91a1f7f041e24638a1276a08bed",
"e2ec0c07e15411564292b5fa75246e4c385f4411",
"e63b72f95aacee40ad087f83afb475645739f669",
"e6b8d5567bc05d9b68f23d562645bc030729abc9",
"e7c796aeecd47cfd01a2d62fd3fb1d41aafa2464",
"e962e3a1946afa0d3ee97f3a0418cb3489a5f84c",
"edec09cc7476cd019560874def4af852bfeaffe3",
"ef79f77e9eed9ad51094ce2747e2c4fdc3a81326",
"fa2b38321419e63cb890f8a8b5c53a1c4728a10a",
"fb449c17f6c34fadea26a5a83e1952e815e001ea",
"fb689ce0e18c2c22f316976d3ae524aed4137773",
"fd042ff1404b495720ad8345404ff5f25acd02a8",
]
.iter()
.map(|str| Id::from_str(str).unwrap())
.collect();
let target = Id::from_str("d1406a3d3a8354d566f21dba8bd06c537cde2a20").unwrap();
let closest = table.closest(target);
let mut closest_ids: Vec<Id> = closest.iter().map(|n| *n.id()).collect();
closest_ids.sort();
assert_eq!(closest_ids, expected_closest_ids);
}
}
}
+1363
View File
File diff suppressed because it is too large Load Diff
+44
View File
@@ -0,0 +1,44 @@
#![doc = include_str!("../README.md")]
//! ## Feature flags
#![doc = document_features::document_features!()]
//!
#![deny(missing_docs)]
#![deny(rustdoc::broken_intra_doc_links)]
#![cfg_attr(not(test), deny(clippy::unwrap_used))]
mod common;
#[cfg(feature = "node")]
mod dht;
mod rpc;
// Public modules
#[cfg(feature = "async")]
pub mod async_dht;
pub use common::{Id, MutableItem, Node, RoutingTable};
#[cfg(feature = "node")]
pub use dht::{Dht, DhtBuilder, Testnet, TestnetBuilder};
#[cfg(feature = "node")]
pub use rpc::{
config::Config,
messages::{MessageType, PutRequestSpecific, RequestSpecific},
server::{RequestFilter, ServerSettings, MAX_INFO_HASHES, MAX_PEERS, MAX_VALUES},
ClosestNodes, GetMutableOutcome, PutOutcome, DEFAULT_BOOTSTRAP_NODES, DEFAULT_REQUEST_TIMEOUT,
};
pub use ed25519_dalek::SigningKey;
pub mod errors {
//! Exported errors
#[cfg(feature = "node")]
pub use super::common::ErrorSpecific;
#[cfg(feature = "node")]
pub use super::dht::PutMutableError;
#[cfg(feature = "node")]
pub use super::rpc::{ConcurrencyError, PutError, PutQueryError};
pub use super::common::DecodeIdError;
pub use super::common::MutableError;
}
+1261
View File
File diff suppressed because it is too large Load Diff
+297
View File
@@ -0,0 +1,297 @@
use std::{collections::HashSet, convert::TryInto};
use crate::{common::MAX_BUCKET_SIZE_K, Id, Node};
#[derive(Debug, Clone)]
/// Manage closest nodes found in a query.
///
/// Useful to estimate the Dht size.
pub struct ClosestNodes {
target: Id,
nodes: Vec<Node>,
}
impl ClosestNodes {
/// Create a new instance of [ClosestNodes].
pub fn new(target: Id) -> Self {
Self {
target,
nodes: Vec::with_capacity(200),
}
}
// === Getters ===
/// Returns the target of the query for these closest nodes.
pub fn target(&self) -> Id {
self.target
}
/// Returns a slice of the nodes array.
pub fn nodes(&self) -> &[Node] {
&self.nodes
}
/// Returns the number of nodes.
pub fn len(&self) -> usize {
self.nodes.len()
}
/// Returns true if there are no nodes.
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
// === Public Methods ===
/// Add a node.
pub fn add(&mut self, node: Node) {
let seek = node.id().xor(&self.target);
if node.already_exists(&self.nodes) {
return;
}
if let Err(pos) = self.nodes.binary_search_by(|prope| {
if prope.is_secure() && !node.is_secure() {
std::cmp::Ordering::Less
} else if !prope.is_secure() && node.is_secure() {
std::cmp::Ordering::Greater
} else if prope.id() == node.id() {
std::cmp::Ordering::Equal
} else {
prope.id().xor(&self.target).cmp(&seek)
}
}) {
self.nodes.insert(pos, node)
}
}
/// Take enough nodes closest to the target, until the following are satisfied:
/// 1. At least the closest `k` nodes (20).
/// 2. The last node should be at a distance `edk` which is the expected distance of the 20th
/// node given previous estimations of the DHT size.
/// 3. The number of subnets with unique 6 bits prefix in nodes ipv4 addresses match or exceeds
/// the average from previous queries.
///
/// If one or more of these conditions are not met, then we just take all responding nodes
/// and store data at them.
pub fn take_until_secure(
&self,
previous_dht_size_estimate: usize,
average_subnets: usize,
) -> &[Node] {
let mut until_secure = 0;
// 20 / dht_size_estimate == expected_dk / ID space
// so expected_dk = 20 * ID space / dht_size_estimate
let expected_dk =
(20.0 * u128::MAX as f64 / (previous_dht_size_estimate as f64 + 1.0)) as u128;
let mut subnets = HashSet::new();
for node in &self.nodes {
let distance = distance(&self.target, node);
subnets.insert(subnet(node));
if distance >= expected_dk && subnets.len() >= average_subnets {
break;
}
until_secure += 1;
}
&self.nodes[0..until_secure.max(MAX_BUCKET_SIZE_K).min(self.nodes().len())]
}
/// Count the number of subnets with unique 6 bits prefix in ipv4
pub fn subnets_count(&self) -> u8 {
if self.nodes.is_empty() {
return 20;
}
let mut subnets = HashSet::new();
for node in self.nodes.iter().take(MAX_BUCKET_SIZE_K) {
subnets.insert(subnet(node));
}
subnets.len() as u8
}
/// An estimation of the Dht from the distribution of closest nodes
/// responding to a query.
///
/// [Read more](https://github.com/pubky/mainline/blob/main/docs/dht_size_estimate.md)
pub fn dht_size_estimate(&self) -> f64 {
dht_size_estimate(
self.nodes
.iter()
.take(MAX_BUCKET_SIZE_K)
.map(|node| distance(&self.target, node)),
)
}
}
fn subnet(node: &Node) -> u8 {
((node.address().ip().to_bits() >> 26) & 0b0011_1111) as u8
}
fn distance(target: &Id, node: &Node) -> u128 {
let xor = node.id().xor(target);
// Round up the lower 4 bytes to get a u128 from u160.
u128::from_be_bytes(xor.as_bytes()[0..16].try_into().expect("infallible"))
}
fn dht_size_estimate<I>(distances: I) -> f64
where
I: IntoIterator<Item = u128>,
{
let mut sum = 0.0;
let mut count = 0;
// Ignoring the first node, as that gives the best result in simulations.
for distance in distances {
count += 1;
sum += count as f64 * distance as f64;
}
if count == 0 {
return 0.0;
}
let lsq_constant = (count * (count + 1) * (2 * count + 1) / 6) as f64;
lsq_constant * u128::MAX as f64 / sum
}
#[cfg(test)]
mod tests {
use std::{collections::BTreeMap, net::SocketAddrV4, str::FromStr, sync::Arc, time::Instant};
use crate::common::NodeInner;
use super::*;
#[test]
fn add_sorted_by_id() {
let target = Id::random();
let mut closest_nodes = ClosestNodes::new(target);
for i in 0..100 {
let node = Node::unique(i);
closest_nodes.add(node.clone());
closest_nodes.add(node);
}
assert_eq!(closest_nodes.nodes().len(), 100);
let distances = closest_nodes
.nodes()
.iter()
.map(|n| n.id().distance(&target))
.collect::<Vec<_>>();
let mut sorted = distances.clone();
sorted.sort();
assert_eq!(sorted, distances);
}
#[test]
fn order_by_secure_id() {
let unsecure = Node::random();
let secure = Node(Arc::new(NodeInner {
id: Id::from_str("5a3ce9c14e7a08645677bbd1cfe7d8f956d53256").unwrap(),
address: SocketAddrV4::new([21, 75, 31, 124].into(), 0),
token: None,
last_seen: Instant::now(),
}));
let mut closest_nodes = ClosestNodes::new(*unsecure.id());
closest_nodes.add(unsecure.clone());
closest_nodes.add(secure.clone());
assert_eq!(closest_nodes.nodes(), vec![secure, unsecure])
}
#[test]
fn take_until_expected_distance_to_20th_node() {
let target = Id::random();
let dht_size_estimate = 200;
let mut closest_nodes = ClosestNodes::new(target);
let target_bytes = target.as_bytes();
for i in 0..dht_size_estimate {
let node = Node::unique(i);
closest_nodes.add(node);
}
let mut sybil = ClosestNodes::new(target);
for _ in 0..20 {
let mut bytes = target_bytes.to_vec();
bytes[18..].copy_from_slice(&Id::random().as_bytes()[18..]);
let node = Node::new(Id::random(), SocketAddrV4::new(0.into(), 0));
sybil.add(node.clone());
closest_nodes.add(node);
}
let closest = closest_nodes.take_until_secure(dht_size_estimate, 0);
assert!((closest.len() - sybil.nodes().len()) > 10);
}
#[test]
fn simulation() {
let lookups = 4;
let acceptable_margin = 0.2;
let sims = 10;
let dht_size = 2500_f64;
let mean = (0..sims)
.map(|_| simulate(dht_size as usize, lookups) as f64)
.sum::<f64>()
/ (sims as f64);
let margin = (mean - dht_size).abs() / dht_size;
assert!(margin <= acceptable_margin);
}
fn simulate(dht_size: usize, lookups: usize) -> usize {
let mut nodes = BTreeMap::new();
for i in 0..dht_size {
let node = Node::unique(i);
nodes.insert(*node.id(), node);
}
(0..lookups)
.map(|_| {
let target = Id::random();
let mut closest_nodes = ClosestNodes::new(target);
for (_, node) in nodes.range(target..).take(100) {
closest_nodes.add(node.clone())
}
for (_, node) in nodes.range(..target).rev().take(100) {
closest_nodes.add(node.clone())
}
let estimate = closest_nodes.dht_size_estimate();
estimate as usize
})
.sum::<usize>()
/ lookups
}
}
+56
View File
@@ -0,0 +1,56 @@
use std::{
net::{Ipv4Addr, SocketAddrV4},
time::Duration,
};
use super::{ServerSettings, DEFAULT_REQUEST_TIMEOUT};
#[derive(Debug, Clone)]
/// Dht Configurations
pub struct Config {
/// Bootstrap nodes
///
/// Defaults to [super::DEFAULT_BOOTSTRAP_NODES]
pub bootstrap: Option<Vec<SocketAddrV4>>,
/// Explicit port to listen on.
///
/// Defaults to None
pub port: Option<u16>,
/// UDP socket request timeout duration.
///
/// The longer this duration is, the longer queries take until they are deemeed "done".
/// The shortet this duration is, the more responses from busy nodes we miss out on,
/// which affects the accuracy of queries trying to find closest nodes to a target.
///
/// Defaults to [DEFAULT_REQUEST_TIMEOUT]
pub request_timeout: Duration,
/// Server to respond to incoming Requests
pub server_settings: ServerSettings,
/// Whether or not to start in server mode from the get go.
///
/// Defaults to false where it will run in [Adaptive mode](https://github.com/pubky/mainline?tab=readme-ov-file#adaptive-mode).
pub server_mode: bool,
/// A known public IPv4 address for this node to generate
/// a secure node Id from according to [BEP_0042](https://www.bittorrent.org/beps/bep_0042.html)
///
/// Defaults to None, where we depend on suggestions from responding nodes.
pub public_ip: Option<Ipv4Addr>,
/// Address to bind to.
///
/// Defaults to 0.0.0.0 (all interfaces)
pub bind_address: Option<Ipv4Addr>,
}
impl Default for Config {
fn default() -> Self {
Self {
bootstrap: None,
port: None,
request_timeout: DEFAULT_REQUEST_TIMEOUT,
server_settings: Default::default(),
server_mode: false,
public_ip: None,
bind_address: None,
}
}
}
+70
View File
@@ -0,0 +1,70 @@
use std::net::SocketAddrV4;
use crate::Id;
use super::Rpc;
/// Information and statistics about this mainline node.
#[derive(Debug, Clone)]
pub struct Info {
id: Id,
local_addr: SocketAddrV4,
public_address: Option<SocketAddrV4>,
firewalled: bool,
dht_size_estimate: (usize, f64),
server_mode: bool,
}
impl Info {
/// This Node's [Id]
pub fn id(&self) -> &Id {
&self.id
}
/// Local UDP Ipv4 socket address that this node is listening on.
pub fn local_addr(&self) -> SocketAddrV4 {
self.local_addr
}
/// Returns the best guess for this node's Public address.
///
/// If [crate::DhtBuilder::public_ip] was set, this is what will be returned
/// (plus the local port), otherwise it will rely on consensus from
/// responding nodes voting on our public IP and port.
pub fn public_address(&self) -> Option<SocketAddrV4> {
self.public_address
}
/// Returns `true` if we can't confirm that [Self::public_address] is publicly addressable.
///
/// If this node is firewalled, it won't switch to server mode if it is in adaptive mode,
/// but if [crate::DhtBuilder::server_mode] was set to true, then whether or not this node is firewalled
/// won't matter.
pub fn firewalled(&self) -> bool {
self.firewalled
}
/// Returns whether or not this node is running in server mode.
pub fn server_mode(&self) -> bool {
self.server_mode
}
/// Returns:
/// 1. Normal Dht size estimate based on all closer `nodes` in query responses.
/// 2. Standard deviaiton as a function of the number of samples used in this estimate.
///
/// [Read more](https://github.com/pubky/mainline/blob/main/docs/dht_size_estimate.md)
pub fn dht_size_estimate(&self) -> (usize, f64) {
self.dht_size_estimate
}
}
impl From<&Rpc> for Info {
fn from(rpc: &Rpc) -> Self {
Self {
id: *rpc.id(),
local_addr: rpc.local_addr(),
dht_size_estimate: rpc.dht_size_estimate(),
public_address: rpc.public_address(),
firewalled: rpc.firewalled(),
server_mode: rpc.server_mode(),
}
}
}
+310
View File
@@ -0,0 +1,310 @@
//! Manage iterative queries and their corresponding request/response.
use std::collections::HashMap;
use std::collections::HashSet;
use std::net::SocketAddrV4;
use tracing::{debug, trace};
use super::{socket::KrpcSocket, ClosestNodes};
use crate::common::{FindNodeRequestArguments, GetPeersRequestArguments, GetValueRequestArguments};
use crate::{
common::{Id, Node, RequestSpecific, RequestTypeSpecific, MAX_BUCKET_SIZE_K},
rpc::Response,
};
/// Aggregate diagnostics for a mutable GET query.
#[derive(Debug, Default, Clone, PartialEq, Eq)]
pub struct GetMutableOutcome {
/// Number of unique DHT nodes queried.
pub queried: u32,
/// Number of valid mutable values returned.
pub values: u32,
/// Number of `NoValues` responses returned.
pub no_values: u32,
/// Number of `NoMoreRecentValue` responses returned.
pub no_more_recent: u32,
/// Number of mutable value responses that failed validation.
pub invalid_values: u32,
/// Number of invalid response shapes returned.
pub invalid_responses: u32,
/// Number of KRPC error responses returned.
pub krpc_errors: u32,
}
impl GetMutableOutcome {
/// Return the number of nodes that returned a GET response before timing out.
pub fn responded(&self) -> u32 {
self.valid_responses() + self.invalid_values + self.invalid_responses + self.krpc_errors
}
/// Return the number of nodes that returned a valid GET response.
pub fn valid_responses(&self) -> u32 {
self.values + self.no_values + self.no_more_recent
}
/// Return the number of queried nodes that did not return a GET response before timeout.
pub fn timed_out(&self) -> u32 {
self.queried.saturating_sub(self.responded())
}
fn record_value(&mut self) {
self.values += 1;
}
fn record_no_values(&mut self) {
self.no_values += 1;
}
fn record_no_more_recent(&mut self) {
self.no_more_recent += 1;
}
fn record_invalid_value(&mut self) {
self.invalid_values += 1;
}
fn record_invalid_response(&mut self) {
self.invalid_responses += 1;
}
fn record_krpc_error(&mut self) {
self.krpc_errors += 1;
}
fn finish(mut self, queried: u32) -> Self {
self.queried = queried;
self
}
}
/// An iterative process of concurrently sending a request to the closest known nodes to
/// the target, updating the routing table with closer nodes discovered in the responses, and
/// repeating this process until no closer nodes (that aren't already queried) are found.
#[derive(Debug)]
pub(crate) struct IterativeQuery {
pub request: RequestSpecific,
closest: ClosestNodes,
responders: ClosestNodes,
inflight_requests: Vec<u32>,
query_requests: Vec<u32>,
visited: HashSet<SocketAddrV4>,
responses: Vec<Response>,
mutable_outcome: GetMutableOutcome,
public_address_votes: HashMap<SocketAddrV4, u16>,
}
#[derive(Debug)]
pub enum GetRequestSpecific {
FindNode(FindNodeRequestArguments),
GetPeers(GetPeersRequestArguments),
GetValue(GetValueRequestArguments),
}
impl GetRequestSpecific {
pub fn target(&self) -> &Id {
match self {
GetRequestSpecific::FindNode(args) => &args.target,
GetRequestSpecific::GetPeers(args) => &args.info_hash,
GetRequestSpecific::GetValue(args) => &args.target,
}
}
}
impl IterativeQuery {
pub fn new(requester_id: Id, target: Id, request: GetRequestSpecific) -> Self {
let request_type = match request {
GetRequestSpecific::FindNode(s) => RequestTypeSpecific::FindNode(s),
GetRequestSpecific::GetPeers(s) => RequestTypeSpecific::GetPeers(s),
GetRequestSpecific::GetValue(s) => RequestTypeSpecific::GetValue(s),
};
trace!(?target, ?request_type, "New Query");
Self {
request: RequestSpecific {
requester_id,
request_type,
},
closest: ClosestNodes::new(target),
responders: ClosestNodes::new(target),
inflight_requests: Vec::new(),
query_requests: Vec::new(),
visited: HashSet::new(),
responses: Vec::new(),
mutable_outcome: GetMutableOutcome::default(),
public_address_votes: HashMap::new(),
}
}
// === Getters ===
pub fn target(&self) -> Id {
self.responders.target()
}
/// Closest nodes according to other nodes.
pub fn closest(&self) -> &ClosestNodes {
&self.closest
}
/// Return the closest responding nodes after the query is done.
pub fn responders(&self) -> &ClosestNodes {
&self.responders
}
pub fn responses(&self) -> &[Response] {
&self.responses
}
pub fn mutable_outcome(&self) -> GetMutableOutcome {
self.mutable_outcome
.clone()
.finish(self.visited.len() as u32)
}
pub fn best_address(&self) -> Option<SocketAddrV4> {
let mut max = 0_u16;
let mut best_addr = None;
for (addr, count) in self.public_address_votes.iter() {
if *count > max {
max = *count;
best_addr = Some(*addr);
};
}
best_addr
}
// === Public Methods ===
/// Force start query traversal by visiting closest nodes.
pub fn start(&mut self, socket: &mut KrpcSocket) {
self.visit_closest(socket);
}
/// Add a candidate node to query on next tick if it is among the closest nodes.
pub fn add_candidate(&mut self, node: Node) {
// ready for a ipv6 routing table?
self.closest.add(node);
}
/// Add a vote for this node's address.
pub fn add_address_vote(&mut self, address: SocketAddrV4) {
self.public_address_votes
.entry(address)
.and_modify(|counter| *counter += 1)
.or_insert(1);
}
/// Visit explicitly given addresses, and add them to the visited set.
/// only used from the Rpc when calling bootstrapping nodes.
pub fn visit(&mut self, socket: &mut KrpcSocket, address: SocketAddrV4) {
let tid = socket.request(address, self.request.clone());
self.inflight_requests.push(tid);
self.query_requests.push(tid);
let tid = socket.request(
address,
RequestSpecific {
requester_id: Id::random(),
request_type: RequestTypeSpecific::Ping,
},
);
self.inflight_requests.push(tid);
self.visited.insert(address);
}
/// Return true if a response (by transaction_id) is expected by this query.
pub fn is_inflight(&self, tid: u32) -> bool {
self.inflight_requests.contains(&tid)
}
/// Return true if the transaction belongs to the primary query request, not the liveness ping.
pub fn is_inflight_query_request(&self, tid: u32) -> bool {
self.query_requests.contains(&tid)
}
/// Add a node that responded with a token as a probable storage node.
pub fn add_responding_node(&mut self, node: Node) {
self.responders.add(node)
}
/// Store received response.
pub fn response(&mut self, from: SocketAddrV4, response: Response) {
let target = self.target();
debug!(?target, ?response, ?from, "Query got response");
self.responses.push(response.to_owned());
}
pub fn record_mutable_value(&mut self) {
self.mutable_outcome.record_value();
}
pub fn record_no_values(&mut self) {
self.mutable_outcome.record_no_values();
}
pub fn record_no_more_recent(&mut self) {
self.mutable_outcome.record_no_more_recent();
}
pub fn record_invalid_response(&mut self) {
self.mutable_outcome.record_invalid_response();
}
pub fn record_krpc_error(&mut self) {
self.mutable_outcome.record_krpc_error();
}
pub fn record_invalid_mutable_value(&mut self) {
self.mutable_outcome.record_invalid_value();
}
/// Query closest nodes for this query's target and message.
///
/// Returns true if it is done.
pub fn tick(&mut self, socket: &mut KrpcSocket) -> bool {
// Visit closest nodes
self.visit_closest(socket);
// If no more inflight_requests are inflight in the socket (not timed out),
// then the query is done.
let done = !self
.inflight_requests
.iter()
.any(|&tid| socket.inflight(tid));
if done {
debug!(id=?self.target(), closest = ?self.closest.len(), visited = ?self.visited.len(), responders = ?self.responders.len(), "Done query");
};
done
}
// === Private Methods ===
/// Visit the closest candidates and remove them as candidates
fn visit_closest(&mut self, socket: &mut KrpcSocket) {
let to_visit = self
.closest
.nodes()
.iter()
.take(MAX_BUCKET_SIZE_K)
.filter(|node| !self.visited.contains(&node.address()))
.map(|node| node.address())
.collect::<Vec<_>>();
for address in to_visit {
self.visit(socket, address);
}
}
}
+309
View File
@@ -0,0 +1,309 @@
use tracing::{debug, trace};
use crate::{
common::{
ErrorSpecific, Id, PutRequest, PutRequestSpecific, RequestSpecific, RequestTypeSpecific,
},
Node,
};
use super::socket::KrpcSocket;
/// Stores data at the closest nodes after an [super::IterativeQuery] is done,
/// or when a previous cached query is available.
///
/// Tracks successful acknowledgements and errors for the PUT query.
#[derive(Debug)]
pub struct PutQuery {
pub target: Id,
/// Nodes that confirmed success
stored_at: u32,
inflight_requests: Vec<u32>,
pub request: PutRequestSpecific,
errors: Vec<(u8, ErrorSpecific)>,
extra_nodes: Box<[Node]>,
}
impl PutQuery {
pub fn new(target: Id, request: PutRequestSpecific, extra_nodes: Option<Box<[Node]>>) -> Self {
Self {
target,
stored_at: 0,
inflight_requests: Vec::new(),
request,
errors: Vec::new(),
extra_nodes: extra_nodes.unwrap_or_default(),
}
}
pub fn start(
&mut self,
socket: &mut KrpcSocket,
closest_nodes: &[Node],
) -> Result<(), PutError> {
assert!(!self.started(), "should not call PutQuery::start() twice");
let target = self.target;
trace!(?target, "PutQuery start");
if closest_nodes.is_empty() {
Err(PutQueryError::NoClosestNodes)?;
}
assert!(
closest_nodes.len() <= u8::MAX as usize,
"should not send PUT query to more than 256 nodes"
);
for node in closest_nodes.iter().chain(self.extra_nodes.iter()) {
// Set correct values to the request placeholders
if let Some(token) = node.token() {
let tid = socket.request(
node.address(),
RequestSpecific {
requester_id: Id::random(),
request_type: RequestTypeSpecific::Put(PutRequest {
token,
put_request_type: self.request.clone(),
}),
},
);
self.inflight_requests.push(tid);
}
}
Ok(())
}
pub fn started(&self) -> bool {
!self.inflight_requests.is_empty()
}
pub fn inflight(&self, tid: u32) -> bool {
self.inflight_requests.contains(&tid)
}
pub fn success(&mut self) {
debug!(target = ?self.target, "PutQuery got success response");
self.stored_at += 1
}
pub fn error(&mut self, error: ErrorSpecific) {
debug!(target = ?self.target, ?error, "PutQuery got error");
if let Some(pos) = self
.errors
.iter()
.position(|(_, err)| error.code == err.code)
{
// Increment the count of the existing error
self.errors[pos].0 += 1;
// Move the updated element to maintain the order (highest count first)
let mut i = pos;
while i > 0 && self.errors[i].0 > self.errors[i - 1].0 {
self.errors.swap(i, i - 1);
i -= 1;
}
} else {
// Add the new error with a count of 1
self.errors.push((1, error));
}
}
/// Check if the query has completed, returning the PUT outcome when complete.
pub fn poll_completion(&self, socket: &KrpcSocket) -> Result<Option<PutOutcome>, PutError> {
if !self.started() {
return Ok(None);
}
if let Some(most_common_error) = self.majority_nodes_rejected_put_mutable() {
debug!(
target = ?self.target,
?most_common_error,
nodes_count = self.inflight_requests.len(),
"PutQuery for MutableItem was rejected by most nodes with 3xx code."
);
return Err(PutError::from(most_common_error));
}
// And all queries got responses or timed out.
if self.is_done(socket) {
let target = self.target;
if self.stored_at == 0 {
let most_common_error = self.most_common_error();
debug!(
?target,
?most_common_error,
nodes_count = self.inflight_requests.len(),
"Put Query: failed"
);
return Err(most_common_error
.map(|(_, error)| error)
.unwrap_or(PutQueryError::Timeout.into()));
}
debug!(?target, stored_at = ?self.stored_at, "PutQuery Done successfully");
return Ok(Some(PutOutcome {
target: self.target,
stored_at: self.stored_at,
}));
}
Ok(None)
}
fn is_done(&self, socket: &KrpcSocket) -> bool {
self.inflight_requests
.iter()
.copied()
.all(|transaction_id| !socket.inflight(transaction_id))
}
fn majority_nodes_rejected_put_mutable(&self) -> Option<ConcurrencyError> {
if !matches!(self.request, PutRequestSpecific::PutMutable(_)) {
return None;
}
let (count, error) = self.most_common_error()?;
let half = ((self.inflight_requests.len() / 2) + 1) as u8;
if count < half {
return None;
}
match error {
PutError::Concurrency(error) => Some(error),
PutError::Query(_) => None,
}
}
fn most_common_error(&self) -> Option<(u8, PutError)> {
self.errors
.first()
.and_then(|(count, error)| match error.code {
301 => Some((*count, PutError::from(ConcurrencyError::CasFailed))),
302 => Some((*count, PutError::from(ConcurrencyError::NotMostRecent))),
_ => None,
})
}
}
/// Result details for a successful PUT query.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PutOutcome {
/// DHT target the request was published under.
pub target: Id,
/// Number of DHT nodes that acknowledged storing the item.
pub stored_at: u32,
}
/// PutQuery errors
#[derive(thiserror::Error, Debug, Clone)]
pub enum PutError {
/// Common PutQuery errors
#[error(transparent)]
Query(#[from] PutQueryError),
#[error(transparent)]
/// PutQuery for [crate::MutableItem] errors
Concurrency(#[from] ConcurrencyError),
}
/// Common PutQuery errors
#[derive(thiserror::Error, Debug, Clone)]
pub enum PutQueryError {
/// Failed to find any nodes close, usually means dht node failed to bootstrap,
/// so the routing table is empty. Check the machine's access to UDP socket,
/// or find better bootstrapping nodes.
#[error("Failed to find any nodes close to store value at")]
NoClosestNodes,
/// Either Put Query failed to store at any nodes, and most nodes responded
/// with a non `301` nor `302` errors.
///
/// Either way; contains the most common error response.
#[error("Query Error Response")]
ErrorResponse(ErrorSpecific),
/// PutQuery timed out with no responses neither success or errors
#[error("PutQuery timed out with no responses neither success or errors")]
Timeout,
}
/// PutQuery for [crate::MutableItem] errors
#[derive(thiserror::Error, Debug, Clone)]
pub enum ConcurrencyError {
/// Trying to PUT mutable items with the same `key`, and `salt` but different `seq`.
///
/// Moreover, the more recent item does _NOT_ mention the the earlier
/// item's `seq` in its `cas` field.
///
/// This risks a [Lost Update Problem](https://en.wikipedia.org/wiki/Write-write_conflict).
///
/// Try reading most recent mutable item before writing again,
/// and make sure to set the `cas` field.
#[error("Conflict risk, try reading most recent item before writing again.")]
ConflictRisk,
/// The [crate::MutableItem::seq] is less than or equal the sequence from another signed item.
///
/// Try reading most recent mutable item before writing again.
#[error("MutableItem::seq is not the most recent, try reading most recent item before writing again.")]
NotMostRecent,
/// The `CAS` condition does not match the `seq` of the most recent known signed item.
#[error("CAS check failed, try reading most recent item before writing again.")]
CasFailed,
}
#[cfg(test)]
mod tests {
use crate::{
common::{PutMutableRequestArguments, PutRequestSpecific},
MutableItem, SigningKey,
};
use super::{ConcurrencyError, PutError, PutQuery};
use crate::common::ErrorSpecific;
use crate::rpc::socket::KrpcSocket;
#[test]
fn mutable_majority_cas_failure_wins_over_completed_success() {
let signer = SigningKey::from_bytes(&[
56, 171, 62, 85, 105, 58, 155, 209, 189, 8, 59, 109, 137, 84, 84, 201, 221, 115, 7,
228, 127, 70, 4, 204, 182, 64, 77, 98, 92, 215, 27, 103,
]);
let item = MutableItem::new(signer, b"value", 1002, None);
let request = PutRequestSpecific::PutMutable(PutMutableRequestArguments::from(
item.clone(),
Some(1000),
));
let mut query = PutQuery::new(*item.target(), request, None);
query.inflight_requests = vec![1, 2, 3];
query.success();
query.error(cas_failed());
query.error(cas_failed());
let socket = KrpcSocket::client().unwrap();
assert!(matches!(
query.poll_completion(&socket),
Err(PutError::Concurrency(ConcurrencyError::CasFailed))
));
}
fn cas_failed() -> ErrorSpecific {
ErrorSpecific {
code: 301,
description: "cas failed".to_string(),
}
}
}
+421
View File
@@ -0,0 +1,421 @@
//! Modules needed only for nodes running in server mode (not read-only).
pub mod peers;
pub mod tokens;
use std::{fmt::Debug, net::SocketAddrV4, num::NonZeroUsize};
use dyn_clone::DynClone;
use lru::LruCache;
use tracing::debug;
use crate::common::{
validate_immutable, AnnouncePeerRequestArguments, ErrorSpecific, FindNodeRequestArguments,
FindNodeResponseArguments, GetImmutableResponseArguments, GetMutableResponseArguments,
GetPeersRequestArguments, GetPeersResponseArguments, GetValueRequestArguments, Id, MutableItem,
NoMoreRecentValueResponseArguments, NoValuesResponseArguments, PingResponseArguments,
PutImmutableRequestArguments, PutMutableRequestArguments, PutRequest, PutRequestSpecific,
RequestTypeSpecific, ResponseSpecific, RoutingTable,
};
use peers::PeersStore;
use tokens::Tokens;
pub use crate::common::{MessageType, RequestSpecific};
/// Default maximum number of info_hashes for which to store peers.
pub const MAX_INFO_HASHES: usize = 2000;
/// Default maximum number of peers to store per info_hash.
pub const MAX_PEERS: usize = 500;
/// Default maximum number of Immutable and Mutable items to store.
pub const MAX_VALUES: usize = 1000;
/// A trait for filtering incoming requests to a DHT node and
/// decide whether to allow handling it or rate limit or ban
/// the requester, or prohibit specific requests' details.
pub trait RequestFilter: Send + Sync + Debug + DynClone {
/// Returns true if the request from this source is allowed.
fn allow_request(&self, request: &RequestSpecific, from: SocketAddrV4) -> bool;
}
dyn_clone::clone_trait_object!(RequestFilter);
#[derive(Debug, Clone)]
struct DefaultFilter;
impl RequestFilter for DefaultFilter {
fn allow_request(&self, _request: &RequestSpecific, _from: SocketAddrV4) -> bool {
true
}
}
#[derive(Debug)]
/// A server that handles incoming requests.
///
/// Supports [BEP_005](https://www.bittorrent.org/beps/bep_0005.html) and [BEP_0044](https://www.bittorrent.org/beps/bep_0044.html).
///
/// But it doesn't implement any rate-limiting or blocking.
pub struct Server {
/// Tokens generator
tokens: Tokens,
/// Peers store
peers: PeersStore,
/// Immutable values store
immutable_values: LruCache<Id, Box<[u8]>>,
/// Mutable values store
mutable_values: LruCache<Id, MutableItem>,
/// Filter requests before handling them.
filter: Box<dyn RequestFilter>,
}
impl Default for Server {
fn default() -> Self {
Self::new(ServerSettings::default())
}
}
#[derive(Debug, Clone)]
/// Settings for the default dht server.
pub struct ServerSettings {
/// The maximum info_hashes for which to store peers.
///
/// Defaults to [MAX_INFO_HASHES]
pub max_info_hashes: usize,
/// The maximum peers to store per info_hash.
///
/// Defaults to [MAX_PEERS]
pub max_peers_per_info_hash: usize,
/// Maximum number of immutable values to store.
///
/// Defaults to [MAX_VALUES]
pub max_immutable_values: usize,
/// Maximum number of mutable values to store.
///
/// Defaults to [MAX_VALUES]
pub max_mutable_values: usize,
/// Filter requests before handling them.
///
/// Defaults to a function that always returns true.
pub filter: Box<dyn RequestFilter>,
}
impl Default for ServerSettings {
fn default() -> Self {
Self {
max_info_hashes: MAX_INFO_HASHES,
max_peers_per_info_hash: MAX_PEERS,
max_mutable_values: MAX_VALUES,
max_immutable_values: MAX_VALUES,
filter: Box::new(DefaultFilter),
}
}
}
impl Server {
/// Creates a new [Server]
pub fn new(settings: ServerSettings) -> Self {
let tokens = Tokens::new();
Self {
tokens,
peers: PeersStore::new(
NonZeroUsize::new(settings.max_info_hashes).unwrap_or(
NonZeroUsize::new(MAX_INFO_HASHES).expect("MAX_PEERS is NonZeroUsize"),
),
NonZeroUsize::new(settings.max_peers_per_info_hash)
.unwrap_or(NonZeroUsize::new(MAX_PEERS).expect("MAX_PEERS is NonZeroUsize")),
),
immutable_values: LruCache::new(
NonZeroUsize::new(settings.max_immutable_values)
.unwrap_or(NonZeroUsize::new(MAX_VALUES).expect("MAX_VALUES is NonZeroUsize")),
),
mutable_values: LruCache::new(
NonZeroUsize::new(settings.max_mutable_values)
.unwrap_or(NonZeroUsize::new(MAX_VALUES).expect("MAX_VALUES is NonZeroUsize")),
),
filter: settings.filter,
}
}
/// Returns an optional response or an error for a request.
///
/// Passed to the Rpc to send back to the requester.
pub fn handle_request(
&mut self,
routing_table: &RoutingTable,
from: SocketAddrV4,
request: RequestSpecific,
) -> Option<MessageType> {
if !self.filter.allow_request(&request, from) {
return None;
}
// Lazily rotate secrets before handling a request
if self.tokens.should_update() {
self.tokens.rotate()
}
let requester_id = request.requester_id;
Some(match request.request_type {
RequestTypeSpecific::Ping => {
MessageType::Response(ResponseSpecific::Ping(PingResponseArguments {
responder_id: *routing_table.id(),
}))
}
RequestTypeSpecific::FindNode(FindNodeRequestArguments { target, .. }) => {
MessageType::Response(ResponseSpecific::FindNode(FindNodeResponseArguments {
responder_id: *routing_table.id(),
nodes: routing_table.closest(target),
}))
}
RequestTypeSpecific::GetPeers(GetPeersRequestArguments { info_hash, .. }) => {
MessageType::Response(match self.peers.get_random_peers(&info_hash) {
Some(peers) => ResponseSpecific::GetPeers(GetPeersResponseArguments {
responder_id: *routing_table.id(),
token: self.tokens.generate_token(from).into(),
nodes: Some(routing_table.closest(info_hash)),
values: peers,
}),
None => ResponseSpecific::NoValues(NoValuesResponseArguments {
responder_id: *routing_table.id(),
token: self.tokens.generate_token(from).into(),
nodes: Some(routing_table.closest(info_hash)),
}),
})
}
RequestTypeSpecific::GetValue(GetValueRequestArguments { target, seq, .. }) => {
if seq.is_some() {
MessageType::Response(self.handle_get_mutable(routing_table, from, target, seq))
} else if let Some(v) = self.immutable_values.get(&target) {
MessageType::Response(ResponseSpecific::GetImmutable(
GetImmutableResponseArguments {
responder_id: *routing_table.id(),
token: self.tokens.generate_token(from).into(),
nodes: Some(routing_table.closest(target)),
v: v.clone(),
},
))
} else {
MessageType::Response(self.handle_get_mutable(routing_table, from, target, seq))
}
}
RequestTypeSpecific::Put(PutRequest {
token,
put_request_type,
}) => match put_request_type {
PutRequestSpecific::AnnouncePeer(AnnouncePeerRequestArguments {
info_hash,
port,
implied_port,
..
}) => {
if !self.tokens.validate(from, &token) {
debug!(
?info_hash,
?requester_id,
?from,
request_type = "announce_peer",
"Invalid token"
);
return Some(MessageType::Error(ErrorSpecific {
code: 203,
description: "Bad token".to_string(),
}));
}
let peer = match implied_port {
Some(true) => from,
_ => SocketAddrV4::new(*from.ip(), port),
};
self.peers
.add_peer(info_hash, (&request.requester_id, peer));
return Some(MessageType::Response(ResponseSpecific::Ping(
PingResponseArguments {
responder_id: *routing_table.id(),
},
)));
}
PutRequestSpecific::PutImmutable(PutImmutableRequestArguments {
v,
target,
..
}) => {
if !self.tokens.validate(from, &token) {
debug!(
?target,
?requester_id,
?from,
request_type = "put_immutable",
"Invalid token"
);
return Some(MessageType::Error(ErrorSpecific {
code: 203,
description: "Bad token".to_string(),
}));
}
if v.len() > 1000 {
debug!(?target, ?requester_id, ?from, size = ?v.len(), "Message (v field) too big.");
return Some(MessageType::Error(ErrorSpecific {
code: 205,
description: "Message (v field) too big.".to_string(),
}));
}
if !validate_immutable(&v, target) {
debug!(?target, ?requester_id, ?from, v = ?v, "Target doesn't match the sha1 hash of v field.");
return Some(MessageType::Error(ErrorSpecific {
code: 203,
description: "Target doesn't match the sha1 hash of v field"
.to_string(),
}));
}
self.immutable_values.put(target, v);
return Some(MessageType::Response(ResponseSpecific::Ping(
PingResponseArguments {
responder_id: *routing_table.id(),
},
)));
}
PutRequestSpecific::PutMutable(PutMutableRequestArguments {
target,
v,
k,
seq,
sig,
salt,
cas,
..
}) => {
if !self.tokens.validate(from, &token) {
debug!(
?target,
?requester_id,
?from,
request_type = "put_mutable",
"Invalid token"
);
return Some(MessageType::Error(ErrorSpecific {
code: 203,
description: "Bad token".to_string(),
}));
}
if v.len() > 1000 {
return Some(MessageType::Error(ErrorSpecific {
code: 205,
description: "Message (v field) too big.".to_string(),
}));
}
if let Some(ref salt) = salt {
if salt.len() > 64 {
return Some(MessageType::Error(ErrorSpecific {
code: 207,
description: "salt (salt field) too big.".to_string(),
}));
}
}
if let Some(previous) = self.mutable_values.get(&target) {
if let Some(cas) = cas {
if previous.seq() != cas {
debug!(
?target,
?requester_id,
?from,
"CAS mismatched, re-read value and try again."
);
return Some(MessageType::Error(ErrorSpecific {
code: 301,
description: "CAS mismatched, re-read value and try again."
.to_string(),
}));
}
};
if seq < previous.seq() {
debug!(
?target,
?requester_id,
?from,
"Sequence number less than current."
);
return Some(MessageType::Error(ErrorSpecific {
code: 302,
description: "Sequence number less than current.".to_string(),
}));
}
}
match MutableItem::from_dht_message(target, &k, v, seq, &sig, salt) {
Ok(item) => {
self.mutable_values.put(target, item);
MessageType::Response(ResponseSpecific::Ping(PingResponseArguments {
responder_id: *routing_table.id(),
}))
}
Err(error) => {
debug!(?target, ?requester_id, ?from, ?error, "Invalid signature");
MessageType::Error(ErrorSpecific {
code: 206,
description: "Invalid signature".to_string(),
})
}
}
}
},
})
}
/// Handle get mutable request
fn handle_get_mutable(
&mut self,
routing_table: &RoutingTable,
from: SocketAddrV4,
target: Id,
seq: Option<i64>,
) -> ResponseSpecific {
match self.mutable_values.get(&target) {
Some(item) => {
let no_more_recent_values = seq.map(|request_seq| item.seq() <= request_seq);
match no_more_recent_values {
Some(true) => {
ResponseSpecific::NoMoreRecentValue(NoMoreRecentValueResponseArguments {
responder_id: *routing_table.id(),
token: self.tokens.generate_token(from).into(),
nodes: Some(routing_table.closest(target)),
seq: item.seq(),
})
}
_ => ResponseSpecific::GetMutable(GetMutableResponseArguments {
responder_id: *routing_table.id(),
token: self.tokens.generate_token(from).into(),
nodes: Some(routing_table.closest(target)),
v: item.value().into(),
k: *item.key(),
seq: item.seq(),
sig: *item.signature(),
}),
}
}
None => ResponseSpecific::NoValues(NoValuesResponseArguments {
responder_id: *routing_table.id(),
token: self.tokens.generate_token(from).into(),
nodes: Some(routing_table.closest(target)),
}),
}
}
}
+176
View File
@@ -0,0 +1,176 @@
//! Manage announced peers for info_hashes
use std::{net::SocketAddrV4, num::NonZeroUsize};
use crate::common::Id;
use lru::LruCache;
const CHANCE_SCALE: f32 = 2.0 * (1u32 << 31) as f32;
#[derive(Debug, Clone)]
/// An LRU cache of "Peers" per info hashes.
///
/// Read [BEP_0005](https://www.bittorrent.org/beps/bep_0005.html) for more information.
pub struct PeersStore {
info_hashes: LruCache<Id, LruCache<Id, SocketAddrV4>>,
max_peers: NonZeroUsize,
}
impl PeersStore {
/// Create a new store of peers announced on info hashes.
pub fn new(max_info_hashes: NonZeroUsize, max_peers: NonZeroUsize) -> Self {
Self {
info_hashes: LruCache::new(max_info_hashes),
max_peers,
}
}
/// Add a peer for an info hash.
pub fn add_peer(&mut self, info_hash: Id, peer: (&Id, SocketAddrV4)) {
if let Some(info_hash_lru) = self.info_hashes.get_mut(&info_hash) {
info_hash_lru.put(*peer.0, peer.1);
} else {
let mut info_hash_lru = LruCache::new(self.max_peers);
info_hash_lru.put(*peer.0, peer.1);
self.info_hashes.put(info_hash, info_hash_lru);
};
}
/// Returns a random set of peers per an info hash.
pub fn get_random_peers(&mut self, info_hash: &Id) -> Option<Vec<SocketAddrV4>> {
if let Some(info_hash_lru) = self.info_hashes.get(info_hash) {
let size = info_hash_lru.len();
let target_size = 10;
if size == 0 {
return None;
}
if size < target_size {
return Some(
info_hash_lru
.iter()
.map(|n| n.1.to_owned())
.collect::<Vec<_>>(),
);
}
let mut results = Vec::with_capacity(10);
let mut chunk = vec![0_u8; info_hash_lru.iter().len() * 4];
getrandom::fill(chunk.as_mut_slice()).expect("getrandom");
for (index, (_, addr)) in info_hash_lru.iter().enumerate() {
// Calculate the chance of adding the current item based on remaining items and slots
let remaining_slots = target_size - results.len();
let remaining_items = info_hash_lru.len() - index;
let current_chance =
((remaining_slots as f32 / remaining_items as f32) * CHANCE_SCALE) as u32;
// Get random integer from the chunk
let rand_int =
u32::from_le_bytes(chunk[index..index + 4].try_into().expect("infallible"));
// Randomly decide to add the item based on the current chance
if rand_int < current_chance {
results.push(*addr);
if results.len() == target_size {
break;
}
}
}
return Some(results);
}
None
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn max_info_hashes() {
let mut store = PeersStore::new(
NonZeroUsize::new(1).unwrap(),
NonZeroUsize::new(100).unwrap(),
);
let info_hash_a = Id::random();
let info_hash_b = Id::random();
store.add_peer(
info_hash_a,
(&info_hash_a, SocketAddrV4::new([127, 0, 1, 1].into(), 0)),
);
store.add_peer(
info_hash_b,
(&info_hash_b, SocketAddrV4::new([127, 0, 1, 1].into(), 0)),
);
assert_eq!(store.info_hashes.len(), 1);
assert_eq!(
store.get_random_peers(&info_hash_b),
Some([SocketAddrV4::new([127, 0, 1, 1].into(), 0)].into())
);
}
#[test]
fn all_peers() {
let mut store =
PeersStore::new(NonZeroUsize::new(1).unwrap(), NonZeroUsize::new(2).unwrap());
let info_hash_a = Id::random();
let info_hash_b = Id::random();
let info_hash_c = Id::random();
store.add_peer(
info_hash_a,
(&info_hash_a, SocketAddrV4::new([127, 0, 1, 1].into(), 0)),
);
store.add_peer(
info_hash_a,
(&info_hash_b, SocketAddrV4::new([127, 0, 1, 2].into(), 0)),
);
store.add_peer(
info_hash_a,
(&info_hash_c, SocketAddrV4::new([127, 0, 1, 3].into(), 0)),
);
assert_eq!(
store.get_random_peers(&info_hash_a),
Some(
[
SocketAddrV4::new([127, 0, 1, 3].into(), 0),
SocketAddrV4::new([127, 0, 1, 2].into(), 0),
]
.into()
)
);
}
#[test]
fn random_peers_subset() {
let mut store = PeersStore::new(
NonZeroUsize::new(1).unwrap(),
NonZeroUsize::new(200).unwrap(),
);
let info_hash = Id::random();
for i in 0..200 {
store.add_peer(
info_hash,
(&Id::random(), SocketAddrV4::new([127, 0, 1, i].into(), 0)),
)
}
assert_eq!(store.info_hashes.get(&info_hash).unwrap().len(), 200);
let sample = store.get_random_peers(&info_hash).unwrap();
assert_eq!(sample.len(), 10);
}
}
+119
View File
@@ -0,0 +1,119 @@
//! Manage tokens for remote client IPs.
use crc::{Crc, CRC_32_ISCSI};
use std::{
fmt::{self, Debug, Formatter},
net::SocketAddrV4,
time::Instant,
};
use tracing::trace;
const SECRET_SIZE: usize = 20;
const TOKEN_SIZE: usize = 4;
const CASTAGNOLI: Crc<u32> = Crc::<u32>::new(&CRC_32_ISCSI);
/// Tokens generator.
///
/// Read [BEP_0005](https://www.bittorrent.org/beps/bep_0005.html) for more information.
#[derive(Clone)]
pub struct Tokens {
prev_secret: [u8; SECRET_SIZE],
curr_secret: [u8; SECRET_SIZE],
last_updated: Instant,
}
impl Debug for Tokens {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
write!(f, "Tokens (_)")
}
}
impl Tokens {
/// Create a Tokens generator.
pub fn new() -> Self {
Tokens {
prev_secret: random(),
curr_secret: random(),
last_updated: Instant::now(),
}
}
// === Public Methods ===
/// Returns `true` if the current secret needs to be updated after an interval.
pub fn should_update(&self) -> bool {
self.last_updated.elapsed() > crate::common::TOKEN_ROTATE_INTERVAL
}
/// Validate that the token was generated within the past 10 minutes
pub fn validate(&mut self, address: SocketAddrV4, token: &[u8]) -> bool {
let prev = self.internal_generate_token(address, self.prev_secret);
let curr = self.internal_generate_token(address, self.curr_secret);
token == curr || token == prev
}
/// Rotate the tokens secret.
pub fn rotate(&mut self) {
trace!("Rotating secrets");
self.prev_secret = self.curr_secret;
self.curr_secret = random();
self.last_updated = Instant::now();
}
/// Generates a new token for a remote peer.
pub fn generate_token(&mut self, address: SocketAddrV4) -> [u8; 4] {
self.internal_generate_token(address, self.curr_secret)
}
// === Private Methods ===
fn internal_generate_token(
&mut self,
address: SocketAddrV4,
secret: [u8; SECRET_SIZE],
) -> [u8; TOKEN_SIZE] {
let mut digest = CASTAGNOLI.digest();
let octets: Box<[u8]> = address.ip().octets().into();
digest.update(&octets);
digest.update(&secret);
let checksum = digest.finalize();
checksum.to_be_bytes()
}
}
impl Default for Tokens {
fn default() -> Self {
Self::new()
}
}
fn random() -> [u8; SECRET_SIZE] {
let mut bytes = [0_u8; SECRET_SIZE];
getrandom::fill(&mut bytes).expect("getrandom");
bytes
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn valid_tokens() {
let mut tokens = Tokens::new();
let address = SocketAddrV4::new([127, 0, 0, 1].into(), 6881);
let token = tokens.generate_token(address);
assert!(tokens.validate(address, &token))
}
}
+428
View File
@@ -0,0 +1,428 @@
//! UDP socket layer managing incoming/outgoing requests and responses.
mod inflight_requests;
use crate::common::{ErrorSpecific, Message, MessageType, RequestSpecific, ResponseSpecific};
use inflight_requests::InflightRequests;
use std::io::ErrorKind;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4, UdpSocket};
use std::time::{Duration, Instant};
use tracing::{debug, trace, warn};
use super::config::Config;
const VERSION: [u8; 4] = [82, 83, 0, 5]; // "RS" version 05
const MTU: usize = 2048;
pub const DEFAULT_PORT: u16 = 6881;
/// Default request timeout before abandoning an inflight request to a non-responding node.
pub const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_millis(2000); // 2 seconds
pub const READ_TIMEOUT: Duration = Duration::from_millis(50);
/// Cleanup interval for expired inflight requests to avoid overhead on every recv
const INFLIGHT_CLEANUP_INTERVAL: Duration = Duration::from_millis(200);
/// A UdpSocket wrapper that formats and correlates DHT requests and responses.
#[derive(Debug)]
pub struct KrpcSocket {
next_tid: u32,
socket: UdpSocket,
pub(crate) server_mode: bool,
inflight_requests: InflightRequests,
last_cleanup: Instant,
local_addr: SocketAddrV4,
// poll_interval: Duration,
}
impl KrpcSocket {
pub(crate) fn new(config: &Config) -> Result<Self, std::io::Error> {
let request_timeout = config.request_timeout;
let port = config.port;
let bind_addr = config.bind_address.unwrap_or(Ipv4Addr::UNSPECIFIED);
let socket = if let Some(port) = port {
UdpSocket::bind(SocketAddr::from((bind_addr, port)))?
} else {
match UdpSocket::bind(SocketAddr::from((bind_addr, DEFAULT_PORT))) {
Ok(socket) => Ok(socket),
Err(_) => UdpSocket::bind(SocketAddr::from((bind_addr, 0))),
}?
};
let local_addr = match socket.local_addr()? {
SocketAddr::V4(addr) => addr,
SocketAddr::V6(_) => unimplemented!("KrpcSocket does not support Ipv6"),
};
socket.set_read_timeout(Some(READ_TIMEOUT))?;
Ok(Self {
socket,
next_tid: 0,
server_mode: config.server_mode,
inflight_requests: InflightRequests::new(request_timeout),
last_cleanup: Instant::now(),
local_addr,
})
}
#[cfg(test)]
pub(crate) fn server() -> Result<Self, std::io::Error> {
Self::new(&Config {
server_mode: true,
bind_address: Some(Ipv4Addr::LOCALHOST),
..Default::default()
})
}
#[cfg(test)]
pub(crate) fn client() -> Result<Self, std::io::Error> {
Self::new(&Config {
bind_address: Some(Ipv4Addr::LOCALHOST),
..Default::default()
})
}
// === Getters ===
/// Returns the address the server is listening to.
#[inline]
pub fn local_addr(&self) -> SocketAddrV4 {
self.local_addr
}
// === Public Methods ===
/// Returns true if this message's transaction_id is still inflight
pub fn inflight(&self, transaction_id: u32) -> bool {
self.inflight_requests.contains(transaction_id)
}
/// Send a request to the given address and return the transaction_id
pub fn request(&mut self, address: SocketAddrV4, request: RequestSpecific) -> u32 {
let message = self.request_message(request);
trace!(context = "socket_message_sending", message = ?message);
let tid = message.transaction_id;
self.inflight_requests.add(tid, address);
let _ = self.send(address, message).map_err(|e| {
debug!(?e, "Error sending request message");
});
tid
}
/// Send a response to the given address.
pub fn response(
&mut self,
address: SocketAddrV4,
transaction_id: u32,
response: ResponseSpecific,
) {
let message =
self.response_message(MessageType::Response(response), address, transaction_id);
trace!(context = "socket_message_sending", message = ?message);
let _ = self.send(address, message).map_err(|e| {
debug!(?e, "Error sending response message");
});
}
/// Send an error to the given address.
pub fn error(&mut self, address: SocketAddrV4, transaction_id: u32, error: ErrorSpecific) {
let message = self.response_message(MessageType::Error(error), address, transaction_id);
let _ = self.send(address, message).map_err(|e| {
debug!(?e, "Error sending error message");
});
}
/// Receives a single krpc message on the socket.
/// On success, returns the dht message and the origin.
pub fn recv_from(&mut self) -> Option<(Message, SocketAddrV4)> {
let mut buf = [0u8; MTU];
let now = Instant::now();
if now.duration_since(self.last_cleanup) > INFLIGHT_CLEANUP_INTERVAL {
self.last_cleanup = now;
self.inflight_requests.cleanup();
}
match self.socket.recv_from(&mut buf) {
Ok((amt, SocketAddr::V4(from))) => {
let bytes = &buf[..amt];
if from.port() == 0 {
trace!(
context = "socket_validation",
message = "Response from port 0"
);
return None;
}
match Message::from_bytes(bytes) {
Ok(message) => {
let should_return = match message.message_type {
MessageType::Request(_) => {
trace!(
context = "socket_message_receiving",
?message,
?from,
"Received request message"
);
true
}
MessageType::Response(_) => {
trace!(
context = "socket_message_receiving",
?message,
?from,
"Received response message"
);
self.is_expected_response(&message, &from)
}
MessageType::Error(_) => {
trace!(
context = "socket_message_receiving",
?message,
?from,
"Received error message"
);
self.is_expected_response(&message, &from)
}
};
if should_return {
return Some((message, from));
}
}
Err(error) => {
trace!(
context = "socket_error",
?error,
?from,
message = ?String::from_utf8_lossy(bytes),
"Received invalid Bencode message."
);
}
}
}
Ok((_, SocketAddr::V6(_))) => {
trace!(
context = "socket_validation",
message = "Received IPv6 packet"
);
}
Err(error) => match error.kind() {
// A read timeout means there was no packet this tick. Unix
// generally returns WouldBlock; Windows returns TimedOut.
ErrorKind::WouldBlock | ErrorKind::TimedOut => {}
_ => {
warn!("IO error {error}")
}
},
}
None
}
// === Private Methods ===
fn is_expected_response(&mut self, message: &Message, from: &SocketAddrV4) -> bool {
// Find and remove the matching inflight request
if let Some(_request) = self.inflight_requests.remove(message.transaction_id, from) {
return true;
} else {
trace!(
context = "socket_validation",
message = "Unexpected response id or wrong address"
);
}
false
}
/// Increments self.next_tid and returns the previous value.
fn tid(&mut self) -> u32 {
// We don't bother much with reusing freed transaction ids,
// since the timeout is so short we are unlikely to run out
// of 4294967295 ids in 2 seconds.
let tid = self.next_tid;
self.next_tid = self.next_tid.wrapping_add(1);
tid
}
/// Set transactin_id, version and read_only
fn request_message(&mut self, message: RequestSpecific) -> Message {
let transaction_id = self.tid();
Message {
transaction_id,
message_type: MessageType::Request(message),
version: Some(VERSION),
read_only: !self.server_mode,
requester_ip: None,
}
}
/// Same as request_message but with request transaction_id and the requester_ip.
fn response_message(
&mut self,
message: MessageType,
requester_ip: SocketAddrV4,
request_tid: u32,
) -> Message {
Message {
transaction_id: request_tid,
message_type: message,
version: Some(VERSION),
read_only: !self.server_mode,
// BEP_0042 Only relevant in responses.
requester_ip: Some(requester_ip),
}
}
/// Send a raw dht message
fn send(&mut self, address: SocketAddrV4, message: Message) -> Result<(), SendMessageError> {
self.socket.send_to(&message.to_bytes()?, address)?;
trace!(context = "socket_message_sending", message = ?message);
Ok(())
}
}
#[derive(thiserror::Error, Debug)]
/// Mainline crate error enum.
pub enum SendMessageError {
/// Errors related to parsing DHT messages.
#[error("Failed to parse packet bytes: {0}")]
BencodeError(#[from] serde_bencode::Error),
#[error(transparent)]
/// Transparent [std::io::Error]
IO(#[from] std::io::Error),
}
#[cfg(test)]
mod test {
use std::thread;
use crate::common::{Id, PingResponseArguments, RequestTypeSpecific};
use super::*;
#[test]
fn tid() {
let mut socket = KrpcSocket::server().unwrap();
assert_eq!(socket.tid(), 0);
assert_eq!(socket.tid(), 1);
assert_eq!(socket.tid(), 2);
socket.next_tid = u32::MAX;
assert_eq!(socket.tid(), 4294967295);
assert_eq!(socket.tid(), 0);
}
#[test]
fn recv_request() {
let mut server = KrpcSocket::server().unwrap();
let server_address = server.local_addr();
let mut client = KrpcSocket::client().unwrap();
client.next_tid = 120;
let client_address = client.local_addr();
let request = RequestSpecific {
requester_id: Id::random(),
request_type: RequestTypeSpecific::Ping,
};
let expected_request = request.clone();
let server_thread = thread::spawn(move || loop {
if let Some((message, from)) = server.recv_from() {
assert_eq!(from.port(), client_address.port());
assert_eq!(message.transaction_id, 120);
assert!(message.read_only, "Read-only should be true");
assert_eq!(message.version, Some(VERSION), "Version should be 'RS'");
assert_eq!(message.message_type, MessageType::Request(expected_request));
break;
}
});
client.request(server_address, request);
server_thread.join().unwrap();
}
#[test]
fn recv_response() {
let (tx, rx) = flume::bounded(1);
let mut client = KrpcSocket::client().unwrap();
let client_address = client.local_addr();
let responder_id = Id::random();
let response = ResponseSpecific::Ping(PingResponseArguments { responder_id });
let server_thread = thread::spawn(move || {
let mut server = KrpcSocket::client().unwrap();
let server_address = server.local_addr();
tx.send(server_address).unwrap();
loop {
server.inflight_requests.add(8, client_address);
if let Some((message, from)) = server.recv_from() {
assert_eq!(from.port(), client_address.port());
assert_eq!(message.transaction_id, 8);
assert!(message.read_only, "Read-only should be true");
assert_eq!(message.version, Some(VERSION), "Version should be 'RS'");
assert_eq!(
message.message_type,
MessageType::Response(ResponseSpecific::Ping(PingResponseArguments {
responder_id,
}))
);
break;
}
}
});
let server_address = rx.recv().unwrap();
client.response(server_address, 8, response);
server_thread.join().unwrap();
}
#[test]
fn ignore_response_from_wrong_address() {
let mut server = KrpcSocket::client().unwrap();
let server_address = server.local_addr();
let mut client = KrpcSocket::client().unwrap();
let client_address = client.local_addr();
server.inflight_requests.add(
8,
SocketAddrV4::new([127, 0, 0, 1].into(), client_address.port() + 1),
);
let response = ResponseSpecific::Ping(PingResponseArguments {
responder_id: Id::random(),
});
let _ = response.clone();
let server_thread = thread::spawn(move || {
thread::sleep(Duration::from_millis(5));
assert!(
server.recv_from().is_none(),
"Should not receive a response from wrong address"
);
});
client.response(server_address, 8, response);
server_thread.join().unwrap();
}
}
+86
View File
@@ -0,0 +1,86 @@
use std::collections::BTreeMap;
use std::net::SocketAddrV4;
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct InflightRequest {
pub to: SocketAddrV4,
pub sent_at: Instant,
}
impl InflightRequest {
pub fn does_match(&self, socket: &SocketAddrV4) -> bool {
if self.to.port() != socket.port() {
return false;
}
if self.to.ip().is_unspecified() {
return true;
}
self.to.ip() == socket.ip()
}
}
#[derive(Debug)]
pub struct InflightRequests {
// BTreeMap provides O(log n) lookup, insertion, and deletion keyed by transaction_id.
requests: BTreeMap<u32, InflightRequest>,
timeout: Duration,
}
impl InflightRequests {
pub fn new(timeout: Duration) -> Self {
Self {
requests: BTreeMap::new(),
timeout,
}
}
/// Add a new inflight request O(log n)
pub fn add(&mut self, transaction_id: u32, to: SocketAddrV4) {
self.requests.insert(
transaction_id,
InflightRequest {
to,
sent_at: Instant::now(),
},
);
}
/// Check if a transaction_id is still inflight and not expired O(log n)
pub fn contains(&self, transaction_id: u32) -> bool {
if let Some(request) = self.requests.get(&transaction_id) {
return request.sent_at.elapsed() < self.timeout;
}
false
}
/// Remove inflight request by transaction_id if it exists and matches the address
/// O(log n)
pub fn remove(&mut self, transaction_id: u32, from: &SocketAddrV4) -> Option<InflightRequest> {
let request = self.requests.get(&transaction_id)?;
// Drop immediately if expired; avoid accepting late responses
if request.sent_at.elapsed() >= self.timeout {
self.requests.remove(&transaction_id);
return None;
}
if !request.does_match(from) {
return None;
}
self.requests.remove(&transaction_id)
}
/// Cleanup expired requests based on timeout
/// O(n) scans all requests to remove expired ones
pub fn cleanup(&mut self) {
let now = Instant::now();
let cutoff = now - self.timeout;
// Remove expired requests in a single pass using retain
self.requests.retain(|_, request| request.sent_at > cutoff);
}
}