From 7c79bff593a104145e1964febac73a023cbf9f70 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=A1=A5=E4=B8=8B=E7=BA=A2=E8=8D=AF?= Date: Sun, 19 Jul 2026 15:26:04 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E9=87=8D=E6=9E=84=20DHT=20=E7=88=AC?= =?UTF-8?q?=E5=8F=96=E4=B8=8E=E5=8F=AF=E8=A7=82=E6=B5=8B=E6=80=A7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- CHANGELOG.md | 38 + Cargo.toml | 5 +- README.md | 468 +++--- docs/metrics.md | 99 ++ examples-jni/README.md | 127 +- .../java/cn/lmcw/dht/model/DHTOptions.java | 76 +- examples/main.rs | 41 +- jni/callbacks.rs | 2 +- jni/env.rs | 10 +- jni/exports.rs | 6 +- jni/server.rs | 2 +- jni/types.rs | 80 +- src/addr.rs | 71 + src/bootstrap.rs | 221 +++ src/budget.rs | 132 ++ src/crawl_config.rs | 137 ++ src/crawl_engine.rs | 928 ++++++++++++ src/error.rs | 7 + src/krpc.rs | 191 +++ src/lib.rs | 49 +- src/metadata.rs | 731 +++++++--- src/node_id.rs | 82 ++ src/node_pool.rs | 390 +++++ src/protocol.rs | 23 +- src/routing_snapshot.rs | 37 + src/runtime_stats.rs | 1063 ++++++++++++++ src/scheduler.rs | 1038 ++++++++++--- src/server.rs | 1283 ++++++++--------- src/sharded.rs | 291 ---- src/types.rs | 284 +++- src/udp_buffer.rs | 79 + src/udp_ingress.rs | 345 +++++ 32 files changed, 6738 insertions(+), 1598 deletions(-) create mode 100644 CHANGELOG.md create mode 100644 docs/metrics.md create mode 100644 src/addr.rs create mode 100644 src/bootstrap.rs create mode 100644 src/budget.rs create mode 100644 src/crawl_config.rs create mode 100644 src/crawl_engine.rs create mode 100644 src/krpc.rs create mode 100644 src/node_id.rs create mode 100644 src/node_pool.rs create mode 100644 src/routing_snapshot.rs create mode 100644 src/runtime_stats.rs delete mode 100644 src/sharded.rs create mode 100644 src/udp_buffer.rs create mode 100644 src/udp_ingress.rs diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..1e6ea13 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,38 @@ +# Changelog + +本项目遵循语义化版本。0.2.0 是包含公开 API 变更的 breaking release。 + +## 0.2.0 - 2026-07-11 + +### Breaking changes + +- 将 `DHTOptions` 的 Metadata 和 crawl 参数改为嵌套结构:`MetadataOptions`、 + `CrawlOptions`、`RateLimitOptions`、`PoolOptions`、`BootstrapOptions`、 + `TargetOptions`、`SchedulerOptions`。 +- 删除旧的 `metadata_timeout`、`max_metadata_queue_size`、 + `max_metadata_worker_count`、`node_queue_capacity` 等扁平字段。 +- 删除旧 active/candidate frontier 和 sharded queue 实现,改用单所有者严格 FIFO + 节点池、recent-probe set 与 responsive-node ring。 +- Metadata 调度改为有界、按 InfoHash 去重、最多三个 Peer、60 秒 freshness TTL。 + +### Added + +- 独立的主动爬取 QPS、新目标、节点替换、回复包/字节、单来源回复、总在途和子网在途限制。 +- 根据 Metadata 队列压力自动降低实际 `find_node` QPS。 +- Bootstrap 来源退避、低水位触发和响应节点快照。 +- `on_torrent_with_ack`、`on_metadata_fetch_complete`、 + `MetadataFetchCompletionStatus` 和真实 Peer `attempts`。 +- 按 `SocketAddr` 缓存 Metadata Peer 的 timeout/connect failure。 +- `DhtRuntimeStats::snapshot()`、`observability_snapshot()` 和三组固定桶直方图。 +- DHT、UDP、节点池、Metadata scheduler/fetcher 的低基数 Prometheus 指标。 + +### Changed + +- Metadata timeout 现在覆盖连接、握手、传输、SHA1 和解析的完整 Peer 尝试。 +- UDP ingress、crawl events 和 Metadata queues 全部有界,并暴露 drop/depth 指标。 +- DHT 回复增加总包、总字节、单来源限流,以及 `ping`/`get_peers` 10% 保底预算。 +- JNI `DHTOptions` 更新为 0.2 配置的扁平化子集;未映射字段继续采用 Rust 默认值。 + +### Migration + +迁移字段表、回调语义和运行时示例见 [README](README.md#从-01-迁移到-02)。 diff --git a/Cargo.toml b/Cargo.toml index bda76bf..28e8fa2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "dht-crawler" -version = "0.1.2" +version = "0.2.0" edition = "2024" authors = ["桥下红药 <1121744186@qq.com>"] description = "高性能的 Rust DHT (Distributed Hash Table) 爬虫库 | A high-performance Rust DHT crawler library for fetching torrent information from the BitTorrent DHT network" @@ -33,9 +33,12 @@ ahash = "0.8" serde_bytes = "0.11.19" metrics = { version = "0.24", optional = true } async-channel = "2.5.0" +crossbeam-queue = "0.3" jni = { version = "0.21", optional = true } +arc-swap = "1.7" [dev-dependencies] +tokio = { version = "1.35", features = ["signal"] } tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt", "tracing-log"] } mimalloc = "0.1" diff --git a/README.md b/README.md index 26633f5..5fc5d45 100644 --- a/README.md +++ b/README.md @@ -2,238 +2,382 @@ [![Crates.io](https://img.shields.io/crates/v/dht-crawler.svg)](https://crates.io/crates/dht-crawler) [![Documentation](https://docs.rs/dht-crawler/badge.svg)](https://docs.rs/dht-crawler) -[![License](https://img.shields.io/crates/l/dht-crawler.svg)](https://github.com/0xddy/dht-crawler/blob/master/LICENSE) +[![License](https://img.shields.io/crates/l/dht-crawler.svg)](LICENSE) -一个基于 Rust 和 Tokio 实现的高性能分布式哈希表 (DHT) 爬虫库。它能够加入 BitTorrent DHT 网络,监听并自动获取种子的元数据(Metadata/InfoHash)。 +基于 Rust、Tokio 的 BitTorrent DHT 爬虫库。它加入 BEP-5 网络,接收有效的 +`announce_peer`,并通过 BEP-9 `ut_metadata` 下载和校验种子元数据。 -## ✨ 核心特性 +当前版本:`0.2.0`。0.2 重做了节点池、主动爬取、Metadata 调度和运行时观测接口, +从 0.1 升级时请先阅读[迁移说明](#从-01-迁移到-02)和 [CHANGELOG](CHANGELOG.md)。 -- **🚀 极致性能**:基于 `Tokio` 异步运行时构建,支持数万级的高并发连接处理。 -- **📦 自动元数据抓取**:内置元数据获取引擎,自动完成从 InfoHash 到种子详情的抓取。 -- **🌐 双栈网络支持**:完美支持 IPv4 和 IPv6(DualStack 模式),扩大节点覆盖范围。 -- **⚡ 高度可配置**:支持自定义并发数、队列大小、超时时间等核心参数。 -- **📊 监控友好**:提供 Prometheus 指标导出接口,轻松监控爬虫状态(可选)。 +## 文档导航 -## 🏗️ 架构与流程 +- [快速开始](#快速开始):最小可运行示例和优雅停机。 +- [架构与背压](#架构与背压):UDP、主动爬取和 Metadata 管道。 +- [配置默认值](#配置默认值):所有公开 `DHTOptions` 字段。 +- [回调与生命周期](#回调与生命周期):抓取准入、交付确认和完成状态。 +- [运行时观测](#运行时观测):无 exporter 快照、Prometheus 和完整[指标表](docs/metrics.md)。 +- [JNI](#jni):Java 集成入口;[0.1 → 0.2 迁移](#从-01-迁移到-02)。 -本库采用了 **Reactor 模式** 与 **Worker Pool** 相结合的高并发架构,确保了在处理海量 UDP 数据包时的吞吐量。 +## 主要能力 -### 系统架构图 +- IPv4、IPv6 和双栈 DHT Socket。 +- 单所有者 crawl actor:严格 FIFO 节点池、最近探测状态、在途请求和所有速率预算 + 由一个 actor 管理,UDP worker 不锁节点池。 +- 查询 QPS、新目标/分钟、节点替换/分钟、总在途、子网在途、回复包、回复字节和 + 单来源回复分别限流。 +- 有界 Metadata 队列按 InfoHash 去重,并保留最多三个新鲜 Peer。 +- Metadata 总超时覆盖 TCP 连接、BitTorrent/扩展握手、分片下载、SHA1 校验和解析。 +- 按 `SocketAddr` 缓存 Peer 的超时/连接失败,避免坏 Peer 反复占用 worker。 +- 传输无关的原子运行时快照和固定桶直方图;可选 `metrics` feature。 +- 可选 JNI 接口和 Java 示例。 -```mermaid -graph TD - %% 网络层 - Network((DHT Network)) <-->|UDP Packets| Socket[UDP Socket] - - %% 接收与分发 - subgraph Receiver [Packet Receiver] - Socket -->|recv_from| Reader[UDP Reader / Dispatcher] - Reader -->|Round Robin| Ch1[Channel 1] - Reader -->|Round Robin| Ch2[Channel 2] - Reader -->|...| ChN[Channel N] - end - - %% 并行处理 - subgraph Processing [Packet Processing Workers] - Ch1 --> W1[Worker 1] - Ch2 --> W2[Worker 2] - ChN --> WN[Worker N] - - W1 & W2 & WN -->|Parse & Logic| Logic{Protocol Logic} - end - - %% 业务逻辑分支 - Logic -->|Discover Node| NodeMgr[Node Queue] - Logic -->|Discover InfoHash| HashQ[Hash Queue] - Logic -->|On Error| ErrorCb[User on_error Callback] - - %% 元数据抓取子系统 - subgraph Metadata [Metadata Subsystem] - HashQ --> Scheduler[Scheduler] - Scheduler -->|Spawn| MetaW1[Meta Worker 1] - Scheduler -->|...| MetaWN[Meta Worker N] - - MetaW1 & MetaWN <-->|TCP / ut_metadata| Peer((Remote Peer)) - end - - MetaW1 & MetaWN -->|Success| Callback[User Callback] -``` - -### 核心流程解析 - -1. **UDP 读取与分发 (Reader & Dispatcher)**: - * 独立的 UDP Reader 任务持续从 Socket 读取数据包。 - * 使用 Round-Robin 策略将数据包分发给 N 个(默认为 CPU 核心数)处理 Channel,实现无锁的负载均衡。 - -2. **并行协议处理 (Packet Workers)**: - * N 个 Packet Worker 并行消费 Channel 中的数据。 - * 负责 Bencode 解码、KRPC 协议解析、消息路由(Query/Response)。 - * 高效处理 `get_peers` 和 `announce_peer` 消息,提取 InfoHash。 - * 运行时错误通过 `on_error` 回调上报,不触发 panic,便于 JNI 等集成场景。 - -3. **元数据调度 (Metadata Subsystem)**: - * 提取出的 InfoHash 进入独立的 Hash Queue。 - * Scheduler 根据配置的并发度(如 1000+)动态启动 Metadata Worker。 - * Worker 通过 TCP 连接 Peer,使用 BEP-0009 协议下载种子元数据。 - -## 📦 安装 - -在你的 `Cargo.toml` 中添加依赖: +## 安装 ```toml [dependencies] -dht-crawler = "0.1" +dht-crawler = "0.2" +tokio = { version = "1", features = ["rt-multi-thread", "macros", "signal"] } ``` -如果需要 **Prometheus 监控支持**: +如果应用需要通过 `metrics` facade 输出指标: ```toml [dependencies] -dht-crawler = { version = "0.1", features = ["metrics"] } +dht-crawler = { version = "0.2", features = ["metrics"] } +metrics-exporter-prometheus = { version = "0.18", default-features = false, features = ["http-listener"] } ``` -## 🚀 快速开始 +`metrics` feature 只负责记录指标,不会在库内启动 HTTP 服务。应用必须自行安装 +recorder/exporter;完整指标清单见 [docs/metrics.md](docs/metrics.md)。 -下面是一个最简的启动示例。它会启动一个 DHT 节点,并在抓取到新种子时打印日志。 +### Cargo features + +| Feature | 默认启用 | 作用 | +|---|---|---| +| `metrics` | 否 | 通过 `metrics` facade 记录低基数指标 | +| `jni` | 否 | 构建 Java JNI 接口和 `cdylib` | +| `mimalloc` | 否 | 将 mimalloc 注册为全局分配器 | + +库的默认 feature 集为空。启用 `mimalloc` 前,请确认最终二进制没有注册其他全局分配器。 + +## 快速开始 ```rust use dht_crawler::prelude::*; #[tokio::main] async fn main() -> Result<()> { - // 1. 配置爬虫参数 let options = DHTOptions { port: 12313, + netmode: NetMode::Ipv4Only, + metadata: MetadataOptions { + timeout_secs: 4, + max_queue_size: 10_000, + max_worker_count: 256, + ..Default::default() + }, + crawl: CrawlOptions { + rate_limit: RateLimitOptions { + max_find_node_rate_per_sec: 200, + burst: 40, + max_in_flight: 512, + ..Default::default() + }, + ..Default::default() + }, ..Default::default() }; - // 2. 初始化 Server let server = DHTServer::new(options).await?; - println!("DHT Server 启动于端口 12313..."); - // 3. 注册错误回调:运行时错误通过回调输出,避免 panic(适合 JNI/嵌入式场景) - server.on_error(|err| { - eprintln!("DHT 错误: {}", err); + // 返回 true 才允许该 InfoHash 进入实际 Peer 下载阶段。 + server.on_metadata_fetch(|_info_hash| async move { true }); + + // 简单回调总是接受交付。 + server.on_torrent(|torrent| { + println!("{}: {}", torrent.info_hash, torrent.name); }); - // 4. 注册回调:成功获取到种子元数据时触发 - server.on_torrent(move |torrent| { - println!("🎉 抓取成功: {} (文件数: {})", torrent.name, torrent.files.len()); + server.on_error(|error| eprintln!("DHT runtime error: {error}")); + + let shutdown_server = server.clone(); + tokio::spawn(async move { + if tokio::signal::ctrl_c().await.is_ok() { + shutdown_server.shutdown(); + } }); - // 5. 可选:在拉取元数据前过滤 info_hash,返回 true 表示允许拉取 - server.on_metadata_fetch(|_hash| async move { true }); - - // 6. 启动服务 - server.start().await?; - Ok(()) + // start() 阻塞到 shutdown() 被调用。 + server.start().await } ``` -*完整的可运行代码请参考 [examples/main.rs](examples/main.rs)* +可运行版本见 [examples/main.rs](examples/main.rs)。如果自己的 Tokio 依赖没有启用 +`signal` feature,可以使用其他取消源调用 `shutdown()`。 -## ⚙️ 配置详解 +## 架构与背压 -`DHTOptions` 提供了丰富的配置项来调整爬虫行为: +```text +UDP sockets + ├─ bounded UDP worker queues ──→ KRPC workers ──→ bounded crawl events + │ │ + │ └─ announce_peer → bounded hash ingress + │ + └─ crawl egress ← single crawl actor ← priority/discovery events + │ + ├─ strict FIFO node pool + recent-probe set + ├─ pending transaction map + subnet counters + └─ ArcSwap responsive-node snapshot -```rust -let options = DHTOptions { - // 监听端口 - port: 12313, - - // 网络模式:Ipv4Only, Ipv6Only, 或 DualStack (默认) - netmode: NetMode::Ipv4Only, - - // 元数据获取超时时间 (秒) - metadata_timeout: 5, - - // 元数据下载队列大小,建议根据内存大小调整 - max_metadata_queue_size: 100000, - - // 同时进行元数据下载的并发任务数 - max_metadata_worker_count: 1000, - - // 节点池容量(DHT 节点队列) - node_queue_capacity: 100000, - - // InfoHash 发现队列容量 - hash_queue_capacity: 10000, - - ..Default::default() -}; +hash ingress → deduplicating Metadata queue → bounded workers → torrent callback ``` -**可选 API**:`server.set_filter(|info_hash_hex| bool)` 可在发现阶段过滤要处理的 info_hash(返回 `true` 表示允许进入元数据队列)。 +所有跨任务入口都是有界队列。达到容量时,事件会被拒绝、淘汰或计入 drop 指标, +不会依靠无限增长的缓冲区掩盖下游过载。主动爬取预算还会随 Metadata 队列压力下降。 -## 错误处理 +### 主动爬取 -库内采用严格的错误处理策略,避免底层异常导致进程崩溃,便于集成 JNI 或嵌入式场景。 +- 新地址通常只发送一次 `find_node`;默认等待回复 `2s`,不做同目标重试。 +- 超时不会自动降低配置 QPS。Metadata 队列压力达到 80% 后才开始自动降速,95% 时 + 降到 `metadata_pressure_floor_percent` 指定的比例;默认下限为配置 QPS 的 25%。 +- 节点池是严格 FIFO。重复地址、无效公网地址和超出 replacement budget 的替换会被拒绝。 +- 响应成功的节点进入一个独立、有界、带 TTL 的 responsive ring,用于回复其他 DHT + 节点和 revisit 查询;它不是第二个爬取池。 +- 节点池低于 `low_watermark` 时触发 bootstrap。失败的 bootstrap 来源按配置退避。 +- UDP 回复总包数、总字节数和单来源包数分别限流,其中 10% 包/字节预算保留给 + `ping` 和 `get_peers` 的保底回复,但不会突破配置的总上限。 -### 错误类型 `DHTError` +### Metadata 调度 -```rust -use dht_crawler::{DHTError, Result}; +- Hash ingress 和 Metadata pending queue 都是有界的。 +- Pending queue 按 InfoHash 去重,每个 Hash 最多保留三个不同且新鲜的 Peer。 +- Pending 项固定在 60 秒后过期。队列满时,较新的 Hash 可以淘汰最旧项;比当前 + 最旧项还旧的事件直接视为 stale。 +- worker 优先分派最新的可用 Hash,以提高 Peer 仍在线的概率。 +- `timeout_secs` 是一次 Peer 尝试的端到端期限,不会在连接、握手和下载阶段重复叠加。 +- 单个 metadata payload 上限为 10 MiB;下载完成后必须通过 SHA1 和 bencode 解析。 +- Peer failure cache 只缓存 `timeout` 和 `connect_failed`,键为完整 `SocketAddr` + (IP + port)。缓存命中不会发起网络请求,也不计入三次真实 Peer 尝试。 +- `peer_failure_cache_capacity = 0` 或 `peer_failure_ttl_secs = 0` 会关闭缓存。 -// 错误变体包括: -// - DHTError::Network(io::Error) — 网络/IO 错误 -// - DHTError::Init(String) — 初始化错误(如 socket、worker) -// - DHTError::Internal(String) — 内部逻辑错误 -// - DHTError::LockPoisoned(String) — 锁中毒(预留) -// - DHTError::Other(String) — 其他 +## 配置默认值 + +库本身不包含 P1/P15 等档位概念。应用如需档位,应将其转换成下列具体选项。 + +### `DHTOptions` 与 Metadata + +| 字段 | 默认值 | 说明 | +|---|---:|---| +| `port` | `6881` | DHT UDP 监听端口 | +| `netmode` | `Ipv4Only` | `DHTOptions::default()` 的网络模式 | +| `hash_queue_capacity` | `10000` | announce 到 Metadata scheduler 的 ingress 容量 | +| `metadata.timeout_secs` | `4` | 单 Peer 端到端超时 | +| `metadata.max_queue_size` | `10000` | 去重 Pending Hash 容量 | +| `metadata.max_worker_count` | `256` | 最大并发 Metadata job 数 | +| `metadata.peer_failure_cache_capacity` | `200000` | 坏 Peer 缓存容量 | +| `metadata.peer_failure_ttl_secs` | `60` | 坏 Peer 缓存 TTL | + +### `crawl.rate_limit` + +| 字段 | 默认值 | +|---|---:| +| `max_find_node_rate_per_sec` | `200` | +| `burst` | `40` | +| `max_in_flight` | `512` | +| `request_timeout_secs` | `2` | +| `max_new_destinations_per_minute` | `10000` | +| `max_replacements_per_minute` | `25000` | +| `max_response_rate_per_sec` | `500` | +| `max_response_bytes_per_sec` | `1048576` | +| `max_response_rate_per_source` | `40` | +| `metadata_pressure_floor_percent` | `25` | +| `max_in_flight_per_subnet` | `8` | + +### Pool、Bootstrap、Target 与 Scheduler + +| 字段 | 默认值 | +|---|---:| +| `pool.capacity` | `100000` | +| `pool.recent_probe_ttl_secs` | `600` | +| `pool.responsive_capacity` | `16384` | +| `pool.responsive_ttl_secs` | `900` | +| `pool.low_watermark` | `10000` | +| `bootstrap.interval_secs` | `300` | +| `bootstrap.max_nodes_per_round` | `3` | +| `bootstrap.source_backoff_base_secs` | `300` | +| `bootstrap.source_backoff_max_secs` | `3600` | +| `target.random_walk_percent` | `70` | +| `target.sparse_bucket_percent` | `30` | +| `target.neighbor_sender_id` | `true` | +| `scheduler.priority_event_channel_capacity` | `8192` | +| `scheduler.discovery_event_channel_capacity` | `16384` | +| `scheduler.event_batch_limit` | `256` | +| `scheduler.node_batch_limit` | `4096` | +| `scheduler.routing_snapshot_size` | `4096` | +| `scheduler.snapshot_refresh_millis` | `1000` | + +默认 bootstrap 来源: + +```text +router.bittorrent.com:6881 +dht.transmissionbt.com:6881 +router.utorrent.com:6881 +dht.aelitis.com:6881 ``` -### 注册错误回调 `on_error` +内部会对不安全的零值和百分比做归一化,例如容量/在途至少为 1、百分比最大为 100、 +`low_watermark` 不超过 pool capacity。建议调用方仍显式传入有效配置,不依赖归一化。 -运行时错误(如协议处理失败)会通过回调上报,而不会 panic: +## 回调与生命周期 + +### `on_metadata_fetch` + +在 Hash 首次准备进入 Peer 下载前调用。返回 `false` 表示 gate reject:不下载、不触发 +`on_torrent`,也不会触发 `on_metadata_fetch_complete`。 + +### `on_torrent` 与 `on_torrent_with_ack` + +- `on_torrent` 适合无需确认交付的调用方;回调返回后视为 `Accepted`。 +- `on_torrent_with_ack` 返回 `true` 表示应用接受交付,返回 `false` 表示 + `DeliveryRejected`。后者代表 Metadata 已成功下载,但业务层没有接收,不等同于 + `FetchFailed`。 +- 如果没有注册任何 torrent callback,成功下载的 Metadata 同样按 `DeliveryRejected` + 结束。 +- Torrent 回调发生 panic 时会被捕获并按拒绝交付处理。 + +### `on_metadata_fetch_complete` + +一个通过 gate 的 Hash 最终只发出一次完成通知: + +| 状态 | 含义 | +|---|---| +| `Accepted` | 下载成功,torrent callback 接受交付 | +| `FetchFailed` | 所有可用 Peer 尝试失败 | +| `DeliveryRejected` | 下载成功,torrent callback 拒绝交付 | + +`attempts` 只统计实际发起的 Peer 网络尝试;failure cache 命中不计入。 + +### 启动与停止 + +- `DHTServer::new()` 创建并绑定所需 Socket,失败直接返回 `Err`。 +- `start()` 启动后台任务并等待取消,因此通常应在应用主任务中 await。 +- `shutdown()` 可从 clone handle 调用,取消 DHT、crawl 和 Metadata 后台任务。 +- 已 shutdown 的同一实例不能重新 start;需要重新构造 `DHTServer`。 +- `on_error` 用于接收运行期协议/worker 错误;初始化错误仍通过 `Result` 返回。 + +## 运行时观测 + +不启用 `metrics` feature 也可以读取运行时快照: ```rust -let server = DHTServer::new(options).await?; +let stats = server.runtime_stats(); -// 将错误输出到 stderr 或接入自己的日志/监控 -server.on_error(|err| { - log::error!("DHT 运行时错误: {}", err); -}); +let runtime = stats.snapshot(); +println!( + "nodes={} metadata={}/{} in_flight={}", + runtime.node_pool_size, + runtime.metadata_queue_depth, + runtime.metadata_queue_max, + runtime.metadata_in_flight, +); -// JNI 场景示例:将错误传回 Java 层 -// server.on_error(|err| { -// jni_callback_on_error(env, err.to_string()); -// }); +let observability = stats.observability_snapshot(); +println!( + "udp rx={}B tx={}B fetch_p95={:?}ms", + observability.udp_rx_bytes, + observability.udp_tx_bytes, + observability.fetch_duration_ms.percentile(0.95), +); ``` -### 同步错误:`Result` 传播 +快照使用 relaxed atomic load,适合监控,不是跨字段事务视图。计数器是进程生命周期 +累计值,调用方通过相邻快照差值计算 rate,并应处理进程重启导致的 counter reset。 -`DHTServer::new()` 和 `server.start()` 返回 `Result`,调用方需处理或传播: +固定桶: + +| 直方图 | 边界/单位 | +|---|---| +| Metadata queue wait | `10/50/100/250/500/1000/2000/5000 ms` | +| Metadata fetch | `250/500/1000/2000/4000/6000/10000 ms` | +| Metadata payload size | `16/32/64/128/256/512/1024 KiB, 10 MiB` | + +`counts[i]` 表示小于等于 `bounds[i]` 的非累计桶计数,`overflow` 表示超过最后边界的 +数量。`percentile()` 返回桶上界;落入 overflow 时只能返回最后一个上界,因此它是 +有界近似值,不是精确分位数。 + +高级快照类型从 crate 根导出;当前 prelude 只重导出 `DhtRuntimeStats` 和 +`DhtRuntimeSnapshot`。 + +## Prometheus / metrics + +启用 `metrics` 后,库通过 `metrics` facade 记录 counter、gauge 和 histogram。 +应用必须在创建/启动 server 前安装全局 recorder。示例: ```rust -#[tokio::main] -async fn main() -> Result<()> { - let server = DHTServer::new(options).await?; // 初始化失败会返回 Err - server.on_error(|e| eprintln!("{}", e)); - server.start().await?; // 启动失败(如 socket 绑定)会返回 Err - Ok(()) -} +use metrics_exporter_prometheus::PrometheusBuilder; + +PrometheusBuilder::new() + .with_http_listener("127.0.0.1:9000".parse().unwrap()) + .install() + .unwrap(); ``` -## 编译示例 +完整名称、标签和单位见 [docs/metrics.md](docs/metrics.md)。 -### 1. 启用 `mimalloc`(内存优化) +## 应用层集成 -在长运行的高并发场景下,使用 `mimalloc` 可降低约 10–30% 内存占用。本库的 `mimalloc` feature 仅用于方便编译/运行示例;若将本库作为依赖使用,请在自己的 bin 项目中单独引入并配置 mimalloc 全局分配器。 +本 crate 只提供 DHT、BEP-9 Metadata、回调和观测能力,不包含 Redis、Manticore、 +HTTP 看板或 `P1`~`P16` 性能档位。同级的 `dht-crawler-node` 项目负责这些应用层策略, +并把档位转换成具体的 `DHTOptions`。开发两个项目时应保持下面的目录关系: + +```text +workspace-parent/ +├── dht-crawler/ +└── dht-crawler-node/ +``` + +## JNI + +启用 `jni` feature 可构建 `cdylib`。Java 示例、线程模型和 JNI 配置字段见 +[examples-jni/README.md](examples-jni/README.md)。Java `DHTOptions` 是 Rust 配置的 +扁平化子集,未暴露的 Bootstrap、Target、Scheduler 和 Peer failure cache 字段使用 +Rust 默认值。 + +## 从 0.1 迁移到 0.2 + +0.2 是 breaking release: + +| 0.1 | 0.2 | +|---|---| +| `metadata_timeout` | `metadata.timeout_secs` | +| `max_metadata_queue_size` | `metadata.max_queue_size` | +| `max_metadata_worker_count` | `metadata.max_worker_count` | +| `node_queue_capacity` | `crawl.pool.capacity` | +| 旧 active/candidate frontier | 单所有者严格 FIFO pool + responsive ring | +| 无交付确认 | `on_torrent_with_ack` + `DeliveryRejected` | +| 粗粒度统计 | `runtime_stats()` 的两类原子快照和固定桶 | + +`DHTOptions` 不提供旧字段兼容层,升级时必须修改构造代码。配置档位属于应用策略, +不在库内实现。 + +## 构建与验证 ```bash -cargo run --release --example dht_crawler_example --features mimalloc +cargo fmt --all --check +cargo check --all-targets --all-features +cargo clippy --all-targets --all-features -- -D warnings +cargo test --all-features +cargo doc --no-deps --all-features +cargo run --release --example dht_crawler_example ``` -### 2. 启用 `metrics`(监控) - -启用后,可通过 HTTP 接口拉取 Prometheus 格式的监控数据。 +JNI: ```bash -cargo run --release --example dht_crawler_example --features metrics +cargo build --release --features jni ``` -监控地址:http://localhost:9000/metrics +## 许可证 -## 📜 许可证 - -MIT License +[MIT](LICENSE) diff --git a/docs/metrics.md b/docs/metrics.md new file mode 100644 index 0000000..675284a --- /dev/null +++ b/docs/metrics.md @@ -0,0 +1,99 @@ +# dht-crawler 指标参考 + +启用 Cargo feature `metrics` 后,库通过 `metrics` facade 记录以下指标。库不安装 +recorder、不监听端口,也不依赖任何特定导出协议;Prometheus exporter 应由最终应用安装。 + +所有 `*_total` 都是进程生命周期累计 counter。Gauge 是当前值。Histogram 的记录值 +使用下表标出的单位。 + +## UDP 与 KRPC + +| 指标 | 类型 | 标签 | 单位/含义 | +|---|---|---|---| +| `dht_udp_bytes_received_total` | counter | — | Socket 接收字节 | +| `dht_udp_packets_received_total` | counter | `status=ok|dropped_size|dropped_magic|queue_full` | UDP ingress 结果 | +| `dht_udp_bytes_sent_total` | counter | — | 成功发送字节 | +| `dht_udp_packets_sent_total` | counter | `type=query|response` | 成功发送包 | +| `dht_udp_query_size_bytes` | histogram | — | find_node query 编码长度,bytes | +| `dht_messages_processed_total` | counter | `type=q|r|e|unknown` | 成功解析的 KRPC 消息类型 | +| `dht_messages_parse_error_total` | counter | — | bencode/KRPC 解析失败 | +| `dht_queries_total` | counter | `q=ping|find_node|get_peers|announce_peer|vote|other_or_invalid` | 入站查询类型 | +| `dht_udp_responses_dropped_total` | counter | `reason=rate_limit` | 最终未发送的限流回复 | +| `dht_udp_responses_priority_reserved_total` | counter | `query=ping|get_peers` | 使用 10% 保底预算的回复 | + +`dht_udp_bytes_received_total` 包含后续被判定为 invalid/queue-full 的 Datagram;发送侧只在 +`send_to` 成功后累计。 + +## 主动爬取与节点池 + +| 指标 | 类型 | 标签 | 含义 | +|---|---|---|---| +| `dht_node_pool_size` | gauge | — | 当前 FIFO 节点数 | +| `dht_node_pool_oldest_age_seconds` | gauge | — | FIFO 最老节点年龄 | +| `dht_node_pool_admissions_total` | counter | — | 新准入节点 | +| `dht_node_pool_replacements_total` | counter | — | 满池替换 | +| `dht_node_pool_dropped_total` | counter | `reason=duplicate|rate_limit|invalid` | 节点拒绝原因 | +| `dht_find_node_in_flight` | gauge | — | 当前在途 find_node | +| `dht_find_node_effective_rate_per_second` | gauge | — | Metadata 压力调整后的实际预算 | +| `dht_crawl_queries_sent_total` | counter | `kind=new|revisit|bootstrap` | 已交给 egress 的查询用途;发送失败另计 | +| `dht_find_node_responses_total` | counter | — | 与 pending transaction 匹配的回复 | +| `dht_find_node_response_unmatched_total` | counter | — | 无匹配 pending 的回复 | +| `dht_find_node_timeouts_total` | counter | — | pending 超时 | +| `dht_find_node_send_failures_total` | counter | — | UDP query 发送失败 | +| `dht_crawl_events_dropped_total` | counter | `kind=discovered|response` | 有界 actor channel 丢弃 | +| `dht_metadata_queue_pressure_ratio` | gauge | — | Metadata depth/capacity,范围 0..1 | + +actor 每秒把内部增量 flush 到 counter,因此 exporter 看到的 counter 可能最多延迟约一秒。 + +## announce 与 Metadata ingress + +| 指标 | 类型 | 标签 | 含义 | +|---|---|---|---| +| `dht_announce_peer_blocked_total` | counter | `reason=invalid_token|filtered` | announce 拒绝原因 | +| `dht_info_hashes_discovered_total` | counter | — | token/hash/filter 校验通过的 InfoHash | +| `dht_metadata_ingress_dropped_total` | counter | `reason=queue_full` | Hash ingress 满导致的丢弃 | + +## Metadata scheduler + +| 指标 | 类型 | 标签 | 单位/含义 | +|---|---|---|---| +| `dht_metadata_queue_depth` | gauge | — | Pending Hash 数 | +| `dht_metadata_in_flight` | gauge | — | 当前 job 数 | +| `dht_metadata_queue_events_total` | counter | `result=inserted|deduplicated|evicted_oldest|stale|expired` | 队列事件 | +| `dht_metadata_queue_wait_seconds` | histogram | — | Hash 从最近发现到首次分派的秒数 | +| `dht_metadata_jobs_dispatched_total` | counter | — | 分派 job 数 | +| `dht_metadata_jobs_completed_total` | counter | `result=accepted|fetch_failed|delivery_rejected|gate_rejected` | job 终态 | +| `dht_metadata_worker_join_error_total` | counter | — | worker task join 失败 | +| `dht_metadata_completion_callback_panics_total` | counter | — | 完成回调 panic | + +## Metadata Peer 下载 + +| 指标 | 类型 | 标签 | 单位/含义 | +|---|---|---|---| +| `dht_metadata_fetch_attempts_total` | counter | — | 实际 Peer 尝试 | +| `dht_metadata_peer_attempts_total` | counter | — | 与 fetch attempts 相同的 Peer 尝试计数 | +| `dht_metadata_fetch_success_total` | counter | — | 成功下载和解析 | +| `dht_metadata_fetch_result_total` | counter | `result=success|failed|timeout` | Peer 尝试结果 | +| `dht_metadata_fetch_fail_total` | counter | `reason=timeout|send_error|size_limit|sha1_mismatch|parse_error` | 详细失败原因 | +| `dht_metadata_connection_result_total` | counter | `result=success|failed` | TCP/BitTorrent connect 结果 | +| `dht_metadata_handshake_result_total` | counter | `result=success|no_extension_support` | 扩展能力/最终校验结果 | +| `dht_metadata_fetch_duration_seconds` | histogram | — | 端到端 Peer 尝试秒数 | +| `dht_metadata_size_bytes` | histogram | — | 完整 bencoded info payload 字节数 | +| `dht_metadata_bytes_downloaded_total` | counter | — | 收到的 Metadata piece 数据字节数 | +| `dht_metadata_peer_failure_cache_hits_total` | counter | `reason=timeout|connect_failed` | 坏 Peer 缓存命中 | +| `dht_metadata_peer_failure_cache_inserts_total` | counter | `reason=timeout|connect_failed` | 坏 Peer 缓存写入 | +| `dht_metadata_peer_failure_cache_entries` | gauge | — | 当前缓存条目数 | + +`dht_metadata_fetch_result_total{result="failed"}` 汇总非 timeout 的失败,不适合单独用于 +分析具体原因;详细原因应结合 `fetch_fail`、connection 和 handshake 指标。 + +## 与原子快照的关系 + +`DhtRuntimeStats` 始终可用,与 `metrics` feature 无关: + +- `snapshot()` 提供队列、节点池、crawl、Peer 和 UDP 运行状态。 +- `observability_snapshot()` 提供 UDP 字节/包、入站查询分类、announce、节点准入、 + Metadata 失败分类、failure cache 分类和固定桶。 + +两套出口在同一事件点更新,但读取时都不是跨字段事务快照;短时间内可能相差一个并发 +事件,Prometheus 的 crawl actor counter 还可能有最多约一秒 flush 延迟。 diff --git a/examples-jni/README.md b/examples-jni/README.md index 215183f..5a9f862 100644 --- a/examples-jni/README.md +++ b/examples-jni/README.md @@ -1,6 +1,6 @@ -# dht-crawler-jni Java 示例 +# dht-crawler JNI Java 示例 -本目录是一个 Gradle 管理的 Java 项目,演示如何通过 JNI 使用 `dht-crawler` Rust 库。 +Gradle 项目,演示通过 JNI 调用 `dht-crawler`(需启用 Cargo feature `jni`)。 ## 项目结构 @@ -9,79 +9,116 @@ examples-jni/ ├── build.gradle ├── settings.gradle └── src/main/java/cn/lmcw/dht/ - ├── model/ # DHTOptions, TorrentInfo, FileInfo - ├── DhtCrawler.java # 面向对象入口(推荐) - ├── DhtCrawlerJni.java # 包内 native 绑定 - ├── DhtListener.java + ├── model/ # DHTOptions, TorrentInfo, FileInfo + ├── DhtCrawler.java # 面向对象入口(推荐) + ├── DhtCrawlerJni.java # native 方法声明 + ├── DhtListener.java # 回调接口 └── DhtCrawlerExample.java ``` -## 快速开始 +Rust JNI 实现位于 **`dht-crawler/jni/`**(`lib` crate-type 含 `cdylib`)。 -### 方式一:直接下载 Release 中的 JAR 运行(推荐) +## 编译 native 库 -从 GitHub Release 页面下载两个文件: - -1. `dht-crawler-jni-example-.jar` — 平台无关的 fat JAR -2. 对应平台的 JNI 动态库 zip(如 `dht_crawler_jni--x86_64-unknown-linux-gnu.zip`),解压得到 `libdht_crawler.so` / `dht_crawler.dll` / `libdht_crawler.dylib` - -将 JAR 和动态库放到同一目录,然后执行: - -```bash -# Linux / macOS -java -Djava.library.path=. -jar dht-crawler-jni-example-.jar - -# Windows(PowerShell/CMD 需对 -D 参数加引号,否则会报找不到主类) -java "-Djava.library.path=." -jar dht-crawler-jni-example-.jar -``` - -### 方式二:从源码编译并运行 - -#### 1. 编译 Rust JNI 动态库 - -在仓库**根目录**执行: +在 **`dht-crawler`** 目录(本 README 的上一级)执行: ```bash cargo build --release --features jni ``` -产物路径:`target/release/` +产物路径(因平台而异): -#### 2. 运行 Java 示例 +- Linux: `target/release/libdht_crawler.so` +- Windows: `target/release/dht_crawler.dll` +- macOS: `target/release/libdht_crawler.dylib` -在本目录(`examples-jni/`)执行: +## 运行示例 + +### 方式一:Release 预编译包(推荐) + +下载 Release 中的 fat JAR 与对应平台 native 库 zip,放在同一目录: ```bash -gradle run -gradle run -Plib.path=/path/to/your/lib +# Linux / macOS +java -Djava.library.path=. -jar dht-crawler-jni-example-.jar + +# Windows +java "-Djava.library.path=." -jar dht-crawler-jni-example-.jar ``` -#### 3. 构建 fat JAR +### 方式二:源码 + Gradle + +```bash +cd examples-jni +gradle run +# 指定 native 库目录(默认为 ../target/release) +gradle run -Plib.path=/path/to/lib +``` + +构建 fat JAR: ```bash gradle shadowJar ``` -## 在自己的项目中集成 +## 集成到自己的 Java 项目 -1. 复制 `cn/lmcw/dht/` 下源码(含 `model/`、`DhtCrawler`、`DhtCrawlerJni`、`DhtListener`)。 -2. 将对应平台的 so/dll/dylib 放入 `java.library.path`。 +1. 复制 `cn/lmcw/dht/` 包(含 `model/`、`DhtCrawler`、`DhtCrawlerJni`、`DhtListener`)。 +2. 将对应平台的 native 库加入 `java.library.path`。 +3. 使用与示例相同的 `dht-crawler` JNI 版本构建 `cdylib`。 ## API(面向对象) ```java DhtCrawler crawler = DhtCrawler.createServer(options, listener); -crawler.start(); +crawler.start(); // 后台启动,不阻塞调用线程 // ... -crawler.stop(); // 或 try-with-resources +crawler.stop(); // 或 try-with-resources ``` -- **`DhtCrawler.createServer(options, listener)`**:创建会话(未启动 DHT;Java 不能用方法名 `new`,故不用 `open`)。 -- **`start()`**:后台启动 DHT,非阻塞;同一会话多次 `start()` 仅首次生效。 -- **`stop()` / `close()`**:停止并释放 Rust 资源,幂等。 -- **`getNodePoolSize()`**:routing table 节点数。 +| 方法 | 说明 | +|------|------| +| `createServer(options, listener)` | 创建 `ServerHandle`(含独立 tokio `Runtime` + `DHTServer`) | +| `start()` | 在 runtime 内 spawn `server.start()`;同一会话多次调用仅首次生效 | +| `stop()` / `close()` | `shutdown()` 后在后台线程 drop Runtime,避免阻塞 JNI 线程 | +| `getNodePoolSize()` | 节点池大小(`DHTServer::get_node_pool_size`) | + +Native 导出类:`cn.lmcw.dht.DhtCrawlerJni`(`createServer` / `startServer` / `stopServer` / `getNodePoolSize`)。 + +## 回调与线程 + +- `onTorrent` / `onError`:在 Rust 工作线程触发,Java 实现须线程安全。 +- `onMetadataFetch`:在阻塞线程池中调用,应尽快返回 boolean。 + +行为与 Rust 库一致:InfoHash 来自 **`announce_peer`**;`start()` 在 Rust 侧仍阻塞至 `shutdown()`,JNI 通过单独 runtime + `spawn` 避免卡住 Java 主流程。 + +当前 JNI listener 暴露 `onTorrent`、`onMetadataFetch` 和 `onError`,不暴露 Rust +`on_torrent_with_ack` / `on_metadata_fetch_complete`。因此 Java `onTorrent` 返回后始终按 +`Accepted` 处理;需要交付确认和最终状态的应用应扩展 JNI callback contract。 + +## JNI 配置映射 + +Java `DHTOptions` 是 Rust 配置的扁平化子集,不是全部 Rust 字段的一一镜像: + +| Java 字段组 | Rust 目标 | +|---|---| +| port / netMode / hashQueueCapacity | `DHTOptions` 顶层 | +| metadataTimeout / maxMetadataQueueSize / maxMetadataWorkerCount | `metadata.*` | +| poolCapacity / recentProbeTtlSeconds / responsive* / poolLowWatermark | `crawl.pool.*` | +| findNode* / requestTimeout* / response* / pressure / replacements / subnet | `crawl.rate_limit.*` | + +以下配置未通过当前 JNI 暴露,使用 Rust `Default`: + +- `metadata.peer_failure_cache_capacity`、`metadata.peer_failure_ttl_secs` +- `crawl.bootstrap.*` +- `crawl.target.*` +- `crawl.scheduler.*` + +Java 字段默认值与当前 Rust 0.2 默认值保持一致;传入 `null` options 时直接使用完整的 +`DHTOptions::default()`。 ## 注意事项 -- 回调在 Rust 工作线程触发,实现需线程安全。 -- `onMetadataFetch` 在阻塞线程池调用,宜快速返回。 +- Java 11+(见 `build.gradle`)。 +- Java/Rust 版本必须一致,避免 JNI 按字段名和签名读取时失败。 +- 停止后勿再使用同一 `long` 句柄。 diff --git a/examples-jni/src/main/java/cn/lmcw/dht/model/DHTOptions.java b/examples-jni/src/main/java/cn/lmcw/dht/model/DHTOptions.java index dad39fd..f9b41e6 100644 --- a/examples-jni/src/main/java/cn/lmcw/dht/model/DHTOptions.java +++ b/examples-jni/src/main/java/cn/lmcw/dht/model/DHTOptions.java @@ -1,8 +1,9 @@ package cn.lmcw.dht.model; /** - * 与 Rust {@code DHTOptions} 一一对应的配置对象。 - *

通过 JNI 传入 Rust 侧,由 Rust 读取各字段构造 {@code DHTOptions}。

+ * Rust {@code DHTOptions} 的 JNI 扁平化子集。 + *

通过 JNI 传入 Rust 侧;未暴露的 Metadata failure cache、Bootstrap、Target 和 + * Scheduler 字段采用 Rust 0.2 默认值。Java 与 native 库版本必须保持一致。

* *

netMode 取值

*
    @@ -17,16 +18,37 @@ public final class DHTOptions { private int port = 6881; /** 获取 metadata 超时(秒),默认 3 */ - private long metadataTimeout = 3L; + private long metadataTimeout = 4L; /** metadata 队列最大容量,默认 100000 */ - private int maxMetadataQueueSize = 100_000; + private int maxMetadataQueueSize = 10_000; /** 并发 metadata 拉取 worker 数量,默认 1000 */ - private int maxMetadataWorkerCount = 1_000; + private int maxMetadataWorkerCount = 256; /** 节点队列容量,默认 100000 */ - private int nodeQueueCapacity = 100_000; + private int poolCapacity = 100_000; + + private int findNodeRatePerSecond = 200; + + private int findNodeBurst = 40; + + private int maxFindNodeInFlight = 512; + + private int maxNewDestinationsPerMinute = 10_000; + + private int maxReplacementsPerMinute = 25_000; + + private long requestTimeoutSeconds = 2L; + private int maxResponseRatePerSecond = 500; + private long maxResponseBytesPerSecond = 1_048_576L; + private int maxResponseRatePerSource = 40; + private int metadataPressureFloorPercent = 25; + private long recentProbeTtlSeconds = 600L; + private int responsiveCapacity = 16_384; + private long responsiveTtlSeconds = 900L; + private int poolLowWatermark = 10_000; + private int maxInFlightPerSubnet = 8; /** hash 队列容量,默认 10000 */ private int hashQueueCapacity = 10_000; @@ -60,12 +82,48 @@ public final class DHTOptions { return this; } - public int getNodeQueueCapacity() { return nodeQueueCapacity; } - public DHTOptions setNodeQueueCapacity(int nodeQueueCapacity) { - this.nodeQueueCapacity = nodeQueueCapacity; + public int getPoolCapacity() { return poolCapacity; } + public DHTOptions setPoolCapacity(int poolCapacity) { + this.poolCapacity = poolCapacity; return this; } + public int getFindNodeRatePerSecond() { return findNodeRatePerSecond; } + public DHTOptions setFindNodeRatePerSecond(int value) { this.findNodeRatePerSecond = value; return this; } + + public int getFindNodeBurst() { return findNodeBurst; } + public DHTOptions setFindNodeBurst(int value) { this.findNodeBurst = value; return this; } + + public int getMaxFindNodeInFlight() { return maxFindNodeInFlight; } + public DHTOptions setMaxFindNodeInFlight(int value) { this.maxFindNodeInFlight = value; return this; } + + public int getMaxNewDestinationsPerMinute() { return maxNewDestinationsPerMinute; } + public DHTOptions setMaxNewDestinationsPerMinute(int value) { this.maxNewDestinationsPerMinute = value; return this; } + + public int getMaxReplacementsPerMinute() { return maxReplacementsPerMinute; } + public DHTOptions setMaxReplacementsPerMinute(int value) { this.maxReplacementsPerMinute = value; return this; } + + public long getRequestTimeoutSeconds() { return requestTimeoutSeconds; } + public DHTOptions setRequestTimeoutSeconds(long value) { this.requestTimeoutSeconds = value; return this; } + public int getMaxResponseRatePerSecond() { return maxResponseRatePerSecond; } + public DHTOptions setMaxResponseRatePerSecond(int value) { this.maxResponseRatePerSecond = value; return this; } + public long getMaxResponseBytesPerSecond() { return maxResponseBytesPerSecond; } + public DHTOptions setMaxResponseBytesPerSecond(long value) { this.maxResponseBytesPerSecond = value; return this; } + public int getMaxResponseRatePerSource() { return maxResponseRatePerSource; } + public DHTOptions setMaxResponseRatePerSource(int value) { this.maxResponseRatePerSource = value; return this; } + public int getMetadataPressureFloorPercent() { return metadataPressureFloorPercent; } + public DHTOptions setMetadataPressureFloorPercent(int value) { this.metadataPressureFloorPercent = value; return this; } + public long getRecentProbeTtlSeconds() { return recentProbeTtlSeconds; } + public DHTOptions setRecentProbeTtlSeconds(long value) { this.recentProbeTtlSeconds = value; return this; } + public int getResponsiveCapacity() { return responsiveCapacity; } + public DHTOptions setResponsiveCapacity(int value) { this.responsiveCapacity = value; return this; } + public long getResponsiveTtlSeconds() { return responsiveTtlSeconds; } + public DHTOptions setResponsiveTtlSeconds(long value) { this.responsiveTtlSeconds = value; return this; } + public int getPoolLowWatermark() { return poolLowWatermark; } + public DHTOptions setPoolLowWatermark(int value) { this.poolLowWatermark = value; return this; } + public int getMaxInFlightPerSubnet() { return maxInFlightPerSubnet; } + public DHTOptions setMaxInFlightPerSubnet(int value) { this.maxInFlightPerSubnet = value; return this; } + public int getHashQueueCapacity() { return hashQueueCapacity; } public DHTOptions setHashQueueCapacity(int hashQueueCapacity) { this.hashQueueCapacity = hashQueueCapacity; diff --git a/examples/main.rs b/examples/main.rs index ef205ba..55eb06d 100644 --- a/examples/main.rs +++ b/examples/main.rs @@ -34,11 +34,14 @@ async fn main() -> Result<()> { let options = DHTOptions { port: 12313, - metadata_timeout: 3, // ✅ 快速超时,快速失败 - max_metadata_queue_size: 100000, // ✅ 大缓冲区(防止饱和) - max_metadata_worker_count: 1000, // ✅ 激进并发(最大化吞吐) - netmode: NetMode::Ipv4Only, // 网络模式:Ipv4Only(仅IPv4)、Ipv6Only(仅IPv6)、DualStack(双栈,默认) - ..Default::default() // 使用默认值填充其他字段(节点队列容量等) + netmode: NetMode::Ipv4Only, + metadata: MetadataOptions { + timeout_secs: 4, + max_queue_size: 10_000, + max_worker_count: 256, + ..MetadataOptions::default() + }, + ..Default::default() }; // 统计计数器 @@ -83,8 +86,19 @@ async fn main() -> Result<()> { // 设置元数据获取前的检查回调 server.on_metadata_fetch(|_hash| async move { true }); + // 一个通过 gate 的 Hash 最终只会收到一次完成状态。 + server.on_metadata_fetch_complete(|completion| { + log::debug!( + "Metadata complete: hash={}, status={:?}, peer_attempts={}", + completion.info_hash, + completion.status, + completion.attempts + ); + }); + // 启动监控任务 let count_monitor = torrent_count.clone(); + let runtime_stats = server.runtime_stats(); tokio::spawn(async move { let mut interval = tokio::time::interval(std::time::Duration::from_secs(5)); let start_time = std::time::Instant::now(); @@ -93,12 +107,17 @@ async fn main() -> Result<()> { interval.tick().await; let success_fetch = count_monitor.load(Ordering::Relaxed); let uptime = start_time.elapsed().as_secs(); + let runtime = runtime_stats.snapshot(); // ✅ 监控:爬虫运行状态 log::info!( - "📊 [监控] 时长: {}s | 成功抓取: ✨ {}", + "📊 [监控] 时长: {}s | 成功抓取: ✨ {} | 节点: {} | Metadata: {}/{} | worker: {}", uptime, - success_fetch + success_fetch, + runtime.node_pool_size, + runtime.metadata_queue_depth, + runtime.metadata_queue_max, + runtime.metadata_in_flight, ); if uptime > 0 && success_fetch > 0 { @@ -108,6 +127,14 @@ async fn main() -> Result<()> { } }); + let shutdown_server = server.clone(); + tokio::spawn(async move { + if tokio::signal::ctrl_c().await.is_ok() { + log::info!("收到 Ctrl-C,正在停止 DHT Server"); + shutdown_server.shutdown(); + } + }); + server.start().await?; Ok(()) } diff --git a/jni/callbacks.rs b/jni/callbacks.rs index ff02f39..bbe0f7a 100644 --- a/jni/callbacks.rs +++ b/jni/callbacks.rs @@ -71,7 +71,7 @@ fn register_on_metadata_fetch(server: &Arc, callback: JavaCallback) { "(Ljava/lang/String;)Z", &[JValue::Object(&j_s.into())], )?; - Ok(out.z()?) + out.z() }) }) .await; diff --git a/jni/env.rs b/jni/env.rs index e2b4ff5..0426e2e 100644 --- a/jni/env.rs +++ b/jni/env.rs @@ -1,5 +1,5 @@ -use jni::{JavaVM, JNIEnv}; use jni::objects::JObject; +use jni::{JNIEnv, JavaVM}; use std::sync::{Arc, Mutex}; /// 跨线程安全的 Java 回调持有者。 @@ -33,10 +33,10 @@ impl JavaCallback { let listener_guard = self.listener.lock().unwrap(); let result = f(&mut guard, listener_guard.as_obj()); // 检查并清除 Java 侧遗留的异常,避免后续 JNI 调用受污染 - if let Err(ref _e) = result { - if guard.exception_check().unwrap_or(false) { - let _ = guard.exception_clear(); - } + if let Err(ref _e) = result + && guard.exception_check().unwrap_or(false) + { + let _ = guard.exception_clear(); } result } diff --git a/jni/exports.rs b/jni/exports.rs index 3624d9e..a42d851 100644 --- a/jni/exports.rs +++ b/jni/exports.rs @@ -1,6 +1,6 @@ use crate::jni_bindings::callbacks::register_callbacks; use crate::jni_bindings::env::JavaCallback; -use crate::jni_bindings::server::{handle_ref, into_handle_ptr, take_handle, ServerHandle}; +use crate::jni_bindings::server::{ServerHandle, handle_ref, into_handle_ptr, take_handle}; use crate::jni_bindings::types::java_to_dht_options_or_default; use jni::JNIEnv; use jni::objects::{JClass, JObject}; @@ -29,7 +29,7 @@ pub extern "system" fn Java_cn_lmcw_dht_DhtCrawlerJni_createServer( let opts = match java_to_dht_options_or_default(&mut env, &options) { Ok(o) => o, Err(e) => { - let _ = env.throw_new("java/lang/IllegalArgumentException", &e.to_string()); + let _ = env.throw_new("java/lang/IllegalArgumentException", e.to_string()); return 0; } }; @@ -48,7 +48,7 @@ pub extern "system" fn Java_cn_lmcw_dht_DhtCrawlerJni_createServer( match JavaCallback::new(&mut env, &listener) { Ok(cb) => register_callbacks(&handle.server, cb), Err(e) => { - let _ = env.throw_new("java/lang/RuntimeException", &e.to_string()); + let _ = env.throw_new("java/lang/RuntimeException", e.to_string()); return 0; } } diff --git a/jni/server.rs b/jni/server.rs index 45dabf8..a618eb5 100644 --- a/jni/server.rs +++ b/jni/server.rs @@ -1,4 +1,4 @@ -use crate::{DHTServer, DHTOptions}; +use crate::{DHTOptions, DHTServer}; use std::sync::Arc; use tokio::runtime::Runtime; diff --git a/jni/types.rs b/jni/types.rs index 5147ed1..556d79d 100644 --- a/jni/types.rs +++ b/jni/types.rs @@ -1,4 +1,4 @@ -use crate::{TorrentInfo, FileInfo, DHTOptions, types::NetMode}; +use crate::{DHTOptions, FileInfo, MetadataOptions, TorrentInfo, types::NetMode}; use jni::JNIEnv; use jni::objects::{JObject, JString, JValue}; use jni::sys::jlong; @@ -90,7 +90,12 @@ fn build_file_list<'local>( let list = env.new_object(&list_cls, "()V", &[])?; for fi in files { let jfi = file_info_to_java(env, fi)?; - env.call_method(&list, "add", "(Ljava/lang/Object;)Z", &[JValue::Object(&jfi)])?; + env.call_method( + &list, + "add", + "(Ljava/lang/Object;)Z", + &[JValue::Object(&jfi)], + )?; env.delete_local_ref(jfi)?; } Ok(list) @@ -105,7 +110,12 @@ fn build_string_list<'local>( let list = env.new_object(&list_cls, "()V", &[])?; for s in strs { let js = rust_str_to_jstring(env, s)?; - env.call_method(&list, "add", "(Ljava/lang/Object;)Z", &[JValue::Object(&js)])?; + env.call_method( + &list, + "add", + "(Ljava/lang/Object;)Z", + &[JValue::Object(&js)], + )?; env.delete_local_ref(js)?; } Ok(list) @@ -118,15 +128,31 @@ fn build_string_list<'local>( /// 从 Java `cn.lmcw.dht.model.DHTOptions` 对象读取字段,构造 Rust `DHTOptions`。 pub fn java_to_dht_options(env: &mut JNIEnv, obj: &JObject) -> jni::errors::Result { let port = env.get_field(obj, "port", "I")?.i()? as u16; - let metadata_timeout = env.get_field(obj, "metadataTimeout", "J")?.j()? as u64; - let max_metadata_queue_size = - env.get_field(obj, "maxMetadataQueueSize", "I")?.i()? as usize; - let max_metadata_worker_count = + let metadata_timeout_secs = env.get_field(obj, "metadataTimeout", "J")?.j()? as u64; + let metadata_max_queue_size = env.get_field(obj, "maxMetadataQueueSize", "I")?.i()? as usize; + let metadata_max_worker_count = env.get_field(obj, "maxMetadataWorkerCount", "I")?.i()? as usize; - let node_queue_capacity = - env.get_field(obj, "nodeQueueCapacity", "I")?.i()? as usize; - let hash_queue_capacity = - env.get_field(obj, "hashQueueCapacity", "I")?.i()? as usize; + let pool_capacity = env.get_field(obj, "poolCapacity", "I")?.i()? as usize; + let find_node_rate = env.get_field(obj, "findNodeRatePerSecond", "I")?.i()? as u32; + let find_node_burst = env.get_field(obj, "findNodeBurst", "I")?.i()? as u32; + let max_find_node_in_flight = env.get_field(obj, "maxFindNodeInFlight", "I")?.i()? as usize; + let max_new_destinations = env + .get_field(obj, "maxNewDestinationsPerMinute", "I")? + .i()? as u32; + let max_replacements = env.get_field(obj, "maxReplacementsPerMinute", "I")?.i()? as u32; + let request_timeout_secs = env.get_field(obj, "requestTimeoutSeconds", "J")?.j()? as u64; + let max_response_rate = env.get_field(obj, "maxResponseRatePerSecond", "I")?.i()? as u32; + let max_response_bytes = env.get_field(obj, "maxResponseBytesPerSecond", "J")?.j()? as u64; + let max_response_per_source = env.get_field(obj, "maxResponseRatePerSource", "I")?.i()? as u32; + let pressure_floor = env + .get_field(obj, "metadataPressureFloorPercent", "I")? + .i()? as u8; + let recent_probe_ttl = env.get_field(obj, "recentProbeTtlSeconds", "J")?.j()? as u64; + let responsive_capacity = env.get_field(obj, "responsiveCapacity", "I")?.i()? as usize; + let responsive_ttl = env.get_field(obj, "responsiveTtlSeconds", "J")?.j()? as u64; + let low_watermark = env.get_field(obj, "poolLowWatermark", "I")?.i()? as usize; + let subnet_in_flight = env.get_field(obj, "maxInFlightPerSubnet", "I")?.i()? as usize; + let hash_queue_capacity = env.get_field(obj, "hashQueueCapacity", "I")?.i()? as usize; let netmode_ord = env.get_field(obj, "netMode", "I")?.i()?; let netmode = match netmode_ord { 0 => NetMode::Ipv4Only, @@ -134,15 +160,35 @@ pub fn java_to_dht_options(env: &mut JNIEnv, obj: &JObject) -> jni::errors::Resu _ => NetMode::DualStack, }; - Ok(DHTOptions { + let mut options = DHTOptions { port, - metadata_timeout, - max_metadata_queue_size, - max_metadata_worker_count, netmode, - node_queue_capacity, hash_queue_capacity, - }) + metadata: MetadataOptions { + timeout_secs: metadata_timeout_secs, + max_queue_size: metadata_max_queue_size, + max_worker_count: metadata_max_worker_count, + ..MetadataOptions::default() + }, + ..DHTOptions::default() + }; + options.crawl.pool.capacity = pool_capacity.max(1); + options.crawl.pool.recent_probe_ttl_secs = recent_probe_ttl; + options.crawl.pool.responsive_capacity = responsive_capacity.max(1); + options.crawl.pool.responsive_ttl_secs = responsive_ttl; + options.crawl.pool.low_watermark = low_watermark.min(pool_capacity); + options.crawl.rate_limit.max_find_node_rate_per_sec = find_node_rate; + options.crawl.rate_limit.burst = find_node_burst; + options.crawl.rate_limit.max_in_flight = max_find_node_in_flight.max(1); + options.crawl.rate_limit.max_new_destinations_per_minute = max_new_destinations; + options.crawl.rate_limit.request_timeout_secs = request_timeout_secs; + options.crawl.rate_limit.max_response_rate_per_sec = max_response_rate; + options.crawl.rate_limit.max_response_bytes_per_sec = max_response_bytes; + options.crawl.rate_limit.max_response_rate_per_source = max_response_per_source; + options.crawl.rate_limit.metadata_pressure_floor_percent = pressure_floor.min(100); + options.crawl.rate_limit.max_replacements_per_minute = max_replacements; + options.crawl.rate_limit.max_in_flight_per_subnet = subnet_in_flight.max(1); + Ok(options) } /// 从 Java `cn.lmcw.dht.model.DHTOptions` 对象读取,或若为 null 则返回默认选项。 diff --git a/src/addr.rs b/src/addr.rs new file mode 100644 index 0000000..42a79d1 --- /dev/null +++ b/src/addr.rs @@ -0,0 +1,71 @@ +use crate::types::NetMode; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; + +pub(crate) fn addr_allowed_by_netmode(addr: &SocketAddr, netmode: NetMode) -> bool { + match netmode { + NetMode::Ipv4Only => addr.is_ipv4(), + NetMode::Ipv6Only => addr.is_ipv6(), + NetMode::DualStack => true, + } +} + +pub(crate) fn is_valid_node_addr(addr: &SocketAddr) -> bool { + if addr.port() == 0 { + return false; + } + + match addr.ip() { + IpAddr::V4(ip) => is_valid_ipv4_node_addr(ip), + IpAddr::V6(ip) => is_valid_ipv6_node_addr(ip), + } +} + +fn is_valid_ipv4_node_addr(ip: Ipv4Addr) -> bool { + let octets = ip.octets(); + let is_cgnat = octets[0] == 100 && (octets[1] & 0b1100_0000) == 64; + let is_benchmark = octets[0] == 198 && (octets[1] == 18 || octets[1] == 19); + let is_reserved = octets[0] >= 240; + + !ip.is_unspecified() + && !ip.is_loopback() + && !ip.is_private() + && !ip.is_link_local() + && !ip.is_multicast() + && !ip.is_broadcast() + && !ip.is_documentation() + && !is_cgnat + && !is_benchmark + && !is_reserved +} + +fn is_valid_ipv6_node_addr(ip: Ipv6Addr) -> bool { + let octets = ip.octets(); + let is_unique_local = (octets[0] & 0xfe) == 0xfc; + let is_unicast_link_local = octets[0] == 0xfe && (octets[1] & 0xc0) == 0x80; + let is_documentation = + octets[0] == 0x20 && octets[1] == 0x01 && octets[2] == 0x0d && octets[3] == 0xb8; + + !ip.is_unspecified() + && !ip.is_loopback() + && !ip.is_multicast() + && !is_unique_local + && !is_unicast_link_local + && !is_documentation +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn invalid_node_addresses_are_filtered() { + assert!(is_valid_node_addr(&"8.8.8.8:6881".parse().unwrap())); + assert!(is_valid_node_addr( + &"[2001:4860:4860::8888]:6881".parse().unwrap() + )); + assert!(!is_valid_node_addr(&"8.8.8.8:0".parse().unwrap())); + assert!(!is_valid_node_addr(&"10.0.0.1:6881".parse().unwrap())); + assert!(!is_valid_node_addr(&"127.0.0.1:6881".parse().unwrap())); + assert!(!is_valid_node_addr(&"[fc00::1]:6881".parse().unwrap())); + } +} diff --git a/src/bootstrap.rs b/src/bootstrap.rs new file mode 100644 index 0000000..e9547cf --- /dev/null +++ b/src/bootstrap.rs @@ -0,0 +1,221 @@ +use crate::addr::{addr_allowed_by_netmode, is_valid_node_addr}; +use crate::crawl_config::ResolvedCrawlConfig; +use crate::types::NetMode; +use std::collections::{HashMap, HashSet}; +use std::net::SocketAddr; +use std::time::{Duration, Instant}; + +pub(crate) struct BootstrapGate { + last_bootstrap: Option, +} + +impl BootstrapGate { + pub(crate) fn new() -> Self { + Self { + last_bootstrap: None, + } + } + + pub(crate) fn should_bootstrap( + &mut self, + pool_len: usize, + config: &ResolvedCrawlConfig, + now: Instant, + ) -> bool { + if config.bootstrap_max_nodes_per_round == 0 { + return false; + } + if pool_len >= config.low_watermark { + return false; + } + if let Some(last) = self.last_bootstrap + && now.checked_duration_since(last).unwrap_or_default() < config.bootstrap_interval + { + return false; + } + self.last_bootstrap = Some(now); + true + } +} + +#[derive(Default)] +struct BootstrapSourceState { + last_attempt: Option, + last_success: Option, + fail_count: u32, + backoff_until: Option, +} + +pub(crate) struct BootstrapSourcePool { + pub(crate) hosts: Vec, + states: HashMap, + backoff_base: Duration, + backoff_max: Duration, +} + +impl BootstrapSourcePool { + pub(crate) fn new(hosts: Vec, backoff_base: Duration, backoff_max: Duration) -> Self { + Self { + hosts, + states: HashMap::new(), + backoff_base, + backoff_max, + } + } + + pub(crate) fn select( + &mut self, + candidates: Vec, + max_nodes: usize, + now: Instant, + ) -> Vec { + let mut selected = Vec::with_capacity(max_nodes); + let mut seen = HashSet::with_capacity(candidates.len()); + let mut earliest_backoff: Option<(SocketAddr, Instant)> = None; + + for addr in candidates { + if !seen.insert(addr) { + continue; + } + + let state = self.states.entry(addr).or_default(); + if let Some(backoff_until) = state.backoff_until + && backoff_until > now + { + if earliest_backoff.is_none_or(|(_, current)| backoff_until < current) { + earliest_backoff = Some((addr, backoff_until)); + } + continue; + } + + selected.push(addr); + if selected.len() >= max_nodes { + return selected; + } + } + + if selected.is_empty() + && max_nodes > 0 + && let Some((addr, _)) = earliest_backoff + { + selected.push(addr); + } + selected + } + + pub(crate) fn mark_attempt(&mut self, addr: SocketAddr, now: Instant) { + self.states.entry(addr).or_default().last_attempt = Some(now); + } + + pub(crate) fn mark_success(&mut self, addr: SocketAddr, now: Instant) { + let state = self.states.entry(addr).or_default(); + state.last_success = Some(now); + state.fail_count = 0; + state.backoff_until = None; + } + + pub(crate) fn mark_timeout(&mut self, addr: SocketAddr, now: Instant) { + let state = self.states.entry(addr).or_default(); + state.fail_count = state.fail_count.saturating_add(1); + let multiplier = 1u32 + .checked_shl(state.fail_count.saturating_sub(1).min(16)) + .unwrap_or(u32::MAX); + let backoff = self + .backoff_base + .saturating_mul(multiplier) + .min(self.backoff_max); + state.backoff_until = Some(now + backoff); + } +} + +pub(crate) async fn resolve_bootstrap_nodes(hosts: &[String], netmode: NetMode) -> Vec { + let mut resolved = Vec::new(); + for host in hosts { + if let Ok(addrs) = tokio::net::lookup_host(host).await { + for addr in addrs { + if !addr_allowed_by_netmode(&addr, netmode) || !is_valid_node_addr(&addr) { + continue; + } + resolved.push(addr); + } + } + } + resolved +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::types::CrawlOptions; + + fn test_config() -> ResolvedCrawlConfig { + ResolvedCrawlConfig::from_options(&CrawlOptions::default()) + } + + #[test] + fn bootstrap_pool_backs_off_dead_sources_without_spending_quota() { + let start = Instant::now(); + let addr1: SocketAddr = "8.8.8.8:6881".parse().unwrap(); + let addr2: SocketAddr = "1.1.1.1:6881".parse().unwrap(); + let mut pool = BootstrapSourcePool::new( + vec!["example.invalid:6881".to_string()], + Duration::from_secs(300), + Duration::from_secs(3600), + ); + + pool.mark_timeout(addr1, start); + let selected = pool.select(vec![addr1, addr2], 1, start + Duration::from_secs(1)); + assert_eq!(selected, vec![addr2]); + + pool.mark_success(addr1, start + Duration::from_secs(2)); + let selected = pool.select(vec![addr1], 1, start + Duration::from_secs(3)); + assert_eq!(selected, vec![addr1]); + } + + #[test] + fn bootstrap_pool_forces_one_retry_when_all_sources_backed_off() { + let start = Instant::now(); + let addr1: SocketAddr = "8.8.8.8:6881".parse().unwrap(); + let addr2: SocketAddr = "1.1.1.1:6881".parse().unwrap(); + let mut pool = BootstrapSourcePool::new( + vec!["example.invalid:6881".to_string()], + Duration::from_secs(300), + Duration::from_secs(3600), + ); + pool.mark_timeout(addr1, start); + pool.mark_timeout(addr2, start); + + let selected = pool.select(vec![addr1, addr2], 3, start + Duration::from_secs(1)); + assert_eq!(selected.len(), 1); + } + + #[test] + fn bootstrap_pool_deduplicates_resolved_addresses() { + let start = Instant::now(); + let addr: SocketAddr = "8.8.8.8:6881".parse().unwrap(); + let mut pool = BootstrapSourcePool::new( + vec!["example.invalid:6881".to_string()], + Duration::from_secs(300), + Duration::from_secs(3600), + ); + + let selected = pool.select(vec![addr, addr], 10, start); + assert_eq!(selected, vec![addr]); + } + + #[test] + fn bootstrap_gate_uses_pool_low_water_mark() { + let start = Instant::now(); + let config = test_config(); + let mut gate = BootstrapGate::new(); + + assert!(!gate.should_bootstrap(config.low_watermark, &config, start)); + assert!(gate.should_bootstrap(0, &config, start)); + assert!(!gate.should_bootstrap( + 999, + &config, + start + config.bootstrap_interval - Duration::from_secs(1) + )); + assert!(gate.should_bootstrap(999, &config, start + config.bootstrap_interval)); + } +} diff --git a/src/budget.rs b/src/budget.rs new file mode 100644 index 0000000..f76bbe8 --- /dev/null +++ b/src/budget.rs @@ -0,0 +1,132 @@ +use std::time::{Duration, Instant}; + +/// Single-owner token bucket. It deliberately contains no atomics or locks. +pub(crate) struct RateBucket { + rate_per_sec: f64, + capacity: f64, + tokens: f64, + last_refill: Instant, +} + +impl RateBucket { + pub(crate) fn per_second( + rate_per_sec: u32, + burst: u32, + initially_full: bool, + now: Instant, + ) -> Self { + Self::new( + f64::from(rate_per_sec), + f64::from(burst), + initially_full, + now, + ) + } + + pub(crate) fn per_minute( + rate_per_minute: u32, + burst: u32, + initially_full: bool, + now: Instant, + ) -> Self { + Self::new( + f64::from(rate_per_minute) / 60.0, + f64::from(burst), + initially_full, + now, + ) + } + + fn new(rate_per_sec: f64, capacity: f64, initially_full: bool, now: Instant) -> Self { + let capacity = if rate_per_sec <= 0.0 { + 0.0 + } else { + capacity.max(1.0) + }; + Self { + rate_per_sec, + capacity, + tokens: if initially_full { capacity } else { 0.0 }, + last_refill: now, + } + } + + pub(crate) fn set_per_second_rate(&mut self, rate_per_sec: u32, now: Instant) { + self.refill(now); + self.rate_per_sec = f64::from(rate_per_sec); + if self.rate_per_sec <= 0.0 { + self.tokens = 0.0; + } else { + self.tokens = self.tokens.min(self.capacity); + } + } + + pub(crate) fn try_take_one(&mut self, now: Instant) -> bool { + self.try_take_exact(1, now) + } + + pub(crate) fn try_take_exact(&mut self, count: usize, now: Instant) -> bool { + if count == 0 { + return true; + } + if self.rate_per_sec <= 0.0 || self.capacity <= 0.0 { + return false; + } + self.refill(now); + if self.tokens < count as f64 { + return false; + } + self.tokens -= count as f64; + true + } + + pub(crate) fn try_take(&mut self, max: usize, now: Instant) -> usize { + if max == 0 || self.rate_per_sec <= 0.0 || self.capacity <= 0.0 { + return 0; + } + self.refill(now); + let taken = (self.tokens.floor() as usize).min(max); + self.tokens -= taken as f64; + taken + } + + pub(crate) fn refund_one(&mut self) { + self.refund(1); + } + + pub(crate) fn refund(&mut self, count: usize) { + self.tokens = (self.tokens + count as f64).min(self.capacity); + } + + fn refill(&mut self, now: Instant) { + let elapsed = now + .checked_duration_since(self.last_refill) + .unwrap_or(Duration::ZERO) + .as_secs_f64(); + if elapsed > 0.0 { + self.tokens = (self.tokens + elapsed * self.rate_per_sec).min(self.capacity); + self.last_refill = now; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn second_bucket_is_smooth_and_bounded() { + let start = Instant::now(); + let mut bucket = RateBucket::per_second(100, 20, false, start); + assert_eq!(bucket.try_take(100, start), 0); + assert_eq!(bucket.try_take(100, start + Duration::from_millis(100)), 10); + assert_eq!(bucket.try_take(100, start + Duration::from_secs(10)), 20); + } + + #[test] + fn minute_bucket_refills_fractionally() { + let start = Instant::now(); + let mut bucket = RateBucket::per_minute(600, 10, false, start); + assert_eq!(bucket.try_take(100, start + Duration::from_millis(500)), 5); + } +} diff --git a/src/crawl_config.rs b/src/crawl_config.rs new file mode 100644 index 0000000..e622ed4 --- /dev/null +++ b/src/crawl_config.rs @@ -0,0 +1,137 @@ +use crate::types::CrawlOptions; +use std::time::Duration; + +#[derive(Debug, Clone)] +pub(crate) struct ResolvedCrawlConfig { + pub(crate) max_find_node_rate_per_sec: u32, + pub(crate) burst: u32, + pub(crate) max_in_flight: usize, + pub(crate) request_timeout: Duration, + pub(crate) max_new_destinations_per_minute: u32, + pub(crate) max_response_rate_per_sec: u32, + pub(crate) max_response_bytes_per_sec: u64, + pub(crate) max_response_rate_per_source: u32, + pub(crate) metadata_pressure_floor_percent: u8, + + pub(crate) pool_capacity: usize, + pub(crate) max_replacements_per_minute: u32, + pub(crate) recent_probe_ttl: Duration, + pub(crate) responsive_capacity: usize, + pub(crate) responsive_ttl: Duration, + pub(crate) low_watermark: usize, + pub(crate) max_in_flight_per_subnet: usize, + + pub(crate) bootstrap_nodes: Vec, + pub(crate) bootstrap_interval: Duration, + pub(crate) bootstrap_max_nodes_per_round: usize, + pub(crate) bootstrap_backoff_base: Duration, + pub(crate) bootstrap_backoff_max: Duration, + + pub(crate) random_walk_percent: u8, + pub(crate) sparse_bucket_percent: u8, + pub(crate) neighbor_sender_id: bool, + + pub(crate) priority_event_channel_capacity: usize, + pub(crate) discovery_event_channel_capacity: usize, + pub(crate) event_batch_limit: usize, + pub(crate) node_batch_limit: usize, + pub(crate) routing_snapshot_size: usize, + pub(crate) snapshot_refresh: Duration, +} + +impl ResolvedCrawlConfig { + pub(crate) fn from_options(options: &CrawlOptions) -> Self { + let capacity = options.pool.capacity.max(1); + Self { + max_find_node_rate_per_sec: options.rate_limit.max_find_node_rate_per_sec, + burst: if options.rate_limit.max_find_node_rate_per_sec == 0 { + 0 + } else { + options.rate_limit.burst.max(1) + }, + max_in_flight: options.rate_limit.max_in_flight.max(1), + request_timeout: Duration::from_secs(options.rate_limit.request_timeout_secs.max(1)), + max_new_destinations_per_minute: options.rate_limit.max_new_destinations_per_minute, + max_response_rate_per_sec: options.rate_limit.max_response_rate_per_sec, + max_response_bytes_per_sec: options.rate_limit.max_response_bytes_per_sec, + max_response_rate_per_source: options.rate_limit.max_response_rate_per_source, + metadata_pressure_floor_percent: options + .rate_limit + .metadata_pressure_floor_percent + .min(100), + + pool_capacity: capacity, + max_replacements_per_minute: options.rate_limit.max_replacements_per_minute, + recent_probe_ttl: Duration::from_secs(options.pool.recent_probe_ttl_secs.max(1)), + responsive_capacity: options.pool.responsive_capacity.max(1), + responsive_ttl: Duration::from_secs(options.pool.responsive_ttl_secs.max(1)), + low_watermark: options.pool.low_watermark.min(capacity), + max_in_flight_per_subnet: options.rate_limit.max_in_flight_per_subnet.max(1), + + bootstrap_nodes: if options.bootstrap.nodes.is_empty() { + crate::types::BootstrapOptions::default().nodes + } else { + options.bootstrap.nodes.clone() + }, + bootstrap_interval: Duration::from_secs(options.bootstrap.interval_secs), + bootstrap_max_nodes_per_round: options.bootstrap.max_nodes_per_round, + bootstrap_backoff_base: Duration::from_secs( + options.bootstrap.source_backoff_base_secs.max(1), + ), + bootstrap_backoff_max: Duration::from_secs( + options.bootstrap.source_backoff_max_secs.max(1), + ), + + random_walk_percent: options.target.random_walk_percent.min(100), + sparse_bucket_percent: options.target.sparse_bucket_percent.min(100), + neighbor_sender_id: options.target.neighbor_sender_id, + + priority_event_channel_capacity: options + .scheduler + .priority_event_channel_capacity + .max(1), + discovery_event_channel_capacity: options + .scheduler + .discovery_event_channel_capacity + .max(1), + event_batch_limit: options.scheduler.event_batch_limit.max(1), + node_batch_limit: options.scheduler.node_batch_limit.max(1), + routing_snapshot_size: options.scheduler.routing_snapshot_size.max(1), + snapshot_refresh: Duration::from_millis( + options.scheduler.snapshot_refresh_millis.max(100), + ), + } + } + + pub(crate) fn rate_for_metadata_pressure(&self, pressure: f64) -> u32 { + let max = f64::from(self.max_find_node_rate_per_sec); + if pressure < 0.80 { + return self.max_find_node_rate_per_sec; + } + let floor = max * f64::from(self.metadata_pressure_floor_percent) / 100.0; + if pressure >= 0.95 { + return floor.round() as u32; + } + let progress = ((pressure - 0.80) / 0.15).clamp(0.0, 1.0); + (max - (max - floor) * progress).round() as u32 + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn resolves_pool_and_rate_limit() { + let mut options = CrawlOptions::default(); + options.rate_limit.max_find_node_rate_per_sec = 200; + options.rate_limit.metadata_pressure_floor_percent = 25; + options.pool.capacity = 42; + let resolved = ResolvedCrawlConfig::from_options(&options); + + assert_eq!(resolved.pool_capacity, 42); + assert_eq!(resolved.rate_for_metadata_pressure(0.79), 200); + assert_eq!(resolved.rate_for_metadata_pressure(0.95), 50); + assert_eq!(resolved.rate_for_metadata_pressure(1.0), 50); + } +} diff --git a/src/crawl_engine.rs b/src/crawl_engine.rs new file mode 100644 index 0000000..747f7dd --- /dev/null +++ b/src/crawl_engine.rs @@ -0,0 +1,928 @@ +use crate::bootstrap::{BootstrapGate, BootstrapSourcePool, resolve_bootstrap_nodes}; +use crate::budget::RateBucket; +use crate::crawl_config::ResolvedCrawlConfig; +use crate::krpc::{for_each_response_node, send_find_node_query}; +use crate::node_id::{ + TransactionId, bucket_index, neighbor_node_id, random_node_id, target_for_bucket, +}; +use crate::node_pool::{AdmissionOutcome, NodePool, ResponsiveReservoir, SubnetKey}; +use crate::protocol::DhtResponse; +use crate::routing_snapshot::RoutingSnapshot; +#[cfg(test)] +use crate::runtime_stats::DhtRuntimeLimits; +use crate::runtime_stats::DhtRuntimeStats; +use crate::types::{NetMode, NodeTuple}; +use ahash::AHashMap; +use arc_swap::ArcSwap; +use bytes::BytesMut; +#[cfg(feature = "metrics")] +use metrics::{counter, gauge}; +use rand::Rng; +use std::collections::{HashMap, VecDeque}; +use std::net::SocketAddr; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; +use tokio::net::UdpSocket; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +const SCHEDULE_INTERVAL: Duration = Duration::from_millis(5); +const MAX_POOL_SCAN_PER_SCHEDULE: usize = 256; + +#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)] +struct PendingKey { + addr: SocketAddr, + tid: TransactionId, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ProbePurpose { + New, + Revisit, + Bootstrap, +} + +#[derive(Debug, Clone, Copy)] +struct PendingRequest { + node: NodeTuple, + purpose: ProbePurpose, + deadline: Instant, + subnet: SubnetKey, +} + +enum PriorityEvent { + Response { + remote_addr: SocketAddr, + tid: TransactionId, + response: DhtResponse, + }, + SendFailed(PendingKey), + BootstrapResolved(Vec), +} + +struct OutboundRequest { + key: PendingKey, + node: NodeTuple, + target: [u8; 20], + sender_id: [u8; 20], +} + +pub(crate) struct CrawlEngine { + config: ResolvedCrawlConfig, + priority_tx: mpsc::Sender, + discovery_tx: mpsc::Sender, + /// One-shot handoff used only by `DHTServer::start`; never touched by the crawl hot path. + receivers: Mutex, mpsc::Receiver)>>, + pub(crate) snapshot: Arc>, + pub(crate) node_count: Arc, + runtime_stats: DhtRuntimeStats, +} + +impl CrawlEngine { + pub(crate) fn new(config: ResolvedCrawlConfig, runtime_stats: DhtRuntimeStats) -> Self { + let (priority_tx, priority_rx) = mpsc::channel(config.priority_event_channel_capacity); + let (discovery_tx, discovery_rx) = mpsc::channel(config.discovery_event_channel_capacity); + Self { + config, + priority_tx, + discovery_tx, + receivers: Mutex::new(Some((priority_rx, discovery_rx))), + snapshot: Arc::new(ArcSwap::from_pointee(RoutingSnapshot::default())), + node_count: Arc::new(AtomicUsize::new(0)), + runtime_stats, + } + } + + pub(crate) fn route_discovered(&self, node: NodeTuple) { + let enqueue_result = self.discovery_tx.try_send(node); + self.runtime_stats.set_crawl_discovery_queue_depth( + self.discovery_tx + .max_capacity() + .saturating_sub(self.discovery_tx.capacity()), + ); + if enqueue_result.is_err() { + self.runtime_stats.crawl_event_dropped_discovered(); + #[cfg(feature = "metrics")] + counter!("dht_crawl_events_dropped_total", "kind" => "discovered").increment(1); + } + } + + pub(crate) fn route_response( + &self, + remote_addr: SocketAddr, + tid: TransactionId, + response: DhtResponse, + ) { + let enqueue_result = self.priority_tx.try_send(PriorityEvent::Response { + remote_addr, + tid, + response, + }); + self.runtime_stats.set_crawl_priority_queue_depth( + self.priority_tx + .max_capacity() + .saturating_sub(self.priority_tx.capacity()), + ); + if enqueue_result.is_err() { + self.runtime_stats.crawl_event_dropped_response(); + #[cfg(feature = "metrics")] + counter!("dht_crawl_events_dropped_total", "kind" => "response").increment(1); + } + } + + pub(crate) fn spawn( + &self, + netmode: NetMode, + local_id: [u8; 20], + sockets: &HashMap>, + metadata_queue_len: Arc, + max_metadata_queue_size: usize, + shutdown: CancellationToken, + ) { + let Some((priority_rx, discovery_rx)) = self + .receivers + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .take() + else { + return; + }; + + let mut egress_v4 = None; + let mut egress_v6 = None; + for (bind_addr, socket) in sockets.iter() { + let (tx, rx) = mpsc::channel(self.config.max_in_flight.max(1)); + spawn_egress( + socket.clone(), + rx, + self.priority_tx.clone(), + self.runtime_stats.clone(), + shutdown.clone(), + ); + if bind_addr.is_ipv4() { + egress_v4 = Some(tx); + } else { + egress_v6 = Some(tx); + } + } + + let actor = CrawlActor::new(CrawlActorInit { + config: self.config.clone(), + netmode, + local_id, + priority_rx, + discovery_rx, + priority_tx: self.priority_tx.clone(), + egress_v4, + egress_v6, + snapshot: self.snapshot.clone(), + node_count: self.node_count.clone(), + metadata_queue_len, + max_metadata_queue_size, + runtime_stats: self.runtime_stats.clone(), + shutdown, + }); + tokio::spawn(actor.run()); + } +} + +fn spawn_egress( + socket: Arc, + mut rx: mpsc::Receiver, + priority_tx: mpsc::Sender, + runtime_stats: DhtRuntimeStats, + shutdown: CancellationToken, +) { + tokio::spawn(async move { + let mut buffer = BytesMut::with_capacity(128); + loop { + tokio::select! { + _ = shutdown.cancelled() => break, + request = rx.recv() => { + let Some(request) = request else { break }; + if !send_find_node_query( + &request.node.addr, + &request.key.tid, + &request.target, + &request.sender_id, + &socket, + &mut buffer, + ).await { + let _ = priority_tx.try_send(PriorityEvent::SendFailed(request.key)); + runtime_stats.set_crawl_priority_queue_depth( + priority_tx + .max_capacity() + .saturating_sub(priority_tx.capacity()), + ); + } else { + runtime_stats.udp_sent(buffer.len()); + } + } + } + } + }); +} + +#[derive(Default)] +struct ActorMetrics { + admitted: u64, + replaced: u64, + duplicate: u64, + admission_limited: u64, + invalid: u64, + queries_new: u64, + queries_revisit: u64, + queries_bootstrap: u64, + responses: u64, + timeouts: u64, + send_failures: u64, + unmatched_responses: u64, +} + +impl ActorMetrics { + fn record_admission(&mut self, outcome: AdmissionOutcome) { + match outcome { + AdmissionOutcome::Admitted => self.admitted += 1, + AdmissionOutcome::Replaced => self.replaced += 1, + AdmissionOutcome::Duplicate => self.duplicate += 1, + AdmissionOutcome::RateLimited => self.admission_limited += 1, + AdmissionOutcome::Invalid => self.invalid += 1, + } + } +} + +fn record_runtime_admission(stats: &DhtRuntimeStats, outcome: AdmissionOutcome) { + match outcome { + AdmissionOutcome::Admitted => stats.node_admitted(), + AdmissionOutcome::Replaced => stats.node_replaced(), + AdmissionOutcome::Duplicate => stats.node_dropped_duplicate(), + AdmissionOutcome::RateLimited => stats.node_dropped_rate_limited(), + AdmissionOutcome::Invalid => stats.node_dropped_invalid(), + } +} + +struct CrawlActorInit { + config: ResolvedCrawlConfig, + netmode: NetMode, + local_id: [u8; 20], + priority_rx: mpsc::Receiver, + discovery_rx: mpsc::Receiver, + priority_tx: mpsc::Sender, + egress_v4: Option>, + egress_v6: Option>, + snapshot: Arc>, + node_count: Arc, + metadata_queue_len: Arc, + max_metadata_queue_size: usize, + runtime_stats: DhtRuntimeStats, + shutdown: CancellationToken, +} + +struct CrawlActor { + config: ResolvedCrawlConfig, + netmode: NetMode, + local_id: [u8; 20], + priority_rx: mpsc::Receiver, + discovery_rx: mpsc::Receiver, + priority_tx: mpsc::Sender, + egress_v4: Option>, + egress_v6: Option>, + pool: NodePool, + responsive: ResponsiveReservoir, + pending: AHashMap, + pending_expiry: VecDeque<(Instant, PendingKey)>, + subnet_in_flight: AHashMap, + query_budget: RateBucket, + destination_budget: RateBucket, + bootstrap_gate: BootstrapGate, + bootstrap_pool: BootstrapSourcePool, + bootstrap_queue: VecDeque, + snapshot: Arc>, + node_count: Arc, + metadata_queue_len: Arc, + max_metadata_queue_size: usize, + next_tid: u64, + metrics: ActorMetrics, + runtime_stats: DhtRuntimeStats, + shutdown: CancellationToken, +} + +impl CrawlActor { + fn new(init: CrawlActorInit) -> Self { + let CrawlActorInit { + config, + netmode, + local_id, + priority_rx, + discovery_rx, + priority_tx, + egress_v4, + egress_v6, + snapshot, + node_count, + metadata_queue_len, + max_metadata_queue_size, + runtime_stats, + shutdown, + } = init; + let now = Instant::now(); + let destination_burst = config.max_new_destinations_per_minute.div_ceil(60).max(1); + Self { + pool: NodePool::new( + config.pool_capacity, + config.max_replacements_per_minute, + config.recent_probe_ttl, + now, + ), + responsive: ResponsiveReservoir::new(config.responsive_capacity, config.responsive_ttl), + pending: AHashMap::with_capacity(config.max_in_flight), + pending_expiry: VecDeque::with_capacity(config.max_in_flight), + subnet_in_flight: AHashMap::new(), + query_budget: RateBucket::per_second( + config.max_find_node_rate_per_sec, + config.burst, + false, + now, + ), + destination_budget: RateBucket::per_minute( + config.max_new_destinations_per_minute, + destination_burst, + false, + now, + ), + bootstrap_pool: BootstrapSourcePool::new( + config.bootstrap_nodes.clone(), + config.bootstrap_backoff_base, + config.bootstrap_backoff_max, + ), + config, + netmode, + local_id, + priority_rx, + discovery_rx, + priority_tx, + egress_v4, + egress_v6, + bootstrap_gate: BootstrapGate::new(), + bootstrap_queue: VecDeque::new(), + snapshot, + node_count, + metadata_queue_len, + max_metadata_queue_size, + next_tid: 1, + metrics: ActorMetrics::default(), + runtime_stats, + shutdown, + } + } + + async fn run(mut self) { + let mut schedule_tick = tokio::time::interval(SCHEDULE_INTERVAL); + schedule_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut snapshot_tick = tokio::time::interval(self.config.snapshot_refresh); + snapshot_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut metrics_tick = tokio::time::interval(Duration::from_secs(1)); + metrics_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + + loop { + tokio::select! { + biased; + _ = self.shutdown.cancelled() => break, + _ = schedule_tick.tick() => self.on_schedule_tick(Instant::now()), + _ = snapshot_tick.tick() => self.publish_snapshot(Instant::now()), + _ = metrics_tick.tick() => self.flush_metrics(Instant::now()), + event = self.priority_rx.recv() => { + self.runtime_stats + .set_crawl_priority_queue_depth(self.priority_rx.len()); + let Some(event) = event else { break }; + self.handle_priority(event, Instant::now()); + self.drain_events(); + } + node = self.discovery_rx.recv() => { + self.runtime_stats + .set_crawl_discovery_queue_depth(self.discovery_rx.len()); + let Some(node) = node else { break }; + self.admit(node, Instant::now()); + self.drain_events(); + } + } + } + self.runtime_stats.set_crawl_priority_queue_depth(0); + self.runtime_stats.set_crawl_discovery_queue_depth(0); + } + + fn drain_events(&mut self) { + let mut events = 1; + let mut nodes = 0; + while events < self.config.event_batch_limit && nodes < self.config.node_batch_limit { + if events % 8 == 0 + && let Ok(node) = self.discovery_rx.try_recv() + { + self.runtime_stats + .set_crawl_discovery_queue_depth(self.discovery_rx.len()); + self.admit(node, Instant::now()); + events += 1; + nodes += 1; + continue; + } + if let Ok(event) = self.priority_rx.try_recv() { + self.runtime_stats + .set_crawl_priority_queue_depth(self.priority_rx.len()); + self.handle_priority(event, Instant::now()); + events += 1; + continue; + } + if let Ok(node) = self.discovery_rx.try_recv() { + self.runtime_stats + .set_crawl_discovery_queue_depth(self.discovery_rx.len()); + self.admit(node, Instant::now()); + events += 1; + nodes += 1; + continue; + } + break; + } + } + + fn handle_priority(&mut self, event: PriorityEvent, now: Instant) { + match event { + PriorityEvent::Response { + remote_addr, + tid, + response, + } => self.handle_response(remote_addr, tid, response, now), + PriorityEvent::SendFailed(key) => { + if let Some(pending) = self.pending.remove(&key) { + self.release_pending(pending); + self.metrics.send_failures += 1; + self.runtime_stats.send_failure(); + if pending.purpose == ProbePurpose::Bootstrap { + self.bootstrap_pool.mark_timeout(pending.node.addr, now); + } + } + } + PriorityEvent::BootstrapResolved(candidates) => { + let selected = self.bootstrap_pool.select( + candidates, + self.config.bootstrap_max_nodes_per_round, + now, + ); + self.bootstrap_queue.extend(selected); + } + } + } + + fn handle_response( + &mut self, + remote_addr: SocketAddr, + tid: TransactionId, + response: DhtResponse, + now: Instant, + ) { + let key = PendingKey { + addr: remote_addr, + tid, + }; + let Some(pending) = self.pending.remove(&key) else { + self.metrics.unmatched_responses += 1; + self.runtime_stats.unmatched_response(); + return; + }; + self.release_pending(pending); + self.metrics.responses += 1; + self.runtime_stats.response(); + if pending.purpose == ProbePurpose::Bootstrap { + self.bootstrap_pool.mark_success(remote_addr, now); + } + + let mut responsive_node = pending.node; + if let Some(id) = response.id.as_ref() + && let Ok(id) = <[u8; 20]>::try_from(id.as_slice()) + { + responsive_node.id = id; + } + self.responsive.record(responsive_node, now); + + let pool = &mut self.pool; + let metrics = &mut self.metrics; + let runtime_stats = &self.runtime_stats; + for_each_response_node(&response, self.netmode, |node| { + let outcome = pool.admit(node, now); + metrics.record_admission(outcome); + record_runtime_admission(runtime_stats, outcome); + }); + self.sync_node_count(); + } + + fn admit(&mut self, node: NodeTuple, now: Instant) { + let outcome = self.pool.admit(node, now); + self.metrics.record_admission(outcome); + record_runtime_admission(&self.runtime_stats, outcome); + self.sync_node_count(); + } + + fn on_schedule_tick(&mut self, now: Instant) { + self.expire_pending(now); + self.maybe_resolve_bootstrap(now); + + let rate = self + .config + .rate_for_metadata_pressure(self.metadata_pressure()); + self.runtime_stats.set_find_node_effective_rate(rate); + self.query_budget.set_per_second_rate(rate, now); + let budget = self.query_budget.try_take(self.config.burst as usize, now); + for _ in 0..budget { + if self.pending.len() >= self.config.max_in_flight { + self.query_budget.refund_one(); + break; + } + if !self.schedule_one(now) { + self.query_budget.refund_one(); + break; + } + } + self.sync_node_count(); + } + + fn schedule_one(&mut self, now: Instant) -> bool { + if self.pool.len() < self.config.low_watermark + && let Some(addr) = self.bootstrap_queue.front().copied() + { + let is_new = !self.pool.contains_recent(&addr, now); + if is_new && !self.destination_budget.try_take_one(now) { + return self.schedule_revisit(now); + } + let node = NodeTuple { + id: self.local_id, + addr, + }; + if self.try_dispatch(node, ProbePurpose::Bootstrap, now) { + self.bootstrap_queue.pop_front(); + self.pool.record_probe(addr, now); + self.bootstrap_pool.mark_attempt(addr, now); + self.metrics.queries_bootstrap += 1; + self.runtime_stats.query_bootstrap(); + return true; + } + if is_new { + self.destination_budget.refund_one(); + } + } + + let scan_limit = self.pool.len().min(MAX_POOL_SCAN_PER_SCHEDULE); + for _ in 0..scan_limit { + let Some(node) = self.pool.front() else { + break; + }; + let subnet = SubnetKey::from_addr(&node.addr); + if self.subnet_count(&subnet) >= self.config.max_in_flight_per_subnet { + self.pool.rotate_front_to_back(); + continue; + } + if !self.destination_budget.try_take_one(now) { + return self.schedule_revisit(now); + } + let node = self + .pool + .take_front_for_probe(now) + .expect("front node exists"); + if self.try_dispatch(node, ProbePurpose::New, now) { + self.metrics.queries_new += 1; + self.runtime_stats.query_new(); + return true; + } + self.pool.restore_front(node, now); + self.destination_budget.refund_one(); + break; + } + self.schedule_revisit(now) + } + + fn schedule_revisit(&mut self, now: Instant) -> bool { + let Some(node) = self.responsive.next_revisit(now) else { + return false; + }; + let was_recent = self.pool.contains_recent(&node.addr, now); + if !was_recent && !self.destination_budget.try_take_one(now) { + return false; + } + if self.try_dispatch(node, ProbePurpose::Revisit, now) { + self.pool.record_probe(node.addr, now); + self.metrics.queries_revisit += 1; + self.runtime_stats.query_revisit(); + return true; + } + if !was_recent { + self.destination_budget.refund_one(); + } + false + } + + fn try_dispatch(&mut self, node: NodeTuple, purpose: ProbePurpose, now: Instant) -> bool { + if self.pending.len() >= self.config.max_in_flight { + return false; + } + let subnet = SubnetKey::from_addr(&node.addr); + if self.subnet_count(&subnet) >= self.config.max_in_flight_per_subnet { + return false; + } + let tx = if node.addr.is_ipv4() { + self.egress_v4.clone() + } else { + self.egress_v6.clone() + }; + let Some(tx) = tx else { + return false; + }; + let Ok(permit) = tx.try_reserve() else { + return false; + }; + + let tid = self.next_tid.to_be_bytes(); + self.next_tid = self.next_tid.wrapping_add(1).max(1); + let key = PendingKey { + addr: node.addr, + tid, + }; + let deadline = now + self.config.request_timeout; + self.pending.insert( + key, + PendingRequest { + node, + purpose, + deadline, + subnet, + }, + ); + self.pending_expiry.push_back((deadline, key)); + *self.subnet_in_flight.entry(subnet).or_insert(0) += 1; + self.runtime_stats + .set_find_node_in_flight(self.pending.len()); + + let sender_id = if self.config.neighbor_sender_id { + let generated = neighbor_node_id(&node.id, &self.local_id); + generated + .as_slice() + .try_into() + .expect("neighbor id is always 20 bytes") + } else { + self.local_id + }; + permit.send(OutboundRequest { + key, + node, + target: self.choose_target(&node), + sender_id, + }); + true + } + + fn choose_target(&self, node: &NodeTuple) -> [u8; 20] { + let total = self + .config + .random_walk_percent + .saturating_add(self.config.sparse_bucket_percent) + .max(1); + if rand::thread_rng().gen_range(0..total) < self.config.sparse_bucket_percent { + target_for_bucket(&self.local_id, bucket_index(&node.id, &self.local_id)) + } else { + random_node_id() + } + } + + fn expire_pending(&mut self, now: Instant) { + while let Some((deadline, key)) = self.pending_expiry.front().copied() { + if deadline > now { + break; + } + self.pending_expiry.pop_front(); + let should_remove = self + .pending + .get(&key) + .is_some_and(|pending| pending.deadline == deadline); + if !should_remove { + continue; + } + let pending = self.pending.remove(&key).expect("pending entry exists"); + self.release_pending(pending); + self.metrics.timeouts += 1; + self.runtime_stats.timeout(); + if pending.purpose == ProbePurpose::Bootstrap { + self.bootstrap_pool.mark_timeout(pending.node.addr, now); + } + } + } + + fn release_pending(&mut self, pending: PendingRequest) { + if let Some(count) = self.subnet_in_flight.get_mut(&pending.subnet) { + *count = count.saturating_sub(1); + if *count == 0 { + self.subnet_in_flight.remove(&pending.subnet); + } + } + self.runtime_stats + .set_find_node_in_flight(self.pending.len()); + } + + fn subnet_count(&self, subnet: &SubnetKey) -> usize { + self.subnet_in_flight.get(subnet).copied().unwrap_or(0) + } + + fn maybe_resolve_bootstrap(&mut self, now: Instant) { + if !self.bootstrap_queue.is_empty() + || !self + .bootstrap_gate + .should_bootstrap(self.pool.len(), &self.config, now) + { + return; + } + let hosts = self.bootstrap_pool.hosts.clone(); + let netmode = self.netmode; + let priority_tx = self.priority_tx.clone(); + let runtime_stats = self.runtime_stats.clone(); + tokio::spawn(async move { + let resolved = resolve_bootstrap_nodes(&hosts, netmode).await; + let _ = priority_tx.try_send(PriorityEvent::BootstrapResolved(resolved)); + runtime_stats.set_crawl_priority_queue_depth( + priority_tx + .max_capacity() + .saturating_sub(priority_tx.capacity()), + ); + }); + } + + fn publish_snapshot(&self, now: Instant) { + let nodes = self + .responsive + .snapshot(self.config.routing_snapshot_size, now); + self.snapshot.store(Arc::new(RoutingSnapshot::from_nodes( + nodes, + self.config.routing_snapshot_size, + ))); + } + + fn metadata_pressure(&self) -> f64 { + if self.max_metadata_queue_size == 0 { + 1.0 + } else { + (self.metadata_queue_len.load(Ordering::Relaxed) as f64 + / self.max_metadata_queue_size as f64) + .min(1.0) + } + } + + fn sync_node_count(&self) { + self.node_count.store(self.pool.len(), Ordering::Relaxed); + self.runtime_stats.set_node_pool_size(self.pool.len()); + } + + fn flush_metrics(&mut self, now: Instant) { + #[cfg(feature = "metrics")] + { + let metadata_pressure = self.metadata_pressure(); + gauge!("dht_node_pool_size").set(self.pool.len() as f64); + gauge!("dht_node_pool_oldest_age_seconds").set(self.pool.oldest_age(now).as_secs_f64()); + gauge!("dht_find_node_in_flight").set(self.pending.len() as f64); + gauge!("dht_metadata_queue_pressure_ratio").set(metadata_pressure); + gauge!("dht_find_node_effective_rate_per_second") + .set(self.config.rate_for_metadata_pressure(metadata_pressure) as f64); + counter!("dht_node_pool_admissions_total").increment(self.metrics.admitted); + counter!("dht_node_pool_replacements_total").increment(self.metrics.replaced); + counter!("dht_node_pool_dropped_total", "reason" => "duplicate") + .increment(self.metrics.duplicate); + counter!("dht_node_pool_dropped_total", "reason" => "rate_limit") + .increment(self.metrics.admission_limited); + counter!("dht_node_pool_dropped_total", "reason" => "invalid") + .increment(self.metrics.invalid); + counter!("dht_crawl_queries_sent_total", "kind" => "new") + .increment(self.metrics.queries_new); + counter!("dht_crawl_queries_sent_total", "kind" => "revisit") + .increment(self.metrics.queries_revisit); + counter!("dht_crawl_queries_sent_total", "kind" => "bootstrap") + .increment(self.metrics.queries_bootstrap); + counter!("dht_find_node_responses_total").increment(self.metrics.responses); + counter!("dht_find_node_timeouts_total").increment(self.metrics.timeouts); + counter!("dht_find_node_send_failures_total").increment(self.metrics.send_failures); + counter!("dht_find_node_response_unmatched_total") + .increment(self.metrics.unmatched_responses); + } + #[cfg(not(feature = "metrics"))] + let _ = now; + self.metrics = ActorMetrics::default(); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::types::CrawlOptions; + + fn node(id: u8, addr: &str) -> NodeTuple { + NodeTuple { + id: [id; 20], + addr: addr.parse().unwrap(), + } + } + + fn test_actor( + config: ResolvedCrawlConfig, + ) -> (CrawlActor, mpsc::Receiver, DhtRuntimeStats) { + let (priority_tx, priority_rx) = mpsc::channel(16); + let (_discovery_tx, discovery_rx) = mpsc::channel(16); + let (egress_tx, egress_rx) = mpsc::channel(config.max_in_flight); + let runtime_stats = DhtRuntimeStats::with_limits(DhtRuntimeLimits { + metadata_queue: 100_000, + node_pool: config.pool_capacity, + node_pool_low_watermark: config.low_watermark, + find_node_in_flight: config.max_in_flight, + initial_find_node_rate: config.max_find_node_rate_per_sec, + hash_ingress_queue: 0, + crawl_priority_queue: config.priority_event_channel_capacity, + crawl_discovery_queue: config.discovery_event_channel_capacity, + }); + let actor = CrawlActor::new(CrawlActorInit { + config, + netmode: NetMode::Ipv4Only, + local_id: [7; 20], + priority_rx, + discovery_rx, + priority_tx, + egress_v4: Some(egress_tx), + egress_v6: None, + snapshot: Arc::new(ArcSwap::from_pointee(RoutingSnapshot::default())), + node_count: Arc::new(AtomicUsize::new(0)), + metadata_queue_len: Arc::new(AtomicUsize::new(0)), + max_metadata_queue_size: 100_000, + runtime_stats: runtime_stats.clone(), + shutdown: CancellationToken::new(), + }); + (actor, egress_rx, runtime_stats) + } + + #[test] + fn saturated_head_subnet_does_not_block_later_node() { + let config = ResolvedCrawlConfig::from_options(&CrawlOptions::default()); + let max_per_subnet = config.max_in_flight_per_subnet; + let (mut actor, mut egress_rx, runtime_stats) = test_actor(config); + let now = Instant::now() + Duration::from_secs(1); + let blocked = node(1, "8.8.8.8:6881"); + let eligible = node(2, "1.1.1.1:6881"); + + assert_eq!(actor.pool.admit(blocked, now), AdmissionOutcome::Admitted); + assert_eq!(actor.pool.admit(eligible, now), AdmissionOutcome::Admitted); + actor + .subnet_in_flight + .insert(SubnetKey::from_addr(&blocked.addr), max_per_subnet); + + assert!(actor.schedule_one(now)); + let request = egress_rx.try_recv().expect("eligible node was dispatched"); + + assert_eq!(request.node, eligible); + assert_eq!(actor.pool.front(), Some(blocked)); + assert_eq!(actor.pool.admit(blocked, now), AdmissionOutcome::Duplicate); + assert!(!actor.pool.contains_recent(&blocked.addr, now)); + assert!(actor.pool.contains_recent(&eligible.addr, now)); + let snapshot = runtime_stats.snapshot(); + assert_eq!(snapshot.queries_new, 1); + assert_eq!(snapshot.find_node_in_flight, 1); + assert_eq!(snapshot.find_node_in_flight_max, 512); + } + + #[test] + fn runtime_stats_count_dropped_crawl_events() { + let mut options = CrawlOptions::default(); + options.scheduler.priority_event_channel_capacity = 1; + options.scheduler.discovery_event_channel_capacity = 1; + let config = ResolvedCrawlConfig::from_options(&options); + let stats = DhtRuntimeStats::with_limits(DhtRuntimeLimits { + metadata_queue: 100, + node_pool: config.pool_capacity, + node_pool_low_watermark: config.low_watermark, + find_node_in_flight: config.max_in_flight, + initial_find_node_rate: config.max_find_node_rate_per_sec, + hash_ingress_queue: 0, + crawl_priority_queue: config.priority_event_channel_capacity, + crawl_discovery_queue: config.discovery_event_channel_capacity, + }); + let engine = CrawlEngine::new(config, stats.clone()); + + engine.route_discovered(node(1, "8.8.8.8:1")); + engine.route_discovered(node(2, "1.1.1.1:2")); + let response = || DhtResponse { + id: None, + nodes: None, + nodes6: None, + }; + engine.route_response("8.8.8.8:1".parse().unwrap(), [1; 8], response()); + engine.route_response("1.1.1.1:2".parse().unwrap(), [2; 8], response()); + + let snapshot = stats.snapshot(); + assert_eq!(snapshot.crawl_events_dropped_discovered, 1); + assert_eq!(snapshot.crawl_events_dropped_response, 1); + assert_eq!(snapshot.crawl_discovery_queue_depth, 1); + assert_eq!(snapshot.crawl_discovery_queue_capacity, 1); + assert_eq!(snapshot.crawl_priority_queue_depth, 1); + assert_eq!(snapshot.crawl_priority_queue_capacity, 1); + } +} diff --git a/src/error.rs b/src/error.rs index 7563444..26b5dde 100644 --- a/src/error.rs +++ b/src/error.rs @@ -1,21 +1,28 @@ use thiserror::Error; #[derive(Error, Debug)] +/// Error returned during DHT initialization or execution. pub enum DHTError { + /// Socket or other network I/O failed. #[error("网络错误: {0}")] Network(#[from] std::io::Error), + /// A shared lock was poisoned. #[error("锁中毒: {0}")] LockPoisoned(String), + /// Server initialization failed, for example because a socket could not bind. #[error("初始化错误: {0}")] Init(String), + /// An internal invariant or worker operation failed. #[error("内部错误: {0}")] Internal(String), + /// Another error represented by a human-readable message. #[error("{0}")] Other(String), } +/// Result type used by the crate's public APIs. pub type Result = std::result::Result; diff --git a/src/krpc.rs b/src/krpc.rs new file mode 100644 index 0000000..b51674c --- /dev/null +++ b/src/krpc.rs @@ -0,0 +1,191 @@ +use crate::addr::is_valid_node_addr; +use crate::node_id::TransactionId; +use crate::protocol::DhtResponse; +use crate::types::{NetMode, NodeTuple}; +use bytes::BytesMut; +#[cfg(feature = "metrics")] +use metrics::{counter, histogram}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; +use std::sync::Arc; +use tokio::net::UdpSocket; + +pub(crate) fn for_each_response_node( + response: &DhtResponse, + netmode: NetMode, + mut visit: impl FnMut(NodeTuple), +) -> usize { + let mut count = 0; + if netmode != NetMode::Ipv6Only + && let Some(nodes) = response.nodes.as_deref() + && nodes.len() % 26 == 0 + { + for chunk in nodes.chunks_exact(26) { + let id: [u8; 20] = chunk[..20].try_into().expect("compact v4 id is 20 bytes"); + let ip = Ipv4Addr::new(chunk[20], chunk[21], chunk[22], chunk[23]); + let port = u16::from_be_bytes([chunk[24], chunk[25]]); + let addr = SocketAddr::new(IpAddr::V4(ip), port); + if is_valid_node_addr(&addr) { + visit(NodeTuple { id, addr }); + count += 1; + } + } + } + if netmode != NetMode::Ipv4Only + && let Some(nodes) = response.nodes6.as_deref() + && nodes.len() % 38 == 0 + { + for chunk in nodes.chunks_exact(38) { + let id: [u8; 20] = chunk[..20].try_into().expect("compact v6 id is 20 bytes"); + let ip_bytes: [u8; 16] = chunk[20..36] + .try_into() + .expect("compact v6 address is 16 bytes"); + let port = u16::from_be_bytes([chunk[36], chunk[37]]); + let addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::from(ip_bytes)), port); + if is_valid_node_addr(&addr) { + visit(NodeTuple { id, addr }); + count += 1; + } + } + } + count +} + +pub(crate) fn encode_find_node_query( + buffer: &mut BytesMut, + tid: &TransactionId, + target: &[u8; 20], + sender_id: &[u8; 20], +) { + buffer.clear(); + buffer.reserve(112); + buffer.extend_from_slice(b"d1:ad2:id20:"); + buffer.extend_from_slice(sender_id); + buffer.extend_from_slice(b"6:target20:"); + buffer.extend_from_slice(target); + buffer.extend_from_slice(b"e1:q9:find_node1:t8:"); + buffer.extend_from_slice(tid); + buffer.extend_from_slice(b"1:y1:qe"); +} + +pub(crate) fn encode_response( + buffer: &mut BytesMut, + tid: &[u8], + node_id: &[u8; 20], + token: &[u8; 8], + nodes: &[NodeTuple], + ipv6: bool, +) { + buffer.clear(); + buffer.reserve(384); + buffer.extend_from_slice(b"d1:rd2:id20:"); + buffer.extend_from_slice(node_id); + + let compact_len = if ipv6 { + nodes.iter().filter(|node| node.addr.is_ipv6()).count() * 38 + } else { + nodes.iter().filter(|node| node.addr.is_ipv4()).count() * 26 + }; + if compact_len > 0 { + if ipv6 { + buffer.extend_from_slice(b"6:nodes6"); + } else { + buffer.extend_from_slice(b"5:nodes"); + } + push_usize(buffer, compact_len); + buffer.extend_from_slice(b":"); + for node in nodes { + match node.addr.ip() { + IpAddr::V4(ip) if !ipv6 => { + buffer.extend_from_slice(&node.id); + buffer.extend_from_slice(&ip.octets()); + buffer.extend_from_slice(&node.addr.port().to_be_bytes()); + } + IpAddr::V6(ip) if ipv6 => { + buffer.extend_from_slice(&node.id); + buffer.extend_from_slice(&ip.octets()); + buffer.extend_from_slice(&node.addr.port().to_be_bytes()); + } + _ => {} + } + } + } + + buffer.extend_from_slice(b"5:token8:"); + buffer.extend_from_slice(token); + buffer.extend_from_slice(b"e1:t"); + push_usize(buffer, tid.len()); + buffer.extend_from_slice(b":"); + buffer.extend_from_slice(tid); + buffer.extend_from_slice(b"1:y1:re"); +} + +fn push_usize(buffer: &mut BytesMut, mut value: usize) { + let mut digits = [0u8; 20]; + let mut cursor = digits.len(); + loop { + cursor -= 1; + digits[cursor] = b'0' + (value % 10) as u8; + value /= 10; + if value == 0 { + break; + } + } + buffer.extend_from_slice(&digits[cursor..]); +} + +pub(crate) async fn send_find_node_query( + addr: &SocketAddr, + tid: &TransactionId, + target: &[u8; 20], + sender_id: &[u8; 20], + socket: &Arc, + buffer: &mut BytesMut, +) -> bool { + encode_find_node_query(buffer, tid, target, sender_id); + match socket.send_to(buffer, addr).await { + Ok(len) => { + #[cfg(feature = "metrics")] + { + counter!("dht_udp_bytes_sent_total").increment(len as u64); + counter!("dht_udp_packets_sent_total", "type" => "query").increment(1); + histogram!("dht_udp_query_size_bytes").record(len as f64); + } + #[cfg(not(feature = "metrics"))] + let _ = len; + true + } + Err(_) => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::protocol::DhtMessage; + + #[test] + fn manual_find_node_encoding_round_trips() { + let mut buffer = BytesMut::new(); + let tid = [1; 8]; + let target = [2; 20]; + let sender = [3; 20]; + encode_find_node_query(&mut buffer, &tid, &target, &sender); + let message: DhtMessage = serde_bencode::from_bytes(&buffer).unwrap(); + assert_eq!(message.t.as_ref(), &tid); + assert_eq!(message.q.as_deref(), Some("find_node")); + assert_eq!(message.a.unwrap().target.unwrap().as_ref(), &target); + } + + #[test] + fn manual_response_encoding_round_trips() { + let mut buffer = BytesMut::new(); + let nodes = [NodeTuple { + id: [4; 20], + addr: "8.8.8.8:6881".parse().unwrap(), + }]; + encode_response(&mut buffer, &[1, 2], &[2; 20], &[3; 8], &nodes, false); + let message: DhtMessage = serde_bencode::from_bytes(&buffer).unwrap(); + let response = message.r.unwrap(); + assert_eq!(response.nodes.unwrap().len(), 26); + } +} diff --git a/src/lib.rs b/src/lib.rs index 64f2e62..d78c863 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,24 +1,61 @@ +//! High-throughput BitTorrent DHT crawler with bounded crawl and Metadata pipelines. +//! +//! [`DHTServer`] is the primary entry point. Configure it with [`DHTOptions`], register +//! callbacks, then await [`DHTServer::start`] until another task calls [`DHTServer::shutdown`]. +//! Runtime counters are available through [`DHTServer::runtime_stats`] without enabling any +//! exporter. See the repository README for scheduling, backpressure and migration details. + +mod addr; +mod bootstrap; +mod budget; +mod crawl_config; +mod crawl_engine; mod error; +mod krpc; +/// BEP-9 Metadata fetch support. pub mod metadata; +mod node_id; +mod node_pool; +/// Serializable BEP-5 KRPC wire types. pub mod protocol; +mod routing_snapshot; +mod runtime_stats; +/// Bounded, deduplicating Metadata scheduler. pub mod scheduler; mod server; -mod sharded; +/// Public configuration, callback payload and network types. pub mod types; +mod udp_buffer; +mod udp_ingress; pub use error::{DHTError, Result}; -pub use scheduler::MetadataScheduler; +pub use runtime_stats::{ + DhtObservabilitySnapshot, DhtRuntimeSnapshot, DhtRuntimeStats, FixedHistogramSnapshot, +}; +pub use scheduler::{MetadataScheduler, MetadataSchedulerCallbacks, MetadataSchedulerLimits}; pub use server::{DHTServer, HashDiscovered}; -pub use sharded::{NodeTuple, ShardedNodeQueue}; -pub use types::{DHTOptions, FileInfo, NetMode, TorrentInfo}; +pub use types::{ + BootstrapOptions, CrawlOptions, DHTOptions, FileInfo, MetadataFetchCompletion, + MetadataFetchCompletionStatus, MetadataOptions, NetMode, NodeTuple, PoolOptions, + RateLimitOptions, SchedulerOptions, TargetOptions, TorrentInfo, +}; +/// Common server, configuration and callback payload imports. pub mod prelude { pub use crate::error::{DHTError, Result}; - pub use crate::scheduler::MetadataScheduler; + pub use crate::runtime_stats::{DhtRuntimeSnapshot, DhtRuntimeStats}; + pub use crate::scheduler::{ + MetadataScheduler, MetadataSchedulerCallbacks, MetadataSchedulerLimits, + }; pub use crate::server::DHTServer; - pub use crate::types::{DHTOptions, FileInfo, NetMode, TorrentInfo}; + pub use crate::types::{ + BootstrapOptions, CrawlOptions, DHTOptions, FileInfo, MetadataFetchCompletion, + MetadataFetchCompletionStatus, MetadataOptions, NetMode, NodeTuple, PoolOptions, + RateLimitOptions, SchedulerOptions, TargetOptions, TorrentInfo, + }; } #[cfg(feature = "jni")] #[path = "../jni/mod.rs"] +/// JNI entry points used by the bundled Java wrapper. pub mod jni_bindings; diff --git a/src/metadata.rs b/src/metadata.rs index d3af2b6..54ee385 100644 --- a/src/metadata.rs +++ b/src/metadata.rs @@ -1,236 +1,611 @@ +use crate::runtime_stats::DhtRuntimeStats; use crate::types::FileInfo; +use ahash::AHashMap; use bytes::Bytes; #[cfg(feature = "metrics")] -use metrics::{counter, histogram}; +use metrics::{counter, gauge, histogram}; use rbit::peer::ExtensionMessage; use rbit::{ ExtensionHandshake, Message, MetadataMessage, MetadataMessageType, PeerConnection, PeerId, metadata_piece_count, }; use sha1::{Digest, Sha1}; -use std::collections::BTreeMap; +use std::collections::{BTreeMap, VecDeque}; use std::net::SocketAddr; -use std::time::Duration; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; use tokio::time::timeout; -#[derive(Clone)] -pub struct RbitFetcher { - timeout: Duration, +pub(crate) type FetchedMetadata = (String, u64, Vec, u64); + +pub(crate) enum MetadataFetchOutcome { + Fetched(FetchedMetadata), + Failed, + SkippedCached, } -impl RbitFetcher { - pub fn new(timeout_secs: u64) -> Self { +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum MetadataFetchFailure { + Connect, + NoExtension, + Send, + SizeLimit, + Sha1, + Parse, + Other, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PeerFailureReason { + Timeout, + ConnectFailed, +} + +#[cfg(feature = "metrics")] +impl PeerFailureReason { + fn as_str(self) -> &'static str { + match self { + Self::Timeout => "timeout", + Self::ConnectFailed => "connect_failed", + } + } +} + +#[derive(Debug, Clone, Copy)] +struct PeerFailureEntry { + expires_at: Instant, + reason: PeerFailureReason, +} + +#[derive(Default)] +struct PeerFailureCacheInner { + entries: AHashMap, + expiry: VecDeque<(Instant, SocketAddr)>, +} + +struct PeerFailureCache { + inner: Mutex, + capacity: usize, + ttl: Duration, +} + +impl PeerFailureCache { + fn new(capacity: usize, ttl: Duration) -> Self { Self { - timeout: Duration::from_secs(if timeout_secs == 0 { 15 } else { timeout_secs }), + inner: Mutex::new(PeerFailureCacheInner { + entries: AHashMap::with_capacity(capacity.min(16_384)), + expiry: VecDeque::with_capacity(capacity.min(16_384)), + }), + capacity, + ttl, } } - pub async fn fetch( + fn get(&self, addr: SocketAddr, now: Instant) -> (Option, usize) { + if self.capacity == 0 || self.ttl.is_zero() { + return (None, 0); + } + let mut inner = self + .inner + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + Self::expire(&mut inner, now); + ( + inner.entries.get(&addr).map(|entry| entry.reason), + inner.entries.len(), + ) + } + + fn insert(&self, addr: SocketAddr, reason: PeerFailureReason, now: Instant) -> usize { + if self.capacity == 0 || self.ttl.is_zero() { + return 0; + } + let mut inner = self + .inner + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + Self::expire(&mut inner, now); + while inner.entries.len() >= self.capacity && !inner.entries.contains_key(&addr) { + let Some((expires_at, oldest_addr)) = inner.expiry.pop_front() else { + break; + }; + if inner + .entries + .get(&oldest_addr) + .is_some_and(|entry| entry.expires_at == expires_at) + { + inner.entries.remove(&oldest_addr); + } + } + + let expires_at = now + self.ttl; + inner + .entries + .insert(addr, PeerFailureEntry { expires_at, reason }); + inner.expiry.push_back((expires_at, addr)); + inner.entries.len() + } + + fn remove(&self, addr: &SocketAddr, now: Instant) -> usize { + if self.capacity == 0 || self.ttl.is_zero() { + return 0; + } + let mut inner = self + .inner + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + Self::expire(&mut inner, now); + inner.entries.remove(addr); + inner.entries.len() + } + + fn expire(inner: &mut PeerFailureCacheInner, now: Instant) { + while let Some((expires_at, addr)) = inner.expiry.front().copied() { + if expires_at > now { + break; + } + inner.expiry.pop_front(); + if inner + .entries + .get(&addr) + .is_some_and(|entry| entry.expires_at == expires_at) + { + inner.entries.remove(&addr); + } + } + } +} + +#[derive(Clone)] +/// BEP-9 Metadata fetcher with an end-to-end timeout and shared Peer failure cache. +pub struct RbitFetcher { + total_timeout: Duration, + runtime_stats: DhtRuntimeStats, + peer_failure_cache: Arc, +} + +impl RbitFetcher { + /// Creates a standalone fetcher with the default failure-cache capacity and TTL. + /// + /// [`DHTServer`](crate::DHTServer) normally constructs this component from + /// [`MetadataOptions`](crate::MetadataOptions). + pub fn new(timeout_secs: u64) -> Self { + Self::new_with_runtime_stats(timeout_secs, 200_000, 60, DhtRuntimeStats::default()) + } + + pub(crate) fn new_with_runtime_stats( + timeout_secs: u64, + peer_failure_cache_capacity: usize, + peer_failure_ttl_secs: u64, + runtime_stats: DhtRuntimeStats, + ) -> Self { + Self { + total_timeout: Duration::from_secs(if timeout_secs == 0 { 15 } else { timeout_secs }), + runtime_stats, + peer_failure_cache: Arc::new(PeerFailureCache::new( + peer_failure_cache_capacity, + Duration::from_secs(peer_failure_ttl_secs), + )), + } + } + + /// Fetch metadata from one peer under a single end-to-end deadline. + /// + /// The deadline covers TCP connect, both BitTorrent handshakes, all metadata + /// piece I/O, hash validation and bencode parsing. Inner library timeouts can + /// therefore never stack on top of the configured metadata timeout. + pub(crate) async fn fetch( &self, info_hash: &[u8; 20], peer_addr: SocketAddr, - ) -> Option<(String, u64, Vec, u64)> { + ) -> MetadataFetchOutcome { + let (cached_reason, cache_entries) = self.peer_failure_cache.get(peer_addr, Instant::now()); + self.set_peer_failure_cache_entries(cache_entries); + if let Some(reason) = cached_reason { + self.runtime_stats.metadata_peer_failure_cache_hit(); + match reason { + PeerFailureReason::Timeout => self.runtime_stats.peer_cache_hit_timeout(), + PeerFailureReason::ConnectFailed => self.runtime_stats.peer_cache_hit_connect(), + } + #[cfg(feature = "metrics")] + counter!("dht_metadata_peer_failure_cache_hits_total", "reason" => reason.as_str()) + .increment(1); + #[cfg(not(feature = "metrics"))] + let _ = reason; + return MetadataFetchOutcome::SkippedCached; + } + + self.runtime_stats.metadata_peer_attempt(); #[cfg(feature = "metrics")] - counter!("dht_metadata_fetch_attempts_total").increment(1); - - let peer_id = PeerId::generate(); - - let mut conn = match timeout( - Duration::from_secs(3), - PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes()), - ) - .await { - Ok(Ok(c)) => { + counter!("dht_metadata_fetch_attempts_total").increment(1); + counter!("dht_metadata_peer_attempts_total").increment(1); + } + + let started = Instant::now(); + let result = timeout( + self.total_timeout, + self.fetch_with_peer(info_hash, peer_addr), + ) + .await; + + #[cfg(feature = "metrics")] + histogram!("dht_metadata_fetch_duration_seconds").record(started.elapsed().as_secs_f64()); + self.runtime_stats.observe_metadata_fetch_duration( + started.elapsed().as_millis().min(u128::from(u64::MAX)) as u64, + ); + + match result { + Ok(Ok(metadata)) => { + let cache_entries = self.peer_failure_cache.remove(&peer_addr, Instant::now()); + self.set_peer_failure_cache_entries(cache_entries); + self.runtime_stats.metadata_peer_succeeded(); + #[cfg(feature = "metrics")] + { + counter!("dht_metadata_fetch_success_total").increment(1); + counter!("dht_metadata_fetch_result_total", "result" => "success").increment(1); + } + MetadataFetchOutcome::Fetched(metadata) + } + Ok(Err(reason)) => { + self.runtime_stats.metadata_peer_failed(); + match reason { + MetadataFetchFailure::Connect => self.runtime_stats.metadata_failure_connect(), + MetadataFetchFailure::NoExtension => { + self.runtime_stats.metadata_failure_no_extension() + } + MetadataFetchFailure::Send => self.runtime_stats.metadata_failure_send(), + MetadataFetchFailure::SizeLimit => { + self.runtime_stats.metadata_failure_size_limit() + } + MetadataFetchFailure::Sha1 => self.runtime_stats.metadata_failure_sha1(), + MetadataFetchFailure::Parse => self.runtime_stats.metadata_failure_parse(), + MetadataFetchFailure::Other => self.runtime_stats.metadata_failure_other(), + } + #[cfg(feature = "metrics")] + counter!("dht_metadata_fetch_result_total", "result" => "failed").increment(1); + MetadataFetchOutcome::Failed + } + Err(_) => { + self.record_peer_failure(peer_addr, PeerFailureReason::Timeout); + self.runtime_stats.metadata_peer_failed(); + self.runtime_stats.metadata_peer_timeout(); + self.runtime_stats.metadata_failure_timeout(); + #[cfg(feature = "metrics")] + { + counter!("dht_metadata_fetch_fail_total", "reason" => "timeout").increment(1); + counter!("dht_metadata_fetch_result_total", "result" => "timeout").increment(1); + } + MetadataFetchOutcome::Failed + } + } + } + + fn record_peer_failure(&self, peer_addr: SocketAddr, reason: PeerFailureReason) { + let cache_entries = self + .peer_failure_cache + .insert(peer_addr, reason, Instant::now()); + self.set_peer_failure_cache_entries(cache_entries); + #[cfg(feature = "metrics")] + counter!("dht_metadata_peer_failure_cache_inserts_total", "reason" => reason.as_str()) + .increment(1); + } + + fn set_peer_failure_cache_entries(&self, count: usize) { + self.runtime_stats + .set_metadata_peer_failure_cache_entries(count); + #[cfg(feature = "metrics")] + gauge!("dht_metadata_peer_failure_cache_entries").set(count as f64); + } + + async fn fetch_with_peer( + &self, + info_hash: &[u8; 20], + peer_addr: SocketAddr, + ) -> Result { + let peer_id = PeerId::generate(); + let mut conn = match PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes()) + .await + { + Ok(conn) => { #[cfg(feature = "metrics")] counter!("dht_metadata_connection_result_total", "result" => "success") .increment(1); - c - } - Ok(Err(_)) => { - #[cfg(feature = "metrics")] - counter!("dht_metadata_connection_result_total", "result" => "failed").increment(1); - return None; + conn } Err(_) => { + self.record_peer_failure(peer_addr, PeerFailureReason::ConnectFailed); + self.runtime_stats.metadata_connect_failed(); #[cfg(feature = "metrics")] - counter!("dht_metadata_connection_result_total", "result" => "timeout") - .increment(1); - return None; + counter!("dht_metadata_connection_result_total", "result" => "failed").increment(1); + return Err(MetadataFetchFailure::Connect); } }; if !conn.supports_extension { + self.runtime_stats.metadata_no_extension(); #[cfg(feature = "metrics")] counter!("dht_metadata_handshake_result_total", "result" => "no_extension_support") .increment(1); - return None; + return Err(MetadataFetchFailure::NoExtension); } let my_ut_metadata_id = 1; let handshake = ExtensionHandshake::with_extensions(&[("ut_metadata", my_ut_metadata_id)]); - - if let Ok(handshake_bytes) = handshake.encode() { - let _ = conn - .send(Message::Extended { - id: 0, - payload: handshake_bytes, - }) - .await; - } else { - return None; + let handshake_bytes = handshake.encode().map_err(|_| MetadataFetchFailure::Send)?; + if conn + .send(Message::Extended { + id: 0, + payload: handshake_bytes, + }) + .await + .is_err() + { + #[cfg(feature = "metrics")] + counter!("dht_metadata_fetch_fail_total", "reason" => "send_error").increment(1); + return Err(MetadataFetchFailure::Send); } let mut metadata_size = 0; let mut remote_ut_metadata_id = 0; let mut pieces: BTreeMap = BTreeMap::new(); + let mut total_received = 0usize; let mut request_sent = false; - let result = timeout(self.timeout, async { - loop { - let msg = conn.receive().await.ok()?; - if let Message::Extended { id, payload } = msg { - if id == 0 { - if let Ok(ExtensionMessage::Handshake(remote_hs)) = ExtensionMessage::decode(id, &payload) { - if let Some(size) = remote_hs.metadata_size { - metadata_size = size as u32; - } - if let Some(ext_id) = remote_hs.get_extension_id("ut_metadata") { - remote_ut_metadata_id = ext_id; - } - } - if metadata_size > 0 && remote_ut_metadata_id > 0 && !request_sent { - if metadata_size > 10 * 1024 * 1024 { - #[cfg(feature = "metrics")] - counter!("dht_metadata_fetch_fail_total", "reason" => "size_limit").increment(1); - return None; - } + let info_bytes = loop { + let msg = conn + .receive() + .await + .map_err(|_| MetadataFetchFailure::Other)?; + let Message::Extended { id, payload } = msg else { + continue; + }; - let count = metadata_piece_count(metadata_size as usize); - for i in 0..count { - let req = MetadataMessage::request(i as u32); - if let Ok(encoded) = req.encode() { - let _ = conn.send(Message::Extended { id: remote_ut_metadata_id, payload: encoded }).await; - } - } - request_sent = true; - } - } else if id == my_ut_metadata_id { - if let Ok(meta_msg) = MetadataMessage::decode(&payload) - && meta_msg.msg_type == MetadataMessageType::Data - && let Some(data) = meta_msg.data { - #[cfg(feature = "metrics")] - counter!("dht_metadata_bytes_downloaded_total").increment(data.len() as u64); - pieces.insert(meta_msg.piece, data); - } - if metadata_size > 0 { - let total_received: usize = pieces.values().map(|p| p.len()).sum(); - if total_received >= metadata_size as usize { - let mut full_data = Vec::with_capacity(metadata_size as usize); - let count = metadata_piece_count(metadata_size as usize); - let mut success = true; - for i in 0..count { - if let Some(p) = pieces.get(&(i as u32)) { - full_data.extend_from_slice(p); - } else { - success = false; break; - } - } - if success { - let info_hash_copy = *info_hash; - let validated = tokio::task::spawn_blocking(move || { - let mut hasher = Sha1::new(); - hasher.update(&full_data); - let digest: [u8; 20] = hasher.finalize().into(); - if digest == info_hash_copy { - Some(full_data) - } else { - None - } - }).await.unwrap_or(None); - - if validated.is_some() { - #[cfg(feature = "metrics")] - counter!("dht_metadata_handshake_result_total", "result" => "success").increment(1); - return validated; - } - #[cfg(feature = "metrics")] - counter!("dht_metadata_fetch_fail_total", "reason" => "sha1_mismatch").increment(1); - return None; - } - } - } - } - } - } - }).await; - - match result { - Ok(Some(info_bytes)) => { - if let Ok(value) = rbit::decode(&info_bytes) - && let Some(dict) = value.as_dict() + if id == 0 { + if let Ok(ExtensionMessage::Handshake(remote_hs)) = + ExtensionMessage::decode(id, &payload) { - let name = dict - .get(&b"name"[..]) - .and_then(|v| v.as_str()) - .unwrap_or("Unknown") - .to_string(); - let piece_length = dict - .get(&b"piece length"[..]) - .and_then(|v| v.as_integer()) - .unwrap_or(0) as u64; - let mut total_size = 0; - let mut file_list = Vec::new(); - if let Some(files) = dict.get(&b"files"[..]).and_then(|v| v.as_list()) { - for file in files { - if let Some(f_dict) = file.as_dict() - && let Some(len) = - f_dict.get(&b"length"[..]).and_then(|v| v.as_integer()) - { - let len = len as u64; - total_size += len; - let mut path_parts = Vec::new(); - if let Some(path_list) = - f_dict.get(&b"path"[..]).and_then(|v| v.as_list()) - { - for p in path_list { - if let Some(p_str) = p.as_str() { - path_parts.push(p_str); - } - } - } - file_list.push(FileInfo { - path: path_parts.join("/"), - size: len, - }); - } - } - } else if let Some(len) = dict.get(&b"length"[..]).and_then(|v| v.as_integer()) - { - total_size = len as u64; - file_list.push(FileInfo { - path: name.clone(), - size: total_size, - }); + if let Some(size) = remote_hs.metadata_size { + metadata_size = size as u32; } - if total_size > 0 { - #[cfg(feature = "metrics")] - { - counter!("dht_metadata_fetch_success_total").increment(1); - histogram!("dht_metadata_size_bytes").record(total_size as f64); - } - return Some((name, total_size, file_list, piece_length)); + if let Some(ext_id) = remote_hs.get_extension_id("ut_metadata") { + remote_ut_metadata_id = ext_id; } } + + if metadata_size > 0 && remote_ut_metadata_id > 0 && !request_sent { + if metadata_size > 10 * 1024 * 1024 { + #[cfg(feature = "metrics")] + counter!("dht_metadata_fetch_fail_total", "reason" => "size_limit") + .increment(1); + return Err(MetadataFetchFailure::SizeLimit); + } + + let count = metadata_piece_count(metadata_size as usize); + for piece in 0..count { + let encoded = MetadataMessage::request(piece as u32) + .encode() + .map_err(|_| MetadataFetchFailure::Send)?; + if conn + .send(Message::Extended { + id: remote_ut_metadata_id, + payload: encoded, + }) + .await + .is_err() + { + #[cfg(feature = "metrics")] + counter!("dht_metadata_fetch_fail_total", "reason" => "send_error") + .increment(1); + return Err(MetadataFetchFailure::Send); + } + } + request_sent = true; + } + continue; + } + + if id != my_ut_metadata_id { + continue; + } + let Ok(meta_msg) = MetadataMessage::decode(&payload) else { + continue; + }; + if meta_msg.msg_type != MetadataMessageType::Data { + continue; + } + let Some(data) = meta_msg.data else { + continue; + }; + + #[cfg(feature = "metrics")] + counter!("dht_metadata_bytes_downloaded_total").increment(data.len() as u64); + self.runtime_stats.metadata_bytes_downloaded(data.len()); + + let data_len = data.len(); + if let Some(previous) = pieces.insert(meta_msg.piece, data) { + total_received = total_received.saturating_sub(previous.len()); + } + total_received = total_received.saturating_add(data_len); + + if metadata_size == 0 || total_received < metadata_size as usize { + continue; + } + + let count = metadata_piece_count(metadata_size as usize); + let mut full_data = Vec::with_capacity(metadata_size as usize); + for piece in 0..count { + let data = pieces + .get(&(piece as u32)) + .ok_or(MetadataFetchFailure::Other)?; + full_data.extend_from_slice(data); + } + + let info_hash_copy = *info_hash; + let validated = tokio::task::spawn_blocking(move || { + let mut hasher = Sha1::new(); + hasher.update(&full_data); + let digest: [u8; 20] = hasher.finalize().into(); + (digest == info_hash_copy).then_some(full_data) + }) + .await + .ok() + .flatten(); + + match validated { + Some(data) => { + #[cfg(feature = "metrics")] + counter!("dht_metadata_handshake_result_total", "result" => "success") + .increment(1); + break data; + } + None => { + #[cfg(feature = "metrics")] + counter!("dht_metadata_fetch_fail_total", "reason" => "sha1_mismatch") + .increment(1); + return Err(MetadataFetchFailure::Sha1); + } + } + }; + + self.runtime_stats.observe_metadata_size(info_bytes.len()); + match parse_metadata(&info_bytes) { + Some(metadata) => { + #[cfg(feature = "metrics")] + histogram!("dht_metadata_size_bytes").record(info_bytes.len() as f64); + Ok(metadata) + } + None => { #[cfg(feature = "metrics")] counter!("dht_metadata_fetch_fail_total", "reason" => "parse_error").increment(1); - None - } - _ => { - #[cfg(feature = "metrics")] - counter!("dht_metadata_fetch_fail_total", "reason" => "timeout").increment(1); - None + Err(MetadataFetchFailure::Parse) } } } } + +fn parse_metadata(info_bytes: &[u8]) -> Option { + let value = rbit::decode(info_bytes).ok()?; + let dict = value.as_dict()?; + let name = dict + .get(&b"name"[..]) + .and_then(|value| value.as_str()) + .unwrap_or("Unknown") + .to_string(); + let piece_length = dict + .get(&b"piece length"[..]) + .and_then(|value| value.as_integer()) + .unwrap_or(0) as u64; + + let mut total_size = 0; + let mut file_list = Vec::new(); + if let Some(files) = dict.get(&b"files"[..]).and_then(|value| value.as_list()) { + for file in files { + let Some(file_dict) = file.as_dict() else { + continue; + }; + let Some(length) = file_dict + .get(&b"length"[..]) + .and_then(|value| value.as_integer()) + else { + continue; + }; + let length = length as u64; + total_size += length; + let path = file_dict + .get(&b"path"[..]) + .and_then(|value| value.as_list()) + .map(|parts| { + parts + .iter() + .filter_map(|part| part.as_str()) + .collect::>() + .join("/") + }) + .unwrap_or_default(); + file_list.push(FileInfo { path, size: length }); + } + } else if let Some(length) = dict + .get(&b"length"[..]) + .and_then(|value| value.as_integer()) + { + total_size = length as u64; + file_list.push(FileInfo { + path: name.clone(), + size: total_size, + }); + } + + (total_size > 0).then_some((name, total_size, file_list, piece_length)) +} + +#[cfg(test)] +mod tests { + use super::*; + use tokio::net::TcpListener; + + #[test] + fn peer_failure_cache_is_socket_specific_and_expires() { + let start = Instant::now(); + let cache = PeerFailureCache::new(10, Duration::from_secs(60)); + let first: SocketAddr = "127.0.0.1:1000".parse().unwrap(); + let same_ip_other_port: SocketAddr = "127.0.0.1:1001".parse().unwrap(); + + assert_eq!(cache.insert(first, PeerFailureReason::Timeout, start), 1); + assert_eq!(cache.get(first, start).0, Some(PeerFailureReason::Timeout)); + assert_eq!(cache.get(same_ip_other_port, start).0, None); + assert_eq!(cache.get(first, start + Duration::from_secs(61)), (None, 0)); + } + + #[test] + fn peer_failure_cache_evicts_oldest_at_capacity() { + let start = Instant::now(); + let cache = PeerFailureCache::new(1, Duration::from_secs(60)); + let first: SocketAddr = "127.0.0.1:1000".parse().unwrap(); + let second: SocketAddr = "127.0.0.1:1001".parse().unwrap(); + + cache.insert(first, PeerFailureReason::Timeout, start); + cache.insert(second, PeerFailureReason::ConnectFailed, start); + + assert_eq!(cache.get(first, start).0, None); + assert_eq!( + cache.get(second, start).0, + Some(PeerFailureReason::ConnectFailed) + ); + } + + #[tokio::test] + async fn total_timeout_covers_peer_handshake() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let accept_task = tokio::spawn(async move { + let (_stream, _) = listener.accept().await.unwrap(); + std::future::pending::<()>().await; + }); + + let stats = DhtRuntimeStats::default(); + let fetcher = RbitFetcher::new_with_runtime_stats(1, 10, 60, stats.clone()); + let started = Instant::now(); + assert!(matches!( + fetcher.fetch(&[7; 20], addr).await, + MetadataFetchOutcome::Failed + )); + assert!(started.elapsed() < Duration::from_secs(2)); + + let cached_started = Instant::now(); + assert!(matches!( + fetcher.fetch(&[8; 20], addr).await, + MetadataFetchOutcome::SkippedCached + )); + assert!(cached_started.elapsed() < Duration::from_millis(100)); + + let snapshot = stats.snapshot(); + assert_eq!(snapshot.metadata_peer_attempts, 1); + assert_eq!(snapshot.metadata_peer_failed, 1); + assert_eq!(snapshot.metadata_peer_timeouts, 1); + assert_eq!(snapshot.metadata_peer_failure_cache_hits, 1); + assert_eq!(snapshot.metadata_peer_failure_cache_entries, 1); + + accept_task.abort(); + } +} diff --git a/src/node_id.rs b/src/node_id.rs new file mode 100644 index 0000000..1a2582b --- /dev/null +++ b/src/node_id.rs @@ -0,0 +1,82 @@ +use rand::Rng; + +pub(crate) type TransactionId = [u8; 8]; + +pub(crate) fn transaction_id_from_bytes(bytes: &[u8]) -> Option { + if bytes.len() != 8 { + return None; + } + + let mut tid = [0u8; 8]; + tid.copy_from_slice(bytes); + Some(tid) +} + +pub(crate) fn random_node_id() -> [u8; 20] { + let mut id = [0u8; 20]; + rand::thread_rng().fill(&mut id); + id +} + +pub(crate) fn neighbor_node_id(remote_id: &[u8], local_id: &[u8]) -> Vec { + let mut id = Vec::with_capacity(20); + let prefix_len = remote_id.len().min(6); + id.extend_from_slice(&remote_id[..prefix_len]); + + if local_id.len() > prefix_len { + id.extend_from_slice(&local_id[prefix_len..]); + } + while id.len() < 20 { + id.push(rand::random()); + } + id.truncate(20); + id +} + +pub(crate) fn bucket_index(id: &[u8], local_id: &[u8; 20]) -> usize { + for bit in 0..160 { + let byte = bit / 8; + if byte >= id.len() { + break; + } + let mask = 1 << (7 - (bit % 8)); + if (id[byte] ^ local_id[byte]) & mask != 0 { + return bit; + } + } + 159 +} + +pub(crate) fn target_for_bucket(local_id: &[u8; 20], bucket: usize) -> [u8; 20] { + let mut id = *local_id; + let bucket = bucket.min(159); + let byte = bucket / 8; + let bit = 7 - (bucket % 8); + id[byte] ^= 1 << bit; + for item in id.iter_mut().skip(byte + 1) { + *item = rand::random(); + } + id +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn transaction_ids_are_eight_bytes() { + let first = 1u64.to_be_bytes(); + assert_eq!(transaction_id_from_bytes(&first), Some(first)); + assert!(transaction_id_from_bytes(&[1, 2]).is_none()); + } + + #[test] + fn neighbor_id_keeps_remote_prefix_and_local_suffix() { + let remote = [1u8; 20]; + let local = [2u8; 20]; + let id = neighbor_node_id(&remote, &local); + + assert_eq!(&id[..6], &[1u8; 6]); + assert_eq!(&id[6..], &[2u8; 14]); + } +} diff --git a/src/node_pool.rs b/src/node_pool.rs new file mode 100644 index 0000000..3a57abc --- /dev/null +++ b/src/node_pool.rs @@ -0,0 +1,390 @@ +use crate::addr::is_valid_node_addr; +use crate::budget::RateBucket; +use crate::types::NodeTuple; +use ahash::{AHashMap, AHashSet}; +use std::collections::VecDeque; +use std::net::{IpAddr, SocketAddr}; +use std::time::{Duration, Instant}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AdmissionOutcome { + Admitted, + Replaced, + Duplicate, + RateLimited, + Invalid, +} + +#[derive(Debug, Clone, Copy)] +struct QueuedNode { + node: NodeTuple, + #[cfg_attr(not(feature = "metrics"), allow(dead_code))] + queued_at: Instant, +} + +pub(crate) struct NodePool { + queue: VecDeque, + queued: AHashSet, + recent: AHashMap, + recent_expiry: VecDeque<(Instant, SocketAddr)>, + replacement_budget: RateBucket, + recent_ttl: Duration, + capacity: usize, + warmed: bool, +} + +impl NodePool { + pub(crate) fn new( + capacity: usize, + replacements_per_minute: u32, + recent_ttl: Duration, + now: Instant, + ) -> Self { + let capacity = capacity.max(1); + let replacement_burst = replacements_per_minute.div_ceil(60).max(1); + Self { + queue: VecDeque::with_capacity(capacity), + queued: AHashSet::with_capacity(capacity), + recent: AHashMap::with_capacity(capacity), + recent_expiry: VecDeque::with_capacity(capacity), + replacement_budget: RateBucket::per_minute( + replacements_per_minute, + replacement_burst, + true, + now, + ), + recent_ttl, + capacity, + warmed: false, + } + } + + pub(crate) fn admit(&mut self, node: NodeTuple, now: Instant) -> AdmissionOutcome { + if !is_valid_node_addr(&node.addr) { + return AdmissionOutcome::Invalid; + } + self.expire_recent(now); + if self.queued.contains(&node.addr) || self.recent.contains_key(&node.addr) { + return AdmissionOutcome::Duplicate; + } + if self.warmed && !self.replacement_budget.try_take_one(now) { + return AdmissionOutcome::RateLimited; + } + + let replaced = if self.queue.len() >= self.capacity { + self.pop_front_internal().is_some() + } else { + false + }; + self.queued.insert(node.addr); + self.queue.push_back(QueuedNode { + node, + queued_at: now, + }); + if self.queue.len() >= self.capacity { + self.warmed = true; + } + if replaced { + AdmissionOutcome::Replaced + } else { + AdmissionOutcome::Admitted + } + } + + pub(crate) fn front(&self) -> Option { + self.queue.front().map(|entry| entry.node) + } + + /// Move the FIFO head behind the remaining queued nodes without marking + /// it as probed. The address stays in `queued` and is not added to `recent`. + pub(crate) fn rotate_front_to_back(&mut self) -> bool { + if self.queue.len() <= 1 { + return false; + } + self.queue.rotate_left(1); + true + } + + pub(crate) fn take_front_for_probe(&mut self, now: Instant) -> Option { + let entry = self.pop_front_internal()?; + let expires_at = now + self.recent_ttl; + self.recent.insert(entry.node.addr, expires_at); + self.recent_expiry.push_back((expires_at, entry.node.addr)); + Some(entry.node) + } + + pub(crate) fn restore_front(&mut self, node: NodeTuple, queued_at: Instant) { + self.recent.remove(&node.addr); + self.queued.insert(node.addr); + self.queue.push_front(QueuedNode { node, queued_at }); + } + + pub(crate) fn contains_recent(&mut self, addr: &SocketAddr, now: Instant) -> bool { + self.expire_recent(now); + self.recent.contains_key(addr) + } + + pub(crate) fn record_probe(&mut self, addr: SocketAddr, now: Instant) { + let expires_at = now + self.recent_ttl; + self.recent.insert(addr, expires_at); + self.recent_expiry.push_back((expires_at, addr)); + } + + pub(crate) fn len(&self) -> usize { + self.queue.len() + } + + #[cfg_attr(not(feature = "metrics"), allow(dead_code))] + pub(crate) fn oldest_age(&self, now: Instant) -> Duration { + self.queue + .front() + .and_then(|entry| now.checked_duration_since(entry.queued_at)) + .unwrap_or_default() + } + + #[cfg(test)] + fn is_warmed(&self) -> bool { + self.warmed + } + + fn pop_front_internal(&mut self) -> Option { + let entry = self.queue.pop_front()?; + self.queued.remove(&entry.node.addr); + Some(entry) + } + + fn expire_recent(&mut self, now: Instant) { + while let Some((expires_at, addr)) = self.recent_expiry.front().copied() { + if expires_at > now { + break; + } + self.recent_expiry.pop_front(); + if self.recent.get(&addr).copied() == Some(expires_at) { + self.recent.remove(&addr); + } + } + } +} + +#[derive(Debug, Clone, Copy)] +struct ResponsiveEntry { + node: NodeTuple, + expires_at: Instant, +} + +/// Fixed-size responsive-node ring. The crawl actor is the only writer. +pub(crate) struct ResponsiveReservoir { + slots: Vec>, + index: AHashMap, + write_cursor: usize, + revisit_cursor: usize, + ttl: Duration, +} + +impl ResponsiveReservoir { + pub(crate) fn new(capacity: usize, ttl: Duration) -> Self { + let capacity = capacity.max(1); + Self { + slots: vec![None; capacity], + index: AHashMap::with_capacity(capacity), + write_cursor: 0, + revisit_cursor: 0, + ttl, + } + } + + pub(crate) fn record(&mut self, node: NodeTuple, now: Instant) { + let entry = ResponsiveEntry { + node, + expires_at: now + self.ttl, + }; + if let Some(slot) = self.index.get(&node.addr).copied() { + self.slots[slot] = Some(entry); + return; + } + + let slot = self.write_cursor; + if let Some(old) = self.slots[slot] + && self.index.get(&old.node.addr).copied() == Some(slot) + { + self.index.remove(&old.node.addr); + } + self.slots[slot] = Some(entry); + self.index.insert(node.addr, slot); + self.write_cursor = (self.write_cursor + 1) % self.slots.len(); + } + + pub(crate) fn next_revisit(&mut self, now: Instant) -> Option { + for _ in 0..self.slots.len() { + let slot = self.revisit_cursor; + self.revisit_cursor = (self.revisit_cursor + 1) % self.slots.len(); + let Some(entry) = self.slots[slot] else { + continue; + }; + if entry.expires_at <= now { + if self.index.get(&entry.node.addr).copied() == Some(slot) { + self.index.remove(&entry.node.addr); + } + self.slots[slot] = None; + continue; + } + return Some(entry.node); + } + None + } + + pub(crate) fn snapshot(&self, limit: usize, now: Instant) -> Vec { + self.slots + .iter() + .filter_map(|entry| { + entry + .filter(|entry| entry.expires_at > now) + .map(|entry| entry.node) + }) + .take(limit) + .collect() + } +} + +#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)] +pub(crate) enum SubnetKey { + V4([u8; 3]), + V6([u8; 8]), +} + +impl SubnetKey { + pub(crate) fn from_addr(addr: &SocketAddr) -> Self { + match addr.ip() { + IpAddr::V4(ip) => { + let octets = ip.octets(); + Self::V4([octets[0], octets[1], octets[2]]) + } + IpAddr::V6(ip) => { + let octets = ip.octets(); + Self::V6(octets[..8].try_into().expect("IPv6 prefix has eight bytes")) + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::net::{IpAddr, Ipv4Addr}; + + fn node(id: u8, addr: &str) -> NodeTuple { + NodeTuple { + id: [id; 20], + addr: addr.parse().unwrap(), + } + } + + #[test] + fn strict_fifo_replaces_oldest_after_warmup() { + let start = Instant::now(); + let mut pool = NodePool::new(2, 600, Duration::from_secs(60), start); + let first = node(1, "8.8.8.8:1"); + let second = node(2, "1.1.1.1:2"); + let third = node(3, "9.9.9.9:3"); + assert_eq!(pool.admit(first, start), AdmissionOutcome::Admitted); + assert_eq!(pool.admit(second, start), AdmissionOutcome::Admitted); + assert!(pool.is_warmed()); + assert_eq!(pool.admit(third, start), AdmissionOutcome::Replaced); + assert_eq!(pool.front(), Some(second)); + } + + #[test] + fn duplicate_does_not_reorder_fifo() { + let start = Instant::now(); + let mut pool = NodePool::new(3, 600, Duration::from_secs(60), start); + let first = node(1, "8.8.8.8:1"); + let second = node(2, "1.1.1.1:2"); + pool.admit(first, start); + pool.admit(second, start); + assert_eq!(pool.admit(first, start), AdmissionOutcome::Duplicate); + assert_eq!(pool.front(), Some(first)); + } + + #[test] + fn rotating_front_preserves_queued_dedup_and_recent_state() { + let start = Instant::now(); + let mut pool = NodePool::new(3, 600, Duration::from_secs(60), start); + let first = node(1, "8.8.8.8:1"); + let second = node(2, "1.1.1.1:2"); + + assert_eq!(pool.admit(first, start), AdmissionOutcome::Admitted); + assert_eq!(pool.admit(second, start), AdmissionOutcome::Admitted); + assert!(pool.rotate_front_to_back()); + + assert_eq!(pool.front(), Some(second)); + assert_eq!(pool.admit(first, start), AdmissionOutcome::Duplicate); + assert!(!pool.contains_recent(&first.addr, start)); + } + + #[test] + fn warmed_pool_enforces_replacement_rate() { + let start = Instant::now(); + let mut pool = NodePool::new(2, 60, Duration::from_secs(60), start); + pool.admit(node(1, "8.8.8.8:1"), start); + pool.admit(node(2, "1.1.1.1:2"), start); + assert_eq!( + pool.admit(node(3, "9.9.9.9:3"), start), + AdmissionOutcome::Replaced + ); + assert_eq!( + pool.admit(node(4, "208.67.222.222:4"), start), + AdmissionOutcome::RateLimited + ); + assert_eq!( + pool.admit(node(4, "208.67.222.222:4"), start + Duration::from_secs(1)), + AdmissionOutcome::Replaced + ); + } + + #[test] + fn recent_probe_blocks_readmission_until_expiry() { + let start = Instant::now(); + let mut pool = NodePool::new(3, 600, Duration::from_secs(10), start); + let first = node(1, "8.8.8.8:1"); + pool.admit(first, start); + assert_eq!(pool.take_front_for_probe(start), Some(first)); + assert_eq!(pool.admit(first, start), AdmissionOutcome::Duplicate); + assert_eq!( + pool.admit(first, start + Duration::from_secs(11)), + AdmissionOutcome::Admitted + ); + } + + #[test] + fn responsive_ring_overwrites_without_growing() { + let start = Instant::now(); + let mut reservoir = ResponsiveReservoir::new(2, Duration::from_secs(10)); + reservoir.record(node(1, "8.8.8.8:1"), start); + reservoir.record(node(2, "1.1.1.1:2"), start); + reservoir.record(node(3, "9.9.9.9:3"), start); + let snapshot = reservoir.snapshot(10, start); + assert_eq!(snapshot.len(), 2); + assert!(!snapshot.iter().any(|entry| entry.id == [1; 20])); + } + + #[test] + #[ignore = "release-only FIFO throughput smoke test"] + fn million_fifo_operations() { + let start = Instant::now(); + let mut pool = NodePool::new(100_000, u32::MAX, Duration::from_secs(600), start); + for value in 0..1_000_000u32 { + let octets = value.to_be_bytes(); + let node = NodeTuple { + id: [octets[3]; 20], + addr: SocketAddr::new( + IpAddr::V4(Ipv4Addr::new(11, octets[1], octets[2], octets[3])), + (value % 65_534 + 1) as u16, + ), + }; + let outcome = pool.admit(node, start); + assert!(!matches!(outcome, AdmissionOutcome::RateLimited)); + } + assert_eq!(pool.len(), 100_000); + eprintln!("1,000,000 FIFO admissions in {:?}", start.elapsed()); + } +} diff --git a/src/protocol.rs b/src/protocol.rs index c9b712c..e12f0ef 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -1,34 +1,51 @@ use serde::Deserialize; -#[derive(Deserialize, Debug)] +#[derive(Deserialize, Debug, Clone)] #[allow(dead_code)] +/// Decoded KRPC envelope. pub struct DhtMessage { + /// Transaction ID bytes. pub t: serde_bytes::ByteBuf, #[allow(dead_code)] + /// Message kind (`q`, `r`, or `e`). pub y: String, #[allow(dead_code)] + /// Query method when `y == q`. pub q: Option, + /// Query arguments. pub a: Option, + /// Response dictionary. pub r: Option, } -#[derive(Deserialize, Debug)] +#[derive(Deserialize, Debug, Clone)] +/// Supported BEP-5 query arguments. pub struct DhtArgs { + /// Sender node ID. pub id: Option, + /// find_node target ID. pub target: Option, + /// get_peers/announce InfoHash. pub info_hash: Option, + /// announce validation token. pub token: Option, + /// Explicit announced Peer port. pub port: Option, + /// Non-zero means use the UDP source port. pub implied_port: Option, } -#[derive(Deserialize, Debug)] +#[derive(Deserialize, Debug, Clone)] +/// Supported BEP-5 response fields. pub struct DhtResponse { #[serde(default)] #[allow(dead_code)] + /// Responder node ID. pub id: Option, #[serde(default)] + /// Compact IPv4 node tuples. pub nodes: Option, #[serde(default)] + /// Compact IPv6 node tuples. pub nodes6: Option, } diff --git a/src/routing_snapshot.rs b/src/routing_snapshot.rs new file mode 100644 index 0000000..b05c1e0 --- /dev/null +++ b/src/routing_snapshot.rs @@ -0,0 +1,37 @@ +use crate::types::NodeTuple; +use rand::seq::SliceRandom; + +#[derive(Default)] +pub(crate) struct RoutingSnapshot { + v4: Vec, + v6: Vec, +} + +impl RoutingSnapshot { + pub(crate) fn from_nodes(nodes: Vec, limit: usize) -> Self { + let mut v4 = Vec::new(); + let mut v6 = Vec::new(); + for node in nodes.into_iter().take(limit) { + if node.addr.is_ipv6() { + v6.push(node); + } else { + v4.push(node); + } + } + Self { v4, v6 } + } + + pub(crate) fn random_nodes(&self, count: usize, filter_ipv6: Option) -> Vec { + let mut rng = rand::thread_rng(); + match filter_ipv6 { + Some(true) => self.v6.choose_multiple(&mut rng, count).cloned().collect(), + Some(false) => self.v4.choose_multiple(&mut rng, count).cloned().collect(), + None => { + let mut all = Vec::with_capacity(self.v4.len() + self.v6.len()); + all.extend_from_slice(&self.v4); + all.extend_from_slice(&self.v6); + all.choose_multiple(&mut rng, count).cloned().collect() + } + } + } +} diff --git a/src/runtime_stats.rs b/src/runtime_stats.rs new file mode 100644 index 0000000..3b648aa --- /dev/null +++ b/src/runtime_stats.rs @@ -0,0 +1,1063 @@ +use std::sync::Arc; +use std::sync::atomic::{AtomicU32, AtomicU64, AtomicUsize, Ordering}; + +pub const METADATA_QUEUE_WAIT_BUCKETS_MS: [u64; 8] = [10, 50, 100, 250, 500, 1_000, 2_000, 5_000]; +pub const METADATA_FETCH_BUCKETS_MS: [u64; 7] = [250, 500, 1_000, 2_000, 4_000, 6_000, 10_000]; +pub const METADATA_SIZE_BUCKETS_BYTES: [u64; 8] = [ + 16 * 1024, + 32 * 1024, + 64 * 1024, + 128 * 1024, + 256 * 1024, + 512 * 1024, + 1024 * 1024, + 10 * 1024 * 1024, +]; + +#[derive(Debug, Clone, PartialEq, Eq)] +/// Snapshot of a non-cumulative fixed-bucket histogram. +pub struct FixedHistogramSnapshot { + /// Inclusive upper bound for each bucket. + pub bounds: Vec, + /// Per-bucket counts; values are not cumulative. + pub counts: Vec, + /// Values larger than the final bound. + pub overflow: u64, + /// Total observations including overflow. + pub count: u64, + /// Wrapping sum of all observed values. + pub sum: u64, +} + +impl FixedHistogramSnapshot { + /// Returns the inclusive bucket bound containing the requested approximate percentile. + /// + /// `percentile` is clamped to `0.0..=1.0`. Overflow observations return the final finite + /// bound because the histogram intentionally stores no dynamic maximum. + pub fn percentile(&self, percentile: f64) -> Option { + if self.count == 0 { + return None; + } + let rank = ((self.count as f64 * percentile.clamp(0.0, 1.0)).ceil() as u64).max(1); + let mut cumulative = 0u64; + for (bound, count) in self.bounds.iter().zip(&self.counts) { + cumulative = cumulative.saturating_add(*count); + if cumulative >= rank { + return Some(*bound); + } + } + self.bounds.last().copied() + } +} + +struct AtomicFixedHistogram { + bounds: [u64; N], + counts: [AtomicU64; N], + overflow: AtomicU64, + count: AtomicU64, + sum: AtomicU64, +} + +impl AtomicFixedHistogram { + fn new(bounds: [u64; N]) -> Self { + Self { + bounds, + counts: std::array::from_fn(|_| AtomicU64::new(0)), + overflow: AtomicU64::new(0), + count: AtomicU64::new(0), + sum: AtomicU64::new(0), + } + } + + fn record(&self, value: u64) { + if let Some(index) = self.bounds.iter().position(|bound| value <= *bound) { + self.counts[index].fetch_add(1, Ordering::Relaxed); + } else { + self.overflow.fetch_add(1, Ordering::Relaxed); + } + self.count.fetch_add(1, Ordering::Relaxed); + self.sum.fetch_add(value, Ordering::Relaxed); + } + + fn snapshot(&self) -> FixedHistogramSnapshot { + FixedHistogramSnapshot { + bounds: self.bounds.to_vec(), + counts: self + .counts + .iter() + .map(|count| count.load(Ordering::Relaxed)) + .collect(), + overflow: self.overflow.load(Ordering::Relaxed), + count: self.count.load(Ordering::Relaxed), + sum: self.sum.load(Ordering::Relaxed), + } + } +} + +impl Default for AtomicFixedHistogram { + fn default() -> Self { + Self::new([0; N]) + } +} + +/// Transport-neutral counters and fixed-bucket histograms used by dashboards. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct DhtObservabilitySnapshot { + /// UDP datagrams received before validation. + pub udp_rx_packets: u64, + /// UDP bytes received before validation. + pub udp_rx_bytes: u64, + /// Successfully sent UDP query and response datagrams. + pub udp_tx_packets: u64, + /// Successfully sent UDP query and response bytes. + pub udp_tx_bytes: u64, + /// Inbound ping queries. + pub inbound_ping: u64, + /// Inbound find_node queries. + pub inbound_find_node: u64, + /// Inbound get_peers queries. + pub inbound_get_peers: u64, + /// Inbound announce_peer queries. + pub inbound_announce_peer: u64, + /// Other or invalid inbound query names. + pub inbound_other: u64, + /// Replies admitted by the regular response budget. + pub response_normal: u64, + /// Priority ping/get_peers replies admitted by the reserve. + pub response_fallback: u64, + /// Replies rejected by final response limits. + pub response_rate_limited: u64, + /// Admitted replies whose `send_to` failed. + pub response_send_failed: u64, + /// Valid, unfiltered announces with a usable Peer port. + pub announce_accepted: u64, + /// Announces with a missing or invalid token. + pub announce_invalid_token: u64, + /// Announces rejected by the application Hash filter. + pub announce_filtered: u64, + /// New nodes admitted to the FIFO pool. + pub node_admitted: u64, + /// Full-pool replacements. + pub node_replaced: u64, + /// Nodes rejected as queued or recently probed duplicates. + pub node_dropped_duplicate: u64, + /// Nodes rejected by the replacement budget. + pub node_dropped_rate_limited: u64, + /// Nodes rejected because their endpoint is not usable. + pub node_dropped_invalid: u64, + /// BEP-9 Metadata piece payload bytes received. + pub metadata_bytes_downloaded: u64, + /// Metadata Peer attempts reaching the end-to-end timeout. + pub metadata_failure_timeout: u64, + /// Metadata Peer connection failures. + pub metadata_failure_connect: u64, + /// Peers without extension-protocol support. + pub metadata_failure_no_extension: u64, + /// Extension handshake or piece-request send failures. + pub metadata_failure_send: u64, + /// Metadata payloads exceeding the 10 MiB limit. + pub metadata_failure_size_limit: u64, + /// Metadata payloads failing InfoHash SHA1 validation. + pub metadata_failure_sha1: u64, + /// Validated payloads that could not be parsed as torrent info. + pub metadata_failure_parse: u64, + /// Other receive, piece or worker failures. + pub metadata_failure_other: u64, + /// Failure-cache hits for previously timed-out Peers. + pub peer_cache_timeout_hits: u64, + /// Failure-cache hits for previous connection failures. + pub peer_cache_connect_hits: u64, + /// Queue wait observations in milliseconds. + pub queue_wait_ms: FixedHistogramSnapshot, + /// End-to-end Peer attempt observations in milliseconds. + pub fetch_duration_ms: FixedHistogramSnapshot, + /// Complete bencoded info payload observations in bytes. + pub metadata_size_bytes: FixedHistogramSnapshot, +} + +/// Cheap, cloneable handle for reading live DHT runtime statistics. +#[derive(Clone, Default)] +pub struct DhtRuntimeStats { + inner: Arc, +} + +/// Point-in-time view returned by [`DhtRuntimeStats::snapshot`]. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct DhtRuntimeSnapshot { + /// Valid announce hashes observed before ingress publication. + pub hashes_received: u64, + /// Hashes dropped because bounded ingress was full. + pub hash_ingress_dropped: u64, + /// Current hash-ingress queue depth. + pub hash_ingress_queue_depth: usize, + /// Configured hash-ingress queue capacity. + pub hash_ingress_queue_capacity: usize, + + /// Current priority crawl-event depth. + pub crawl_priority_queue_depth: usize, + /// Configured priority crawl-event capacity. + pub crawl_priority_queue_capacity: usize, + /// Current discovered-node event depth. + pub crawl_discovery_queue_depth: usize, + /// Configured discovered-node event capacity. + pub crawl_discovery_queue_capacity: usize, + + /// Deduplicated pending Metadata Hashes. + pub metadata_queue_depth: usize, + /// Configured pending Metadata capacity. + pub metadata_queue_max: usize, + /// Current Metadata jobs. + pub metadata_in_flight: usize, + /// Unique Hash insertions, including insertions that evicted an older Hash. + pub metadata_queue_inserted: u64, + /// Duplicate Hash updates merged into existing entries. + pub metadata_queue_deduplicated: u64, + /// Oldest Hashes evicted by newer events. + pub metadata_queue_evicted: u64, + /// Events rejected or expired because they were stale. + pub metadata_queue_stale: u64, + /// Real Peer network attempts. + pub metadata_peer_attempts: u64, + /// Successful Peer downloads and parses. + pub metadata_peer_succeeded: u64, + /// Failed real Peer attempts. + pub metadata_peer_failed: u64, + /// End-to-end Peer timeouts. + pub metadata_peer_timeouts: u64, + /// Peer connection failures. + pub metadata_connect_failed: u64, + /// Peers without extension-protocol support. + pub metadata_no_extension: u64, + /// Peer failure-cache skips. + pub metadata_peer_failure_cache_hits: u64, + /// Current Peer failure-cache entries. + pub metadata_peer_failure_cache_entries: usize, + + /// Current strict FIFO crawl-pool size. + pub node_pool_size: usize, + /// Configured FIFO capacity. + pub node_pool_capacity: usize, + /// Bootstrap low-water mark. + pub node_pool_low_watermark: usize, + + /// Current pending find_node transactions. + pub find_node_in_flight: usize, + /// Configured total find_node in-flight limit. + pub find_node_in_flight_max: usize, + /// Metadata-pressure-adjusted find_node budget per second. + pub find_node_effective_rate_per_sec: u32, + /// Queries sent to never-before-probed destinations. + pub queries_new: u64, + /// Queries sent to responsive revisit nodes. + pub queries_revisit: u64, + /// Queries sent to bootstrap endpoints. + pub queries_bootstrap: u64, + /// Replies matched to pending transactions. + pub responses: u64, + /// Replies without a matching pending transaction. + pub unmatched_responses: u64, + /// Pending transactions that expired. + pub timeouts: u64, + /// Outbound find_node send failures. + pub send_failures: u64, + /// Response events dropped before reaching the actor. + pub crawl_events_dropped_response: u64, + /// Discovered-node events dropped before reaching the actor. + pub crawl_events_dropped_discovered: u64, + + /// UDP datagrams received before validation. + pub udp_received: u64, + /// Valid datagrams dropped because every worker queue was full. + pub udp_queue_full: u64, + /// Empty, oversized or non-bencoded UDP datagrams. + pub udp_invalid: u64, + /// DHT replies rejected by final rate limits. + pub udp_responses_rate_limited: u64, + /// Priority replies admitted by the ping/get_peers reserve. + pub udp_responses_priority_reserved: u64, +} + +struct DhtRuntimeStatsInner { + hashes_received: AtomicU64, + hash_ingress_dropped: AtomicU64, + hash_ingress_queue_depth: AtomicUsize, + hash_ingress_queue_capacity: usize, + + crawl_priority_queue_depth: AtomicUsize, + crawl_priority_queue_capacity: usize, + crawl_discovery_queue_depth: AtomicUsize, + crawl_discovery_queue_capacity: usize, + + metadata_queue_depth: AtomicUsize, + metadata_queue_max: usize, + metadata_in_flight: AtomicUsize, + metadata_queue_inserted: AtomicU64, + metadata_queue_deduplicated: AtomicU64, + metadata_queue_evicted: AtomicU64, + metadata_queue_stale: AtomicU64, + metadata_peer_attempts: AtomicU64, + metadata_peer_succeeded: AtomicU64, + metadata_peer_failed: AtomicU64, + metadata_peer_timeouts: AtomicU64, + metadata_connect_failed: AtomicU64, + metadata_no_extension: AtomicU64, + metadata_peer_failure_cache_hits: AtomicU64, + metadata_peer_failure_cache_entries: AtomicUsize, + + node_pool_size: AtomicUsize, + node_pool_capacity: usize, + node_pool_low_watermark: usize, + + find_node_in_flight: AtomicUsize, + find_node_in_flight_max: usize, + find_node_effective_rate_per_sec: AtomicU32, + queries_new: AtomicU64, + queries_revisit: AtomicU64, + queries_bootstrap: AtomicU64, + responses: AtomicU64, + unmatched_responses: AtomicU64, + timeouts: AtomicU64, + send_failures: AtomicU64, + crawl_events_dropped_response: AtomicU64, + crawl_events_dropped_discovered: AtomicU64, + + udp_received: AtomicU64, + udp_queue_full: AtomicU64, + udp_invalid: AtomicU64, + udp_responses_rate_limited: AtomicU64, + udp_responses_priority_reserved: AtomicU64, + + udp_rx_bytes: AtomicU64, + udp_tx_packets: AtomicU64, + udp_tx_bytes: AtomicU64, + inbound_ping: AtomicU64, + inbound_find_node: AtomicU64, + inbound_get_peers: AtomicU64, + inbound_announce_peer: AtomicU64, + inbound_other: AtomicU64, + response_normal: AtomicU64, + response_send_failed: AtomicU64, + announce_accepted: AtomicU64, + announce_invalid_token: AtomicU64, + announce_filtered: AtomicU64, + node_admitted: AtomicU64, + node_replaced: AtomicU64, + node_dropped_duplicate: AtomicU64, + node_dropped_rate_limited: AtomicU64, + node_dropped_invalid: AtomicU64, + metadata_bytes_downloaded: AtomicU64, + metadata_failure_timeout: AtomicU64, + metadata_failure_connect: AtomicU64, + metadata_failure_no_extension: AtomicU64, + metadata_failure_send: AtomicU64, + metadata_failure_size_limit: AtomicU64, + metadata_failure_sha1: AtomicU64, + metadata_failure_parse: AtomicU64, + metadata_failure_other: AtomicU64, + peer_cache_timeout_hits: AtomicU64, + peer_cache_connect_hits: AtomicU64, + queue_wait_ms: AtomicFixedHistogram<8>, + fetch_duration_ms: AtomicFixedHistogram<7>, + metadata_size_bytes: AtomicFixedHistogram<8>, +} + +impl Default for DhtRuntimeStatsInner { + fn default() -> Self { + Self { + queue_wait_ms: AtomicFixedHistogram::new(METADATA_QUEUE_WAIT_BUCKETS_MS), + fetch_duration_ms: AtomicFixedHistogram::new(METADATA_FETCH_BUCKETS_MS), + metadata_size_bytes: AtomicFixedHistogram::new(METADATA_SIZE_BUCKETS_BYTES), + hashes_received: AtomicU64::new(0), + hash_ingress_dropped: AtomicU64::new(0), + hash_ingress_queue_depth: AtomicUsize::new(0), + hash_ingress_queue_capacity: 0, + crawl_priority_queue_depth: AtomicUsize::new(0), + crawl_priority_queue_capacity: 0, + crawl_discovery_queue_depth: AtomicUsize::new(0), + crawl_discovery_queue_capacity: 0, + metadata_queue_depth: AtomicUsize::new(0), + metadata_queue_max: 0, + metadata_in_flight: AtomicUsize::new(0), + metadata_queue_inserted: AtomicU64::new(0), + metadata_queue_deduplicated: AtomicU64::new(0), + metadata_queue_evicted: AtomicU64::new(0), + metadata_queue_stale: AtomicU64::new(0), + metadata_peer_attempts: AtomicU64::new(0), + metadata_peer_succeeded: AtomicU64::new(0), + metadata_peer_failed: AtomicU64::new(0), + metadata_peer_timeouts: AtomicU64::new(0), + metadata_connect_failed: AtomicU64::new(0), + metadata_no_extension: AtomicU64::new(0), + metadata_peer_failure_cache_hits: AtomicU64::new(0), + metadata_peer_failure_cache_entries: AtomicUsize::new(0), + node_pool_size: AtomicUsize::new(0), + node_pool_capacity: 0, + node_pool_low_watermark: 0, + find_node_in_flight: AtomicUsize::new(0), + find_node_in_flight_max: 0, + find_node_effective_rate_per_sec: AtomicU32::new(0), + queries_new: AtomicU64::new(0), + queries_revisit: AtomicU64::new(0), + queries_bootstrap: AtomicU64::new(0), + responses: AtomicU64::new(0), + unmatched_responses: AtomicU64::new(0), + timeouts: AtomicU64::new(0), + send_failures: AtomicU64::new(0), + crawl_events_dropped_response: AtomicU64::new(0), + crawl_events_dropped_discovered: AtomicU64::new(0), + udp_received: AtomicU64::new(0), + udp_queue_full: AtomicU64::new(0), + udp_invalid: AtomicU64::new(0), + udp_responses_rate_limited: AtomicU64::new(0), + udp_responses_priority_reserved: AtomicU64::new(0), + udp_rx_bytes: AtomicU64::new(0), + udp_tx_packets: AtomicU64::new(0), + udp_tx_bytes: AtomicU64::new(0), + inbound_ping: AtomicU64::new(0), + inbound_find_node: AtomicU64::new(0), + inbound_get_peers: AtomicU64::new(0), + inbound_announce_peer: AtomicU64::new(0), + inbound_other: AtomicU64::new(0), + response_normal: AtomicU64::new(0), + response_send_failed: AtomicU64::new(0), + announce_accepted: AtomicU64::new(0), + announce_invalid_token: AtomicU64::new(0), + announce_filtered: AtomicU64::new(0), + node_admitted: AtomicU64::new(0), + node_replaced: AtomicU64::new(0), + node_dropped_duplicate: AtomicU64::new(0), + node_dropped_rate_limited: AtomicU64::new(0), + node_dropped_invalid: AtomicU64::new(0), + metadata_bytes_downloaded: AtomicU64::new(0), + metadata_failure_timeout: AtomicU64::new(0), + metadata_failure_connect: AtomicU64::new(0), + metadata_failure_no_extension: AtomicU64::new(0), + metadata_failure_send: AtomicU64::new(0), + metadata_failure_size_limit: AtomicU64::new(0), + metadata_failure_sha1: AtomicU64::new(0), + metadata_failure_parse: AtomicU64::new(0), + metadata_failure_other: AtomicU64::new(0), + peer_cache_timeout_hits: AtomicU64::new(0), + peer_cache_connect_hits: AtomicU64::new(0), + } + } +} + +#[derive(Debug, Clone, Copy, Default)] +pub(crate) struct DhtRuntimeLimits { + pub metadata_queue: usize, + pub node_pool: usize, + pub node_pool_low_watermark: usize, + pub find_node_in_flight: usize, + pub initial_find_node_rate: u32, + pub hash_ingress_queue: usize, + pub crawl_priority_queue: usize, + pub crawl_discovery_queue: usize, +} + +impl DhtRuntimeStats { + pub(crate) fn with_limits(limits: DhtRuntimeLimits) -> Self { + Self { + inner: Arc::new(DhtRuntimeStatsInner { + metadata_queue_max: limits.metadata_queue, + node_pool_capacity: limits.node_pool, + node_pool_low_watermark: limits.node_pool_low_watermark, + find_node_in_flight_max: limits.find_node_in_flight, + find_node_effective_rate_per_sec: AtomicU32::new(limits.initial_find_node_rate), + hash_ingress_queue_capacity: limits.hash_ingress_queue, + crawl_priority_queue_capacity: limits.crawl_priority_queue, + crawl_discovery_queue_capacity: limits.crawl_discovery_queue, + ..DhtRuntimeStatsInner::default() + }), + } + } + + /// Load every exposed value once using relaxed atomics. + /// Loads the core runtime counters and gauges using relaxed atomics. + pub fn snapshot(&self) -> DhtRuntimeSnapshot { + let inner = &self.inner; + DhtRuntimeSnapshot { + hashes_received: inner.hashes_received.load(Ordering::Relaxed), + hash_ingress_dropped: inner.hash_ingress_dropped.load(Ordering::Relaxed), + hash_ingress_queue_depth: inner.hash_ingress_queue_depth.load(Ordering::Relaxed), + hash_ingress_queue_capacity: inner.hash_ingress_queue_capacity, + crawl_priority_queue_depth: inner.crawl_priority_queue_depth.load(Ordering::Relaxed), + crawl_priority_queue_capacity: inner.crawl_priority_queue_capacity, + crawl_discovery_queue_depth: inner.crawl_discovery_queue_depth.load(Ordering::Relaxed), + crawl_discovery_queue_capacity: inner.crawl_discovery_queue_capacity, + metadata_queue_depth: inner.metadata_queue_depth.load(Ordering::Relaxed), + metadata_queue_max: inner.metadata_queue_max, + metadata_in_flight: inner.metadata_in_flight.load(Ordering::Relaxed), + metadata_queue_inserted: inner.metadata_queue_inserted.load(Ordering::Relaxed), + metadata_queue_deduplicated: inner.metadata_queue_deduplicated.load(Ordering::Relaxed), + metadata_queue_evicted: inner.metadata_queue_evicted.load(Ordering::Relaxed), + metadata_queue_stale: inner.metadata_queue_stale.load(Ordering::Relaxed), + metadata_peer_attempts: inner.metadata_peer_attempts.load(Ordering::Relaxed), + metadata_peer_succeeded: inner.metadata_peer_succeeded.load(Ordering::Relaxed), + metadata_peer_failed: inner.metadata_peer_failed.load(Ordering::Relaxed), + metadata_peer_timeouts: inner.metadata_peer_timeouts.load(Ordering::Relaxed), + metadata_connect_failed: inner.metadata_connect_failed.load(Ordering::Relaxed), + metadata_no_extension: inner.metadata_no_extension.load(Ordering::Relaxed), + metadata_peer_failure_cache_hits: inner + .metadata_peer_failure_cache_hits + .load(Ordering::Relaxed), + metadata_peer_failure_cache_entries: inner + .metadata_peer_failure_cache_entries + .load(Ordering::Relaxed), + node_pool_size: inner.node_pool_size.load(Ordering::Relaxed), + node_pool_capacity: inner.node_pool_capacity, + node_pool_low_watermark: inner.node_pool_low_watermark, + find_node_in_flight: inner.find_node_in_flight.load(Ordering::Relaxed), + find_node_in_flight_max: inner.find_node_in_flight_max, + find_node_effective_rate_per_sec: inner + .find_node_effective_rate_per_sec + .load(Ordering::Relaxed), + queries_new: inner.queries_new.load(Ordering::Relaxed), + queries_revisit: inner.queries_revisit.load(Ordering::Relaxed), + queries_bootstrap: inner.queries_bootstrap.load(Ordering::Relaxed), + responses: inner.responses.load(Ordering::Relaxed), + unmatched_responses: inner.unmatched_responses.load(Ordering::Relaxed), + timeouts: inner.timeouts.load(Ordering::Relaxed), + send_failures: inner.send_failures.load(Ordering::Relaxed), + crawl_events_dropped_response: inner + .crawl_events_dropped_response + .load(Ordering::Relaxed), + crawl_events_dropped_discovered: inner + .crawl_events_dropped_discovered + .load(Ordering::Relaxed), + udp_received: inner.udp_received.load(Ordering::Relaxed), + udp_queue_full: inner.udp_queue_full.load(Ordering::Relaxed), + udp_invalid: inner.udp_invalid.load(Ordering::Relaxed), + udp_responses_rate_limited: inner.udp_responses_rate_limited.load(Ordering::Relaxed), + udp_responses_priority_reserved: inner + .udp_responses_priority_reserved + .load(Ordering::Relaxed), + } + } + + /// Loads detailed transport counters, failure categories and fixed histograms. + pub fn observability_snapshot(&self) -> DhtObservabilitySnapshot { + let inner = &self.inner; + DhtObservabilitySnapshot { + udp_rx_packets: inner.udp_received.load(Ordering::Relaxed), + udp_rx_bytes: inner.udp_rx_bytes.load(Ordering::Relaxed), + udp_tx_packets: inner.udp_tx_packets.load(Ordering::Relaxed), + udp_tx_bytes: inner.udp_tx_bytes.load(Ordering::Relaxed), + inbound_ping: inner.inbound_ping.load(Ordering::Relaxed), + inbound_find_node: inner.inbound_find_node.load(Ordering::Relaxed), + inbound_get_peers: inner.inbound_get_peers.load(Ordering::Relaxed), + inbound_announce_peer: inner.inbound_announce_peer.load(Ordering::Relaxed), + inbound_other: inner.inbound_other.load(Ordering::Relaxed), + response_normal: inner.response_normal.load(Ordering::Relaxed), + response_fallback: inner + .udp_responses_priority_reserved + .load(Ordering::Relaxed), + response_rate_limited: inner.udp_responses_rate_limited.load(Ordering::Relaxed), + response_send_failed: inner.response_send_failed.load(Ordering::Relaxed), + announce_accepted: inner.announce_accepted.load(Ordering::Relaxed), + announce_invalid_token: inner.announce_invalid_token.load(Ordering::Relaxed), + announce_filtered: inner.announce_filtered.load(Ordering::Relaxed), + node_admitted: inner.node_admitted.load(Ordering::Relaxed), + node_replaced: inner.node_replaced.load(Ordering::Relaxed), + node_dropped_duplicate: inner.node_dropped_duplicate.load(Ordering::Relaxed), + node_dropped_rate_limited: inner.node_dropped_rate_limited.load(Ordering::Relaxed), + node_dropped_invalid: inner.node_dropped_invalid.load(Ordering::Relaxed), + metadata_bytes_downloaded: inner.metadata_bytes_downloaded.load(Ordering::Relaxed), + metadata_failure_timeout: inner.metadata_failure_timeout.load(Ordering::Relaxed), + metadata_failure_connect: inner.metadata_failure_connect.load(Ordering::Relaxed), + metadata_failure_no_extension: inner + .metadata_failure_no_extension + .load(Ordering::Relaxed), + metadata_failure_send: inner.metadata_failure_send.load(Ordering::Relaxed), + metadata_failure_size_limit: inner.metadata_failure_size_limit.load(Ordering::Relaxed), + metadata_failure_sha1: inner.metadata_failure_sha1.load(Ordering::Relaxed), + metadata_failure_parse: inner.metadata_failure_parse.load(Ordering::Relaxed), + metadata_failure_other: inner.metadata_failure_other.load(Ordering::Relaxed), + peer_cache_timeout_hits: inner.peer_cache_timeout_hits.load(Ordering::Relaxed), + peer_cache_connect_hits: inner.peer_cache_connect_hits.load(Ordering::Relaxed), + queue_wait_ms: inner.queue_wait_ms.snapshot(), + fetch_duration_ms: inner.fetch_duration_ms.snapshot(), + metadata_size_bytes: inner.metadata_size_bytes.snapshot(), + } + } + + pub(crate) fn hash_received(&self) { + self.inner.hashes_received.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn hash_ingress_dropped(&self) { + self.inner + .hash_ingress_dropped + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn set_hash_ingress_queue_depth(&self, depth: usize) { + self.inner + .hash_ingress_queue_depth + .store(depth, Ordering::Relaxed); + } + + pub(crate) fn set_crawl_priority_queue_depth(&self, depth: usize) { + self.inner + .crawl_priority_queue_depth + .store(depth, Ordering::Relaxed); + } + + pub(crate) fn set_crawl_discovery_queue_depth(&self, depth: usize) { + self.inner + .crawl_discovery_queue_depth + .store(depth, Ordering::Relaxed); + } + + pub(crate) fn set_metadata_queue(&self, depth: usize, in_flight: usize) { + self.inner + .metadata_queue_depth + .store(depth, Ordering::Relaxed); + self.inner + .metadata_in_flight + .store(in_flight, Ordering::Relaxed); + } + + pub(crate) fn metadata_queue_deduplicated(&self) { + self.inner + .metadata_queue_deduplicated + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_queue_inserted(&self) { + self.inner + .metadata_queue_inserted + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_queue_evicted(&self) { + self.inner + .metadata_queue_evicted + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_queue_stale(&self, count: usize) { + self.inner + .metadata_queue_stale + .fetch_add(count as u64, Ordering::Relaxed); + } + + pub(crate) fn metadata_peer_attempt(&self) { + self.inner + .metadata_peer_attempts + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_peer_succeeded(&self) { + self.inner + .metadata_peer_succeeded + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_peer_failed(&self) { + self.inner + .metadata_peer_failed + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_peer_timeout(&self) { + self.inner + .metadata_peer_timeouts + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_connect_failed(&self) { + self.inner + .metadata_connect_failed + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_no_extension(&self) { + self.inner + .metadata_no_extension + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_peer_failure_cache_hit(&self) { + self.inner + .metadata_peer_failure_cache_hits + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn set_metadata_peer_failure_cache_entries(&self, count: usize) { + self.inner + .metadata_peer_failure_cache_entries + .store(count, Ordering::Relaxed); + } + + pub(crate) fn set_node_pool_size(&self, size: usize) { + self.inner.node_pool_size.store(size, Ordering::Relaxed); + } + + pub(crate) fn set_find_node_in_flight(&self, count: usize) { + self.inner + .find_node_in_flight + .store(count, Ordering::Relaxed); + } + + pub(crate) fn set_find_node_effective_rate(&self, rate: u32) { + self.inner + .find_node_effective_rate_per_sec + .store(rate, Ordering::Relaxed); + } + + pub(crate) fn query_new(&self) { + self.inner.queries_new.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn query_revisit(&self) { + self.inner.queries_revisit.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn query_bootstrap(&self) { + self.inner.queries_bootstrap.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn response(&self) { + self.inner.responses.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn unmatched_response(&self) { + self.inner + .unmatched_responses + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn timeout(&self) { + self.inner.timeouts.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn send_failure(&self) { + self.inner.send_failures.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn crawl_event_dropped_response(&self) { + self.inner + .crawl_events_dropped_response + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn crawl_event_dropped_discovered(&self) { + self.inner + .crawl_events_dropped_discovered + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn udp_received(&self) { + self.inner.udp_received.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn udp_received_bytes(&self, bytes: usize) { + self.inner + .udp_rx_bytes + .fetch_add(bytes as u64, Ordering::Relaxed); + } + + pub(crate) fn udp_sent(&self, bytes: usize) { + self.inner.udp_tx_packets.fetch_add(1, Ordering::Relaxed); + self.inner + .udp_tx_bytes + .fetch_add(bytes as u64, Ordering::Relaxed); + } + + pub(crate) fn inbound_query(&self, query: &str) { + let counter = match query { + "ping" => &self.inner.inbound_ping, + "find_node" => &self.inner.inbound_find_node, + "get_peers" => &self.inner.inbound_get_peers, + "announce_peer" => &self.inner.inbound_announce_peer, + _ => &self.inner.inbound_other, + }; + counter.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn response_normal(&self) { + self.inner.response_normal.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn response_send_failed(&self) { + self.inner + .response_send_failed + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn announce_accepted(&self) { + self.inner.announce_accepted.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn announce_invalid_token(&self) { + self.inner + .announce_invalid_token + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn announce_filtered(&self) { + self.inner.announce_filtered.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn node_admitted(&self) { + self.inner.node_admitted.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn node_replaced(&self) { + self.inner.node_replaced.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn node_dropped_duplicate(&self) { + self.inner + .node_dropped_duplicate + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn node_dropped_rate_limited(&self) { + self.inner + .node_dropped_rate_limited + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn node_dropped_invalid(&self) { + self.inner + .node_dropped_invalid + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_bytes_downloaded(&self, bytes: usize) { + self.inner + .metadata_bytes_downloaded + .fetch_add(bytes as u64, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_timeout(&self) { + self.inner + .metadata_failure_timeout + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_connect(&self) { + self.inner + .metadata_failure_connect + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_no_extension(&self) { + self.inner + .metadata_failure_no_extension + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_send(&self) { + self.inner + .metadata_failure_send + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_size_limit(&self) { + self.inner + .metadata_failure_size_limit + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_sha1(&self) { + self.inner + .metadata_failure_sha1 + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_parse(&self) { + self.inner + .metadata_failure_parse + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn metadata_failure_other(&self) { + self.inner + .metadata_failure_other + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn peer_cache_hit_timeout(&self) { + self.inner + .peer_cache_timeout_hits + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn peer_cache_hit_connect(&self) { + self.inner + .peer_cache_connect_hits + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn observe_metadata_queue_wait(&self, millis: u64) { + self.inner.queue_wait_ms.record(millis); + } + + pub(crate) fn observe_metadata_fetch_duration(&self, millis: u64) { + self.inner.fetch_duration_ms.record(millis); + } + + pub(crate) fn observe_metadata_size(&self, bytes: usize) { + self.inner.metadata_size_bytes.record(bytes as u64); + } + + pub(crate) fn udp_queue_full(&self) { + self.inner.udp_queue_full.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn udp_invalid(&self) { + self.inner.udp_invalid.fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn udp_response_rate_limited(&self) { + self.inner + .udp_responses_rate_limited + .fetch_add(1, Ordering::Relaxed); + } + + pub(crate) fn udp_response_priority_reserved(&self) { + self.inner + .udp_responses_priority_reserved + .fetch_add(1, Ordering::Relaxed); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn cloned_handle_updates_one_snapshot() { + let stats = DhtRuntimeStats::with_limits(DhtRuntimeLimits { + metadata_queue: 100, + node_pool: 1_000, + node_pool_low_watermark: 10, + find_node_in_flight: 512, + initial_find_node_rate: 200, + hash_ingress_queue: 20, + crawl_priority_queue: 30, + crawl_discovery_queue: 40, + }); + let writer = stats.clone(); + + writer.hash_received(); + writer.hash_ingress_dropped(); + writer.set_hash_ingress_queue_depth(11); + writer.set_crawl_priority_queue_depth(12); + writer.set_crawl_discovery_queue_depth(13); + writer.set_metadata_queue(42, 7); + writer.metadata_queue_inserted(); + writer.metadata_queue_deduplicated(); + writer.metadata_queue_evicted(); + writer.metadata_queue_stale(3); + writer.metadata_peer_attempt(); + writer.metadata_peer_succeeded(); + writer.metadata_peer_failed(); + writer.metadata_peer_timeout(); + writer.metadata_connect_failed(); + writer.metadata_no_extension(); + writer.metadata_peer_failure_cache_hit(); + writer.set_metadata_peer_failure_cache_entries(9); + writer.set_node_pool_size(321); + writer.set_find_node_in_flight(12); + writer.set_find_node_effective_rate(150); + writer.query_new(); + writer.query_revisit(); + writer.query_bootstrap(); + writer.response(); + writer.unmatched_response(); + writer.timeout(); + writer.send_failure(); + writer.crawl_event_dropped_response(); + writer.crawl_event_dropped_discovered(); + writer.udp_received(); + writer.udp_queue_full(); + writer.udp_invalid(); + writer.udp_response_rate_limited(); + writer.udp_response_priority_reserved(); + + assert_eq!( + stats.snapshot(), + DhtRuntimeSnapshot { + hashes_received: 1, + hash_ingress_dropped: 1, + hash_ingress_queue_depth: 11, + hash_ingress_queue_capacity: 20, + crawl_priority_queue_depth: 12, + crawl_priority_queue_capacity: 30, + crawl_discovery_queue_depth: 13, + crawl_discovery_queue_capacity: 40, + metadata_queue_depth: 42, + metadata_queue_max: 100, + metadata_in_flight: 7, + metadata_queue_inserted: 1, + metadata_queue_deduplicated: 1, + metadata_queue_evicted: 1, + metadata_queue_stale: 3, + metadata_peer_attempts: 1, + metadata_peer_succeeded: 1, + metadata_peer_failed: 1, + metadata_peer_timeouts: 1, + metadata_connect_failed: 1, + metadata_no_extension: 1, + metadata_peer_failure_cache_hits: 1, + metadata_peer_failure_cache_entries: 9, + node_pool_size: 321, + node_pool_capacity: 1_000, + node_pool_low_watermark: 10, + find_node_in_flight: 12, + find_node_in_flight_max: 512, + find_node_effective_rate_per_sec: 150, + queries_new: 1, + queries_revisit: 1, + queries_bootstrap: 1, + responses: 1, + unmatched_responses: 1, + timeouts: 1, + send_failures: 1, + crawl_events_dropped_response: 1, + crawl_events_dropped_discovered: 1, + udp_received: 1, + udp_queue_full: 1, + udp_invalid: 1, + udp_responses_rate_limited: 1, + udp_responses_priority_reserved: 1, + } + ); + } + + #[test] + fn observability_snapshot_tracks_fixed_histograms_and_categories() { + let stats = DhtRuntimeStats::default(); + stats.udp_received(); + stats.udp_received_bytes(128); + stats.udp_sent(64); + stats.inbound_query("ping"); + stats.inbound_query("unknown"); + stats.node_admitted(); + stats.metadata_bytes_downloaded(1_024); + stats.metadata_failure_parse(); + stats.observe_metadata_queue_wait(75); + stats.observe_metadata_queue_wait(600); + stats.observe_metadata_fetch_duration(1_500); + stats.observe_metadata_size(70_000); + + let snapshot = stats.observability_snapshot(); + assert_eq!(snapshot.udp_rx_packets, 1); + assert_eq!(snapshot.udp_rx_bytes, 128); + assert_eq!(snapshot.udp_tx_packets, 1); + assert_eq!(snapshot.udp_tx_bytes, 64); + assert_eq!(snapshot.inbound_ping, 1); + assert_eq!(snapshot.inbound_other, 1); + assert_eq!(snapshot.node_admitted, 1); + assert_eq!(snapshot.metadata_bytes_downloaded, 1_024); + assert_eq!(snapshot.metadata_failure_parse, 1); + assert_eq!(snapshot.queue_wait_ms.percentile(0.50), Some(100)); + assert_eq!(snapshot.queue_wait_ms.percentile(0.95), Some(1_000)); + assert_eq!(snapshot.fetch_duration_ms.percentile(0.95), Some(2_000)); + assert_eq!(snapshot.metadata_size_bytes.percentile(0.50), Some(131_072)); + } +} diff --git a/src/scheduler.rs b/src/scheduler.rs index f2d9b3b..7def3f7 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -1,214 +1,628 @@ -use crate::metadata::RbitFetcher; +use crate::metadata::{MetadataFetchOutcome, RbitFetcher}; +#[cfg(test)] +use crate::runtime_stats::DhtRuntimeLimits; +use crate::runtime_stats::DhtRuntimeStats; use crate::server::HashDiscovered; -use crate::types::TorrentInfo; +use crate::types::{MetadataFetchCompletion, MetadataFetchCompletionStatus, TorrentInfo}; +use arc_swap::ArcSwapOption; +#[cfg(feature = "metrics")] +use metrics::{counter, gauge, histogram}; +use std::collections::{BTreeMap, HashMap, HashSet}; +use std::net::SocketAddr; +use std::panic::{AssertUnwindSafe, catch_unwind}; +use std::sync::Arc; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; -use std::sync::{Arc, RwLock}; +use std::time::{Duration, Instant}; use tokio::sync::mpsc; +use tokio::task::JoinSet; use tokio_util::sync::CancellationToken; -type TorrentCallback = Arc; -type MetadataFetchCallback = Arc< +const MAX_PEERS_PER_HASH: usize = 3; +const HASH_QUEUE_TTL: Duration = Duration::from_secs(60); + +/// Callback returning whether the application accepted a downloaded torrent. +pub type TorrentAckCallback = Box bool + Send + Sync + 'static>; +/// Callback invoked once when an admitted Metadata job reaches a terminal state. +pub type MetadataCompletionCallback = Box; +/// Asynchronous InfoHash admission callback. +pub type MetadataFetchCallback = Box< dyn Fn(String) -> std::pin::Pin + Send>> + Send - + Sync, + + Sync + + 'static, >; +#[derive(Debug, Clone, Copy)] +/// Pending-queue capacity and maximum concurrent Metadata jobs. +pub struct MetadataSchedulerLimits { + /// Maximum deduplicated pending InfoHashes. + pub queue_size: usize, + /// Maximum concurrently spawned jobs. + pub concurrency: usize, +} + +#[derive(Clone)] +/// Atomically replaceable callbacks shared with [`MetadataScheduler`]. +pub struct MetadataSchedulerCallbacks { + /// Torrent delivery callback. + pub torrent: Arc>, + /// Pre-download admission callback. + pub fetch_gate: Arc>, + /// Terminal completion callback. + pub completion: Arc>, +} + +#[derive(Debug, Clone, Copy)] +struct PeerCandidate { + addr: SocketAddr, + discovered_at: Instant, +} + +#[derive(Debug)] +struct QueuedHash { + info_hash: String, + peers: Vec, + latest_at: Instant, + order_key: (Instant, u64), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum QueuePushKind { + Inserted, + Updated, + EvictedOldest, + Stale, +} + +#[derive(Debug)] +struct PendingHashQueue { + capacity: usize, + ttl: Duration, + entries: HashMap, + order: BTreeMap<(Instant, u64), String>, + sequence: u64, +} + +impl PendingHashQueue { + fn new(capacity: usize, ttl: Duration) -> Self { + Self { + capacity: capacity.max(1), + ttl, + entries: HashMap::new(), + order: BTreeMap::new(), + sequence: 0, + } + } + + fn len(&self) -> usize { + self.entries.len() + } + + fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + #[cfg(test)] + fn contains(&self, info_hash: &str) -> bool { + self.entries.contains_key(info_hash) + } + + fn next_order_key(&mut self, at: Instant) -> (Instant, u64) { + self.sequence = self.sequence.wrapping_add(1); + (at, self.sequence) + } + + fn push(&mut self, event: HashDiscovered, now: Instant) -> QueuePushKind { + if now + .checked_duration_since(event.discovered_at) + .unwrap_or_default() + > self.ttl + { + return QueuePushKind::Stale; + } + + if let Some(mut entry) = self.entries.remove(&event.info_hash) { + self.order.remove(&entry.order_key); + let discovered_at = entry + .peers + .iter() + .find(|peer| peer.addr == event.peer_addr) + .map(|peer| peer.discovered_at.max(event.discovered_at)) + .unwrap_or(event.discovered_at); + entry.peers.retain(|peer| peer.addr != event.peer_addr); + entry.peers.push(PeerCandidate { + addr: event.peer_addr, + discovered_at, + }); + entry + .peers + .sort_unstable_by(|left, right| right.discovered_at.cmp(&left.discovered_at)); + entry.peers.truncate(MAX_PEERS_PER_HASH); + entry.latest_at = entry.latest_at.max(event.discovered_at); + entry.order_key = self.next_order_key(entry.latest_at); + self.order.insert(entry.order_key, entry.info_hash.clone()); + self.entries.insert(entry.info_hash.clone(), entry); + return QueuePushKind::Updated; + } + + let mut result = QueuePushKind::Inserted; + if self.entries.len() >= self.capacity { + let Some((&oldest_key, oldest_hash)) = self.order.first_key_value() else { + return QueuePushKind::Stale; + }; + if event.discovered_at <= oldest_key.0 { + return QueuePushKind::Stale; + } + let oldest_hash = oldest_hash.clone(); + self.order.remove(&oldest_key); + self.entries.remove(&oldest_hash); + result = QueuePushKind::EvictedOldest; + } + + let order_key = self.next_order_key(event.discovered_at); + let info_hash = event.info_hash; + self.order.insert(order_key, info_hash.clone()); + self.entries.insert( + info_hash.clone(), + QueuedHash { + info_hash, + peers: vec![PeerCandidate { + addr: event.peer_addr, + discovered_at: event.discovered_at, + }], + latest_at: event.discovered_at, + order_key, + }, + ); + result + } + + fn remove(&mut self, info_hash: &str) -> Option { + let entry = self.entries.remove(info_hash)?; + self.order.remove(&entry.order_key); + Some(entry) + } + + fn pop_newest_available( + &mut self, + in_flight: &HashMap, + ) -> Option { + let info_hash = self + .order + .iter() + .rev() + .find_map(|(_, hash)| (!in_flight.contains_key(hash)).then(|| hash.clone()))?; + self.remove(&info_hash) + } + + fn expire(&mut self, now: Instant) -> usize { + let mut expired = 0; + loop { + let Some((&oldest_key, oldest_hash)) = self.order.first_key_value() else { + break; + }; + if now.checked_duration_since(oldest_key.0).unwrap_or_default() <= self.ttl { + break; + } + let oldest_hash = oldest_hash.clone(); + self.order.remove(&oldest_key); + self.entries.remove(&oldest_hash); + expired += 1; + } + expired + } +} + +#[derive(Debug)] +struct MetadataJob { + info_hash: String, + peers: Vec, + already_admitted: bool, +} + +#[derive(Debug)] +struct InFlightState { + attempts: usize, + scheduled_peers: HashSet, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum JobOutcome { + GateRejected, + Completed(MetadataFetchCompletionStatus), +} + +#[derive(Debug)] +struct JobResult { + info_hash: String, + attempts: usize, + outcome: JobOutcome, +} + +/// Bounded, deduplicating, freshness-aware Metadata job scheduler. pub struct MetadataScheduler { hash_rx: mpsc::Receiver, max_queue_size: usize, max_concurrent: usize, fetcher: Arc, - callback: Arc>>, - on_metadata_fetch: Arc>>, + callback: Arc>, + on_metadata_fetch: Arc>, + completion_callback: Arc>, total_received: Arc, total_dropped: Arc, total_dispatched: Arc, + total_completed: Arc, queue_len: Arc, + runtime_stats: DhtRuntimeStats, shutdown: CancellationToken, } impl MetadataScheduler { - #[allow(clippy::too_many_arguments)] + /// Creates a scheduler using local runtime statistics. pub fn new( hash_rx: mpsc::Receiver, fetcher: Arc, - max_queue_size: usize, - max_concurrent: usize, - callback: Arc>>, - on_metadata_fetch: Arc>>, + limits: MetadataSchedulerLimits, + callbacks: MetadataSchedulerCallbacks, queue_len: Arc, shutdown: CancellationToken, + ) -> Self { + Self::new_with_runtime_stats( + hash_rx, + fetcher, + limits, + callbacks, + queue_len, + shutdown, + DhtRuntimeStats::default(), + ) + } + + pub(crate) fn new_with_runtime_stats( + hash_rx: mpsc::Receiver, + fetcher: Arc, + limits: MetadataSchedulerLimits, + callbacks: MetadataSchedulerCallbacks, + queue_len: Arc, + shutdown: CancellationToken, + runtime_stats: DhtRuntimeStats, ) -> Self { Self { hash_rx, - max_queue_size, - max_concurrent, + max_queue_size: limits.queue_size.max(1), + max_concurrent: limits.concurrency.max(1), fetcher, - callback, - on_metadata_fetch, + callback: callbacks.torrent, + on_metadata_fetch: callbacks.fetch_gate, + completion_callback: callbacks.completion, total_received: Arc::new(AtomicU64::new(0)), total_dropped: Arc::new(AtomicU64::new(0)), total_dispatched: Arc::new(AtomicU64::new(0)), + total_completed: Arc::new(AtomicU64::new(0)), queue_len, + runtime_stats, shutdown, } } - pub fn set_callback(&mut self, callback: TorrentCallback) { - if let Ok(mut guard) = self.callback.try_write() { - *guard = Some(callback); - } + /// Replaces the torrent delivery callback. + pub fn set_callback(&mut self, callback: Arc) { + self.callback.store(Some(callback)); } - pub fn set_metadata_fetch_callback(&mut self, callback: MetadataFetchCallback) { - if let Ok(mut guard) = self.on_metadata_fetch.try_write() { - *guard = Some(callback); - } + /// Replaces the pre-download admission callback. + pub fn set_metadata_fetch_callback(&mut self, callback: Arc) { + self.on_metadata_fetch.store(Some(callback)); } + /// Replaces the terminal completion callback. + pub fn set_completion_callback(&mut self, callback: Arc) { + self.completion_callback.store(Some(callback)); + } + + /// Runs until cancellation and drains completed jobs before returning when input closes. pub async fn run(mut self) { - let (task_tx, task_rx) = async_channel::bounded::(self.max_queue_size); + let mut queue = PendingHashQueue::new(self.max_queue_size, HASH_QUEUE_TTL); + let mut in_flight = HashMap::::new(); + let mut tasks = JoinSet::::new(); + let mut maintenance = tokio::time::interval(Duration::from_secs(1)); + let mut stats_interval = tokio::time::interval(Duration::from_secs(60)); + let mut input_closed = false; - let shutdown = self.shutdown.clone(); - #[cfg_attr(not(debug_assertions), allow(unused_variables))] - for worker_id in 0..self.max_concurrent { - let task_rx = task_rx.clone(); - let fetcher = self.fetcher.clone(); - let callback = self.callback.clone(); - let on_metadata_fetch = self.on_metadata_fetch.clone(); - let total_dispatched = self.total_dispatched.clone(); - let queue_len = self.queue_len.clone(); - let shutdown_worker = shutdown.clone(); - - tokio::spawn(async move { - #[cfg(debug_assertions)] - log::trace!("Worker {} 启动", worker_id); - - loop { - tokio::select! { - _ = shutdown_worker.cancelled() => { - #[cfg(debug_assertions)] - log::trace!("Worker {} 收到关闭信号,退出", worker_id); - break; - } - result = task_rx.recv() => { - let hash = match result { - Ok(h) => { - queue_len.fetch_sub(1, Ordering::Relaxed); - h - } - Err(_) => break, - }; - - total_dispatched.fetch_add(1, Ordering::Relaxed); - - Self::process_hash( - hash, - &fetcher, - &callback, - &on_metadata_fetch, - ).await; - } - } - } - - #[cfg(debug_assertions)] - log::trace!("Worker {} 退出", worker_id); - }); - } - - let mut stats_interval = if cfg!(debug_assertions) { - Some(tokio::time::interval(std::time::Duration::from_secs(60))) - } else { - None - }; - if let Some(ref mut interval) = stats_interval { - interval.tick().await; - } - - let shutdown = self.shutdown.clone(); loop { - tokio::select! { - _ = shutdown.cancelled() => { - #[cfg(debug_assertions)] - log::trace!("MetadataScheduler 主循环收到关闭信号,退出"); - break; - } - result = self.hash_rx.recv() => { - match result { - Some(hash) => { - self.total_received.fetch_add(1, Ordering::Relaxed); + self.dispatch_ready(&mut queue, &mut in_flight, &mut tasks); + self.sync_queue_len(queue.len(), in_flight.len()); - match task_tx.try_send(hash) { - Ok(_) => { - self.queue_len.fetch_add(1, Ordering::Relaxed); - } - Err(async_channel::TrySendError::Full(_)) => { - self.total_dropped.fetch_add(1, Ordering::Relaxed); - } - Err(_) => break, - } + if input_closed && queue.is_empty() && tasks.is_empty() { + break; + } + + tokio::select! { + _ = self.shutdown.cancelled() => break, + result = self.hash_rx.recv(), if !input_closed => { + self.runtime_stats + .set_hash_ingress_queue_depth(self.hash_rx.len()); + match result { + Some(hash) => self.enqueue(&mut queue, hash), + None => input_closed = true, + } + } + result = tasks.join_next(), if !tasks.is_empty() => { + match result { + Some(Ok(result)) => { + self.handle_job_result(result, &mut queue, &mut in_flight, &mut tasks); } - None => break, + Some(Err(error)) => { + #[cfg(feature = "metrics")] + counter!("dht_metadata_worker_join_error_total").increment(1); + log::warn!("Metadata worker task failed: {error}"); + } + None => {} } } - _ = async { - match stats_interval.as_mut() { - Some(interval) => interval.tick().await, - None => std::future::pending().await, + _ = maintenance.tick() => { + let expired = queue.expire(Instant::now()); + if expired > 0 { + self.total_dropped.fetch_add(expired as u64, Ordering::Relaxed); + self.runtime_stats.metadata_queue_stale(expired); + #[cfg(feature = "metrics")] + counter!("dht_metadata_queue_events_total", "result" => "expired") + .increment(expired as u64); } - } => { - self.print_stats_inline(); } + _ = stats_interval.tick() => self.print_stats_inline(), } } - drop(task_tx); - #[cfg(debug_assertions)] - log::trace!("MetadataScheduler 主循环退出,等待 worker 任务完成"); + tasks.abort_all(); + while tasks.join_next().await.is_some() {} + self.queue_len.store(0, Ordering::Relaxed); + self.runtime_stats.set_hash_ingress_queue_depth(0); + self.runtime_stats.set_metadata_queue(0, 0); + #[cfg(feature = "metrics")] + gauge!("dht_metadata_queue_depth").set(0.0); } - async fn process_hash( - hash: HashDiscovered, - fetcher: &Arc, - callback: &Arc>>, - on_metadata_fetch: &Arc>>, - ) { - let info_hash = hash.info_hash.clone(); - let peer_addr = hash.peer_addr; - - let maybe_check_fn = { - match on_metadata_fetch.read() { - Ok(guard) => guard.clone(), - Err(_) => return, + fn enqueue(&self, queue: &mut PendingHashQueue, hash: HashDiscovered) { + self.total_received.fetch_add(1, Ordering::Relaxed); + let kind = queue.push(hash, Instant::now()); + let result = match kind { + QueuePushKind::Inserted => { + self.runtime_stats.metadata_queue_inserted(); + "inserted" + } + QueuePushKind::Updated => { + self.runtime_stats.metadata_queue_deduplicated(); + "deduplicated" + } + QueuePushKind::EvictedOldest => { + self.total_dropped.fetch_add(1, Ordering::Relaxed); + self.runtime_stats.metadata_queue_inserted(); + self.runtime_stats.metadata_queue_evicted(); + "evicted_oldest" + } + QueuePushKind::Stale => { + self.total_dropped.fetch_add(1, Ordering::Relaxed); + self.runtime_stats.metadata_queue_stale(1); + "stale" } }; + #[cfg(feature = "metrics")] + counter!("dht_metadata_queue_events_total", "result" => result).increment(1); + #[cfg(not(feature = "metrics"))] + let _ = result; + } - if let Some(f) = maybe_check_fn - && !f(info_hash.clone()).await - { + fn dispatch_ready( + &self, + queue: &mut PendingHashQueue, + in_flight: &mut HashMap, + tasks: &mut JoinSet, + ) { + while tasks.len() < self.max_concurrent { + let Some(entry) = queue.pop_newest_available(in_flight) else { + break; + }; + let queue_wait = Instant::now() + .checked_duration_since(entry.latest_at) + .unwrap_or_default(); + self.runtime_stats.observe_metadata_queue_wait( + queue_wait.as_millis().min(u128::from(u64::MAX)) as u64, + ); + #[cfg(feature = "metrics")] + histogram!("dht_metadata_queue_wait_seconds").record(queue_wait.as_secs_f64()); + let state = InFlightState { + attempts: 0, + scheduled_peers: HashSet::new(), + }; + in_flight.insert(entry.info_hash.clone(), state); + self.spawn_job(entry, false, in_flight, tasks); + } + } + + fn spawn_job( + &self, + mut entry: QueuedHash, + already_admitted: bool, + in_flight: &mut HashMap, + tasks: &mut JoinSet, + ) { + let Some(state) = in_flight.get_mut(&entry.info_hash) else { + return; + }; + let remaining = MAX_PEERS_PER_HASH.saturating_sub(state.attempts); + let now = Instant::now(); + entry.peers.retain(|peer| { + !state.scheduled_peers.contains(&peer.addr) + && now + .checked_duration_since(peer.discovered_at) + .unwrap_or_default() + <= HASH_QUEUE_TTL + }); + entry.peers.truncate(remaining); + for peer in &entry.peers { + state.scheduled_peers.insert(peer.addr); + } + + if entry.peers.is_empty() { + let attempts = state.attempts; + if already_admitted { + self.finish( + entry.info_hash.clone(), + MetadataFetchCompletionStatus::FetchFailed, + attempts, + ); + } else { + self.total_dropped.fetch_add(1, Ordering::Relaxed); + self.runtime_stats.metadata_queue_stale(1); + #[cfg(feature = "metrics")] + counter!("dht_metadata_queue_events_total", "result" => "expired").increment(1); + } + in_flight.remove(&entry.info_hash); return; } - let info_hash_bytes: [u8; 20] = match hex::decode(&info_hash) { + self.total_dispatched.fetch_add(1, Ordering::Relaxed); + #[cfg(feature = "metrics")] + counter!("dht_metadata_jobs_dispatched_total").increment(1); + + let job = MetadataJob { + info_hash: entry.info_hash, + peers: entry.peers, + already_admitted, + }; + let fetcher = self.fetcher.clone(); + let callback = self.callback.clone(); + let on_metadata_fetch = self.on_metadata_fetch.clone(); + tasks.spawn(async move { + Self::process_hash(job, &fetcher, &callback, &on_metadata_fetch).await + }); + } + + fn handle_job_result( + &self, + result: JobResult, + queue: &mut PendingHashQueue, + in_flight: &mut HashMap, + tasks: &mut JoinSet, + ) { + let Some(state) = in_flight.get_mut(&result.info_hash) else { + return; + }; + state.attempts = state.attempts.saturating_add(result.attempts); + + match result.outcome { + JobOutcome::GateRejected => { + queue.remove(&result.info_hash); + in_flight.remove(&result.info_hash); + #[cfg(feature = "metrics")] + counter!("dht_metadata_jobs_completed_total", "result" => "gate_rejected") + .increment(1); + } + JobOutcome::Completed(MetadataFetchCompletionStatus::FetchFailed) + if state.attempts < MAX_PEERS_PER_HASH => + { + if let Some(entry) = queue.remove(&result.info_hash) { + self.spawn_job(entry, true, in_flight, tasks); + } else { + let attempts = state.attempts; + in_flight.remove(&result.info_hash); + self.finish( + result.info_hash, + MetadataFetchCompletionStatus::FetchFailed, + attempts, + ); + } + } + JobOutcome::Completed(status) => { + let attempts = state.attempts; + queue.remove(&result.info_hash); + in_flight.remove(&result.info_hash); + self.finish(result.info_hash, status, attempts); + } + } + } + + fn finish(&self, info_hash: String, status: MetadataFetchCompletionStatus, attempts: usize) { + self.total_completed.fetch_add(1, Ordering::Relaxed); + let result = match status { + MetadataFetchCompletionStatus::Accepted => "accepted", + MetadataFetchCompletionStatus::FetchFailed => "fetch_failed", + MetadataFetchCompletionStatus::DeliveryRejected => "delivery_rejected", + }; + #[cfg(feature = "metrics")] + counter!("dht_metadata_jobs_completed_total", "result" => result).increment(1); + #[cfg(not(feature = "metrics"))] + let _ = result; + + if let Some(callback) = self.completion_callback.load_full() { + let completion = MetadataFetchCompletion { + info_hash, + status, + attempts, + }; + if catch_unwind(AssertUnwindSafe(|| callback(completion))).is_err() { + #[cfg(feature = "metrics")] + counter!("dht_metadata_completion_callback_panics_total").increment(1); + log::warn!("Metadata completion callback panicked"); + } + } + } + + async fn process_hash( + job: MetadataJob, + fetcher: &Arc, + callback: &Arc>, + on_metadata_fetch: &Arc>, + ) -> JobResult { + if !job.already_admitted + && let Some(gate) = on_metadata_fetch.load_full() + && !gate(job.info_hash.clone()).await + { + return JobResult { + info_hash: job.info_hash, + attempts: 0, + outcome: JobOutcome::GateRejected, + }; + } + + let info_hash_bytes: [u8; 20] = match hex::decode(&job.info_hash) { Ok(bytes) if bytes.len() == 20 => { - let mut arr = [0u8; 20]; - arr.copy_from_slice(&bytes); - arr + let mut hash = [0u8; 20]; + hash.copy_from_slice(&bytes); + hash + } + _ => { + return JobResult { + info_hash: job.info_hash, + attempts: 0, + outcome: JobOutcome::Completed(MetadataFetchCompletionStatus::FetchFailed), + }; } - _ => return, }; - if let Some((name, total_size, files, piece_length)) = - fetcher.fetch(&info_hash_bytes, peer_addr).await - { + let mut attempts = 0; + for peer in job.peers { + let (name, total_size, files, piece_length) = + match fetcher.fetch(&info_hash_bytes, peer.addr).await { + MetadataFetchOutcome::Fetched(metadata) => { + attempts += 1; + metadata + } + MetadataFetchOutcome::Failed => { + attempts += 1; + continue; + } + MetadataFetchOutcome::SkippedCached => continue, + }; + let metadata = TorrentInfo { - info_hash, + info_hash: job.info_hash.clone(), name, total_size, files, - magnet_link: format!("magnet:?xt=urn:btih:{}", hash.info_hash), - peers: vec![peer_addr.to_string()], + magnet_link: format!("magnet:?xt=urn:btih:{}", job.info_hash), + peers: vec![peer.addr.to_string()], piece_length, timestamp: std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -216,16 +630,36 @@ impl MetadataScheduler { .as_secs(), }; - let maybe_torrent_cb = { - match callback.read() { - Ok(guard) => guard.clone(), - Err(_) => return, - } + let accepted = callback.load_full().is_some_and(|callback| { + catch_unwind(AssertUnwindSafe(|| callback(metadata))).unwrap_or(false) + }); + let status = if accepted { + MetadataFetchCompletionStatus::Accepted + } else { + MetadataFetchCompletionStatus::DeliveryRejected }; + return JobResult { + info_hash: job.info_hash, + attempts, + outcome: JobOutcome::Completed(status), + }; + } - if let Some(cb) = maybe_torrent_cb { - cb(metadata); - } + JobResult { + info_hash: job.info_hash, + attempts, + outcome: JobOutcome::Completed(MetadataFetchCompletionStatus::FetchFailed), + } + } + + fn sync_queue_len(&self, queued: usize, in_flight: usize) { + let depth = queued.saturating_add(in_flight); + self.queue_len.store(depth, Ordering::Relaxed); + self.runtime_stats.set_metadata_queue(depth, in_flight); + #[cfg(feature = "metrics")] + { + gauge!("dht_metadata_queue_depth").set(depth as f64); + gauge!("dht_metadata_in_flight").set(in_flight as f64); } } @@ -235,39 +669,283 @@ impl MetadataScheduler { let received = self.total_received.load(Ordering::Relaxed); let dropped = self.total_dropped.load(Ordering::Relaxed); let dispatched = self.total_dispatched.load(Ordering::Relaxed); - - let drop_rate = if received > 0 { - dropped as f64 / received as f64 * 100.0 - } else { - 0.0 - }; - + let completed = self.total_completed.load(Ordering::Relaxed); let queue_len = self.queue_len.load(Ordering::Relaxed); - let queue_pressure = (queue_len as f64 / self.max_queue_size as f64) * 100.0; - - if queue_pressure > 80.0 { - log::warn!( - "Metadata 队列高压:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", - queue_len, - self.max_queue_size, - queue_pressure, - received, - dispatched, - dropped, - drop_rate - ); - } else { - log::info!( - "Metadata 调度器统计:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", - queue_len, - self.max_queue_size, - queue_pressure, - received, - dispatched, - dropped, - drop_rate - ); - } + let queue_pressure = queue_len as f64 / self.max_queue_size as f64 * 100.0; + log::info!( + "Metadata scheduler: depth={}/{}, pressure={:.1}%, received={}, dispatched={}, completed={}, evicted_or_expired={}", + queue_len, + self.max_queue_size, + queue_pressure, + received, + dispatched, + completed, + dropped, + ); } } } + +#[cfg(test)] +mod tests { + use super::*; + + fn empty_torrent_callback() -> Arc> { + Arc::new(ArcSwapOption::empty()) + } + + fn fetch_gate(callback: F) -> Arc> + where + F: Fn(String) -> Fut + Send + Sync + 'static, + Fut: std::future::Future + Send + 'static, + { + let holder = Arc::new(ArcSwapOption::empty()); + let callback: Arc = + Arc::new(Box::new(move |hash| Box::pin(callback(hash)))); + holder.store(Some(callback)); + holder + } + + fn completion_callback(callback: F) -> Arc> + where + F: Fn(MetadataFetchCompletion) + Send + Sync + 'static, + { + let holder = Arc::new(ArcSwapOption::empty()); + let callback: Arc = Arc::new(Box::new(callback)); + holder.store(Some(callback)); + holder + } + + fn event(hash: &str, port: u16, discovered_at: Instant) -> HashDiscovered { + HashDiscovered { + info_hash: hash.to_string(), + peer_addr: SocketAddr::from(([127, 0, 0, 1], port)), + discovered_at, + } + } + + #[test] + fn runtime_stats_track_queue_outcomes_and_depth() { + let (_hash_tx, hash_rx) = mpsc::channel(4); + let stats = DhtRuntimeStats::with_limits(DhtRuntimeLimits { + metadata_queue: 2, + ..DhtRuntimeLimits::default() + }); + let scheduler = MetadataScheduler::new_with_runtime_stats( + hash_rx, + Arc::new(RbitFetcher::new(1)), + MetadataSchedulerLimits { + queue_size: 2, + concurrency: 1, + }, + MetadataSchedulerCallbacks { + torrent: empty_torrent_callback(), + fetch_gate: Arc::new(ArcSwapOption::empty()), + completion: Arc::new(ArcSwapOption::empty()), + }, + Arc::new(AtomicUsize::new(0)), + CancellationToken::new(), + stats.clone(), + ); + let start = Instant::now(); + let mut queue = PendingHashQueue::new(2, HASH_QUEUE_TTL); + + scheduler.enqueue(&mut queue, event("old", 1000, start)); + scheduler.enqueue( + &mut queue, + event("old", 1001, start + Duration::from_secs(1)), + ); + scheduler.enqueue( + &mut queue, + event("middle", 1002, start + Duration::from_secs(2)), + ); + scheduler.enqueue( + &mut queue, + event("new", 1003, start + Duration::from_secs(3)), + ); + scheduler.enqueue( + &mut queue, + event( + "stale", + 1004, + start + .checked_sub(HASH_QUEUE_TTL + Duration::from_secs(1)) + .unwrap(), + ), + ); + scheduler.sync_queue_len(queue.len(), 1); + + let snapshot = stats.snapshot(); + assert_eq!(snapshot.metadata_queue_inserted, 3); + assert_eq!(snapshot.metadata_queue_deduplicated, 1); + assert_eq!(snapshot.metadata_queue_evicted, 1); + assert_eq!(snapshot.metadata_queue_stale, 1); + assert_eq!(snapshot.metadata_queue_depth, 3); + assert_eq!(snapshot.metadata_queue_max, 2); + assert_eq!(snapshot.metadata_in_flight, 1); + } + + #[test] + fn queue_deduplicates_hash_and_keeps_newest_unique_peers() { + let start = Instant::now(); + let mut queue = PendingHashQueue::new(10, Duration::from_secs(60)); + for offset in 0..4 { + assert_ne!( + queue.push( + event( + "hash", + 1000 + offset, + start + Duration::from_secs(offset as u64) + ), + start + Duration::from_secs(offset as u64), + ), + QueuePushKind::Stale + ); + } + queue.push( + event("hash", 1003, start + Duration::from_secs(5)), + start + Duration::from_secs(5), + ); + + assert_eq!(queue.len(), 1); + let entry = queue.remove("hash").unwrap(); + let ports: Vec<_> = entry.peers.iter().map(|peer| peer.addr.port()).collect(); + assert_eq!(ports, vec![1003, 1002, 1001]); + } + + #[test] + fn full_queue_evicts_oldest_for_newer_hash() { + let start = Instant::now(); + let mut queue = PendingHashQueue::new(2, Duration::from_secs(60)); + queue.push(event("old", 1000, start), start); + queue.push( + event("middle", 1001, start + Duration::from_secs(1)), + start + Duration::from_secs(1), + ); + let result = queue.push( + event("new", 1002, start + Duration::from_secs(2)), + start + Duration::from_secs(2), + ); + + assert_eq!(result, QueuePushKind::EvictedOldest); + assert!(!queue.contains("old")); + assert!(queue.contains("middle")); + assert!(queue.contains("new")); + } + + #[test] + fn queue_expires_stale_hashes_and_pops_newest_first() { + let start = Instant::now(); + let mut queue = PendingHashQueue::new(4, Duration::from_secs(10)); + queue.push(event("older", 1000, start), start); + queue.push( + event("newer", 1001, start + Duration::from_secs(1)), + start + Duration::from_secs(1), + ); + + let in_flight = HashMap::new(); + let newest = queue.pop_newest_available(&in_flight).unwrap(); + assert_eq!(newest.info_hash, "newer"); + assert_eq!(queue.expire(start + Duration::from_secs(11)), 1); + assert!(queue.is_empty()); + } + + #[tokio::test] + async fn gate_rejection_does_not_emit_completion() { + let (hash_tx, hash_rx) = mpsc::channel(4); + let gate_calls = Arc::new(AtomicUsize::new(0)); + let gate_calls_for_callback = gate_calls.clone(); + let gate = fetch_gate(move |_| { + gate_calls_for_callback.fetch_add(1, Ordering::Relaxed); + async { false } + }); + let (completion_tx, mut completion_rx) = mpsc::unbounded_channel(); + let completion = completion_callback(move |result| { + let _ = completion_tx.send(result); + }); + let shutdown = CancellationToken::new(); + let scheduler = MetadataScheduler::new( + hash_rx, + Arc::new(RbitFetcher::new(1)), + MetadataSchedulerLimits { + queue_size: 4, + concurrency: 1, + }, + MetadataSchedulerCallbacks { + torrent: empty_torrent_callback(), + fetch_gate: gate, + completion, + }, + Arc::new(AtomicUsize::new(0)), + shutdown.clone(), + ); + let scheduler_task = tokio::spawn(scheduler.run()); + hash_tx + .send(event("bad", 1000, Instant::now())) + .await + .unwrap(); + + tokio::time::timeout(Duration::from_secs(1), async { + while gate_calls.load(Ordering::Relaxed) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(50), completion_rx.recv()) + .await + .is_err() + ); + + shutdown.cancel(); + scheduler_task.await.unwrap(); + } + + #[tokio::test] + async fn admitted_failure_emits_one_completion() { + let (hash_tx, hash_rx) = mpsc::channel(4); + let gate = fetch_gate(|_| async { true }); + let (completion_tx, mut completion_rx) = mpsc::unbounded_channel(); + let completion = completion_callback(move |result| { + let _ = completion_tx.send(result); + }); + let shutdown = CancellationToken::new(); + let scheduler = MetadataScheduler::new( + hash_rx, + Arc::new(RbitFetcher::new(1)), + MetadataSchedulerLimits { + queue_size: 4, + concurrency: 1, + }, + MetadataSchedulerCallbacks { + torrent: empty_torrent_callback(), + fetch_gate: gate, + completion, + }, + Arc::new(AtomicUsize::new(0)), + shutdown.clone(), + ); + let scheduler_task = tokio::spawn(scheduler.run()); + hash_tx + .send(event("bad", 1000, Instant::now())) + .await + .unwrap(); + + let result = tokio::time::timeout(Duration::from_secs(1), completion_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(result.info_hash, "bad"); + assert_eq!(result.status, MetadataFetchCompletionStatus::FetchFailed); + assert_eq!(result.attempts, 0); + assert!( + tokio::time::timeout(Duration::from_millis(50), completion_rx.recv()) + .await + .is_err() + ); + + shutdown.cancel(); + scheduler_task.await.unwrap(); + } +} diff --git a/src/server.rs b/src/server.rs index 03d1659..d52f354 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,62 +1,261 @@ +use crate::addr::is_valid_node_addr; +use crate::budget::RateBucket; +use crate::crawl_config::ResolvedCrawlConfig; +use crate::crawl_engine::CrawlEngine; use crate::error::Result; +use crate::krpc::encode_response; use crate::metadata::RbitFetcher; -use crate::protocol::{DhtArgs, DhtMessage, DhtResponse}; -use crate::scheduler::MetadataScheduler; -use crate::sharded::{NodeTuple, ShardedNodeQueue}; -use crate::types::{DHTOptions, NetMode, TorrentInfo}; +use crate::node_id::{neighbor_node_id, random_node_id, transaction_id_from_bytes}; +use crate::protocol::{DhtArgs, DhtMessage}; +use crate::runtime_stats::{DhtRuntimeLimits, DhtRuntimeStats}; +use crate::scheduler::{ + MetadataCompletionCallback, MetadataFetchCallback, MetadataScheduler, + MetadataSchedulerCallbacks, MetadataSchedulerLimits, TorrentAckCallback, +}; +use crate::types::{DHTOptions, MetadataFetchCompletion, NetMode, NodeTuple, TorrentInfo}; +use crate::udp_buffer::UdpBufferPool; +use crate::udp_ingress::{WorkerHandle, spawn_udp_listener}; +use ahash::AHashMap; +use arc_swap::ArcSwapOption; +use bytes::BytesMut; #[cfg(feature = "metrics")] -use metrics::{counter, gauge}; +use metrics::counter; use rand::Rng; use socket2::{Domain, Protocol, Socket, Type}; -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; use std::future::Future; use std::hash::{Hash, Hasher}; use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; -use std::pin::Pin; +use std::sync::Arc; use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::{Arc, RwLock}; -use std::time::Duration; +use std::time::{Duration, Instant}; use tokio::net::UdpSocket; -use tokio::sync::{Semaphore, mpsc}; +use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; -const BOOTSTRAP_NODES: &[&str] = &[ - "router.bittorrent.com:6881", - "dht.transmissionbt.com:6881", - "router.utorrent.com:6881", - "dht.aelitis.com:6881", -]; +type FilterCallback = Box bool + Send + Sync + 'static>; +type ErrorCallback = Box; -pub type BoxedBoolFuture = Pin + Send>>; -pub type MetadataFetchCallback = Arc BoxedBoolFuture + Send + Sync>; -type WorkerHandle = mpsc::Sender<(Box<[u8]>, SocketAddr, SocketAddr)>; +struct QueryResponse<'a> { + transaction_id: &'a [u8], + remote_addr: SocketAddr, + local_addr: SocketAddr, + query_type: &'a str, + sender_id: Option<&'a [u8]>, + target_id: Option<&'a [u8]>, +} + +#[derive(Clone, Copy)] +struct SourceResponseWindow { + started_at: Instant, + last_seen: Instant, + count: u32, +} + +struct WorkerResponseLimiter { + regular_packets: RateBucket, + regular_bytes: RateBucket, + priority_packets: RateBucket, + priority_bytes: RateBucket, + per_source_rate: u32, + sources: AHashMap, + source_expiry: VecDeque<(Instant, SocketAddr)>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ResponsePermit { + Regular, + PriorityReserve, + Rejected, +} + +fn reserve_quota(total: u32) -> u32 { + if total == 0 { + 0 + } else { + total.div_ceil(10).min(total) + } +} + +impl WorkerResponseLimiter { + fn new(packet_rate: u32, byte_rate: u64, per_source_rate: u32, now: Instant) -> Self { + let byte_rate = byte_rate.min(u32::MAX as u64) as u32; + let priority_packet_rate = reserve_quota(packet_rate); + let priority_byte_rate = reserve_quota(byte_rate); + let packet_rate = packet_rate.saturating_sub(priority_packet_rate); + let byte_rate = byte_rate.saturating_sub(priority_byte_rate); + Self { + regular_packets: RateBucket::per_second( + packet_rate, + packet_rate.div_ceil(5).max(1), + true, + now, + ), + regular_bytes: RateBucket::per_second( + byte_rate, + byte_rate.div_ceil(5).max(512), + true, + now, + ), + priority_packets: RateBucket::per_second( + priority_packet_rate, + priority_packet_rate.div_ceil(5).max(1), + true, + now, + ), + priority_bytes: RateBucket::per_second( + priority_byte_rate, + priority_byte_rate.div_ceil(5).max(512), + true, + now, + ), + per_source_rate, + sources: AHashMap::new(), + source_expiry: VecDeque::new(), + } + } + + fn acquire( + &mut self, + addr: SocketAddr, + encoded_len: usize, + is_priority: bool, + now: Instant, + ) -> ResponsePermit { + self.expire_sources(now); + if Self::take_budget( + &mut self.regular_packets, + &mut self.regular_bytes, + encoded_len, + now, + ) { + if self.acquire_source_slot(addr, now) { + return ResponsePermit::Regular; + } + Self::refund_budget( + &mut self.regular_packets, + &mut self.regular_bytes, + encoded_len, + ); + return ResponsePermit::Rejected; + } + + if is_priority + && Self::take_budget( + &mut self.priority_packets, + &mut self.priority_bytes, + encoded_len, + now, + ) + { + if self.acquire_source_slot(addr, now) { + return ResponsePermit::PriorityReserve; + } + Self::refund_budget( + &mut self.priority_packets, + &mut self.priority_bytes, + encoded_len, + ); + } + + ResponsePermit::Rejected + } + + fn take_budget( + packets: &mut RateBucket, + bytes: &mut RateBucket, + encoded_len: usize, + now: Instant, + ) -> bool { + if !packets.try_take_one(now) { + return false; + } + if !bytes.try_take_exact(encoded_len, now) { + packets.refund_one(); + return false; + } + true + } + + fn refund_budget(packets: &mut RateBucket, bytes: &mut RateBucket, encoded_len: usize) { + packets.refund_one(); + bytes.refund(encoded_len); + } + + fn acquire_source_slot(&mut self, addr: SocketAddr, now: Instant) -> bool { + if self.per_source_rate == 0 { + return false; + } + + let entry = self.sources.entry(addr).or_insert(SourceResponseWindow { + started_at: now, + last_seen: now, + count: 0, + }); + if now + .checked_duration_since(entry.started_at) + .unwrap_or_default() + >= Duration::from_secs(1) + { + entry.started_at = now; + entry.count = 0; + } + if entry.count >= self.per_source_rate { + return false; + } + entry.count += 1; + entry.last_seen = now; + self.source_expiry + .push_back((now + Duration::from_secs(60), addr)); + true + } + + fn expire_sources(&mut self, now: Instant) { + while let Some((deadline, addr)) = self.source_expiry.front().copied() { + if deadline > now { + break; + } + self.source_expiry.pop_front(); + if self.sources.get(&addr).is_some_and(|entry| { + now.checked_duration_since(entry.last_seen) + .unwrap_or_default() + >= Duration::from_secs(60) + }) { + self.sources.remove(&addr); + } + } + } +} #[derive(Debug, Clone)] +/// InfoHash and announcing Peer submitted to the Metadata scheduler. pub struct HashDiscovered { + /// Lowercase hexadecimal InfoHash. pub info_hash: String, + /// Peer endpoint derived from announce `port`/`implied_port`. pub peer_addr: SocketAddr, + /// Monotonic discovery time used for freshness and queue ordering. pub discovered_at: std::time::Instant, } -type TorrentCallback = Arc; -type FilterCallback = Arc bool + Send + Sync>; -type ErrorCallback = Arc; - #[derive(Clone)] +/// Cloneable BEP-5 server handle and primary crate entry point. pub struct DHTServer { - #[allow(dead_code)] options: DHTOptions, + crawl_config: ResolvedCrawlConfig, node_id: [u8; 20], - socket_providers: Arc>>, + sockets_by_bind_addr: Arc>>, token_secret: [u8; 10], - callback: Arc>>, - filter: Arc>>, - on_metadata_fetch: Arc>>, - on_error_cb: Arc>>, - node_queue: Arc, - hash_tx: mpsc::Sender, + torrent_callback: Arc>, + hash_filter: Arc>, + on_metadata_fetch: Arc>, + metadata_completion_callback: Arc>, + error_callback: Arc>, + crawl_engine: Arc, + hash_events_tx: mpsc::Sender, metadata_queue_len: Arc, max_metadata_queue_size: usize, + runtime_stats: DhtRuntimeStats, shutdown: CancellationToken, } @@ -71,298 +270,325 @@ fn create_udp_sock(domain: Domain, ty: Type, addr: SocketAddr) -> std::io::Resul } let _ = sock.set_reuse_address(true); sock.set_nonblocking(true)?; - let _ = sock.set_recv_buffer_size(32 * 1024 * 1024); let _ = sock.set_send_buffer_size(8 * 1024 * 1024); - sock.bind(&addr.into())?; UdpSocket::from_std(sock.into()) } +fn split_u32_quota(total: u32, workers: usize, worker: usize) -> u32 { + let workers = workers.max(1) as u32; + total / workers + u32::from((worker as u32) < total % workers) +} + +fn split_u64_quota(total: u64, workers: usize, worker: usize) -> u64 { + let workers = workers.max(1) as u64; + total / workers + u64::from((worker as u64) < total % workers) +} + impl DHTServer { + /// Validates options, binds configured UDP sockets and constructs bounded pipelines. + /// + /// Background Metadata scheduling begins during construction. Active crawling and UDP receive + /// loops begin when [`Self::start`] is awaited. pub async fn new(options: DHTOptions) -> Result { + let crawl_config = ResolvedCrawlConfig::from_options(&options.crawl); + let runtime_stats = DhtRuntimeStats::with_limits(DhtRuntimeLimits { + metadata_queue: options.metadata.max_queue_size.max(1), + node_pool: crawl_config.pool_capacity, + node_pool_low_watermark: crawl_config.low_watermark, + find_node_in_flight: crawl_config.max_in_flight, + initial_find_node_rate: crawl_config.max_find_node_rate_per_sec, + hash_ingress_queue: options.hash_queue_capacity, + crawl_priority_queue: crawl_config.priority_event_channel_capacity, + crawl_discovery_queue: crawl_config.discovery_event_channel_capacity, + }); const ANY_V4_ADDR: SocketAddr = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0)), 8080); const ANY_V6_ADDR: SocketAddr = SocketAddr::new(IpAddr::V6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 0)), 8080); - let mut socket_providers = HashMap::new(); - // TODO: Check address reachability + + let mut sockets_by_bind_addr = HashMap::new(); match options.netmode { NetMode::Ipv4Only => { let mut addr = ANY_V4_ADDR; addr.set_port(options.port); let sock = create_udp_sock(Domain::IPV4, Type::DGRAM, addr)?; - socket_providers.insert(addr, Arc::new(sock)); + sockets_by_bind_addr.insert(addr, Arc::new(sock)); } NetMode::Ipv6Only => { let mut addr = ANY_V6_ADDR; addr.set_port(options.port); let sock = create_udp_sock(Domain::IPV6, Type::DGRAM, addr)?; - socket_providers.insert(addr, Arc::new(sock)); + sockets_by_bind_addr.insert(addr, Arc::new(sock)); } NetMode::DualStack => { - let mut addr = ANY_V4_ADDR; - addr.set_port(options.port); - let sock = create_udp_sock(Domain::IPV4, Type::DGRAM, addr)?; - socket_providers.insert(addr, Arc::new(sock)); - let mut addr = ANY_V6_ADDR; - addr.set_port(options.port); - let sock = create_udp_sock(Domain::IPV6, Type::DGRAM, addr)?; - socket_providers.insert(addr, Arc::new(sock)); - } - }; + let mut addr_v4 = ANY_V4_ADDR; + addr_v4.set_port(options.port); + let sock_v4 = create_udp_sock(Domain::IPV4, Type::DGRAM, addr_v4)?; + sockets_by_bind_addr.insert(addr_v4, Arc::new(sock_v4)); - let node_id = generate_random_id(); + let mut addr_v6 = ANY_V6_ADDR; + addr_v6.set_port(options.port); + let sock_v6 = create_udp_sock(Domain::IPV6, Type::DGRAM, addr_v6)?; + sockets_by_bind_addr.insert(addr_v6, Arc::new(sock_v6)); + } + } + + let node_id = random_node_id(); let mut token_secret = [0u8; 10]; rand::thread_rng().fill(&mut token_secret); - let node_queue = ShardedNodeQueue::new(options.node_queue_capacity); - - let (hash_tx, hash_rx) = mpsc::channel::(options.hash_queue_capacity); - - let fetcher = Arc::new(RbitFetcher::new(options.metadata_timeout)); - - let callback = Arc::new(RwLock::new(None)); - let on_metadata_fetch = Arc::new(RwLock::new(None)); - + let (hash_events_tx, hash_rx) = + mpsc::channel::(options.hash_queue_capacity); + let fetcher = Arc::new(RbitFetcher::new_with_runtime_stats( + options.metadata.timeout_secs, + options.metadata.peer_failure_cache_capacity, + options.metadata.peer_failure_ttl_secs, + runtime_stats.clone(), + )); + let torrent_callback = Arc::new(ArcSwapOption::empty()); + let on_metadata_fetch = Arc::new(ArcSwapOption::empty()); + let metadata_completion_callback = Arc::new(ArcSwapOption::empty()); let metadata_queue_len = Arc::new(AtomicUsize::new(0)); - let shutdown = CancellationToken::new(); - let shutdown_for_scheduler = shutdown.clone(); - let scheduler = MetadataScheduler::new( + let scheduler = MetadataScheduler::new_with_runtime_stats( hash_rx, fetcher, - options.max_metadata_queue_size, - options.max_metadata_worker_count, - callback.clone(), - on_metadata_fetch.clone(), + MetadataSchedulerLimits { + queue_size: options.metadata.max_queue_size, + concurrency: options.metadata.max_worker_count, + }, + MetadataSchedulerCallbacks { + torrent: torrent_callback.clone(), + fetch_gate: on_metadata_fetch.clone(), + completion: metadata_completion_callback.clone(), + }, metadata_queue_len.clone(), - shutdown_for_scheduler, + shutdown.clone(), + runtime_stats.clone(), ); + tokio::spawn(scheduler.run()); - tokio::spawn(async move { - scheduler.run().await; - }); - - let max_metadata_queue_size = options.max_metadata_queue_size; - let server = Self { + let max_metadata_queue_size = options.metadata.max_queue_size; + Ok(Self { options, + crawl_config: crawl_config.clone(), node_id, - socket_providers: Arc::new(socket_providers), + sockets_by_bind_addr: Arc::new(sockets_by_bind_addr), token_secret, - callback, + torrent_callback, + hash_filter: Arc::new(ArcSwapOption::empty()), on_metadata_fetch, - node_queue: Arc::new(node_queue), - filter: Arc::new(RwLock::new(None)), - on_error_cb: Arc::new(RwLock::new(None)), - hash_tx, + metadata_completion_callback, + error_callback: Arc::new(ArcSwapOption::empty()), + crawl_engine: Arc::new(CrawlEngine::new(crawl_config, runtime_stats.clone())), + hash_events_tx, metadata_queue_len, max_metadata_queue_size, + runtime_stats, shutdown, - }; - - Ok(server) + }) } + /// Registers the asynchronous admission gate invoked before the first real Peer attempt. + /// + /// Returning `false` rejects the job without downloading Metadata and without emitting a + /// [`MetadataFetchCompletion`]. If no gate is registered, jobs are admitted. pub fn on_metadata_fetch(&self, callback: F) where F: Fn(String) -> Fut + Send + Sync + 'static, Fut: Future + Send + 'static, { - *self.on_metadata_fetch.write().unwrap_or_else(|e| e.into_inner()) = - Some(Arc::new(move |hash| Box::pin(callback(hash)))); + let callback: Arc = + Arc::new(Box::new(move |hash| Box::pin(callback(hash)))); + self.on_metadata_fetch.store(Some(callback)); } + /// Registers a torrent callback whose return is implicitly treated as accepted delivery. + /// + /// Registering a new torrent callback replaces the previous one. pub fn on_torrent(&self, callback: F) where F: Fn(TorrentInfo) + Send + Sync + 'static, { - *self.callback.write().unwrap_or_else(|e| e.into_inner()) = Some(Arc::new(callback)); + let callback: Arc = Arc::new(Box::new(move |torrent| { + callback(torrent); + true + })); + self.torrent_callback.store(Some(callback)); } + /// Registers a torrent callback that acknowledges application delivery. + /// + /// Returning `true` produces [`crate::MetadataFetchCompletionStatus::Accepted`]. Returning + /// `false` produces [`crate::MetadataFetchCompletionStatus::DeliveryRejected`]: the Metadata + /// download was + /// valid, but the application did not accept it. A callback panic is caught and treated as + /// rejected delivery. Registering this callback replaces any previous torrent callback. + pub fn on_torrent_with_ack(&self, callback: F) + where + F: Fn(TorrentInfo) -> bool + Send + Sync + 'static, + { + let callback: Arc = Arc::new(Box::new(callback)); + self.torrent_callback.store(Some(callback)); + } + + /// Registers a callback invoked exactly once when an admitted Metadata job terminates. + /// + /// Gate rejection does not emit a completion. `attempts` counts real Peer network attempts; + /// failure-cache skips are excluded. Registering a new callback replaces the previous one. + pub fn on_metadata_fetch_complete(&self, callback: F) + where + F: Fn(MetadataFetchCompletion) + Send + Sync + 'static, + { + let callback: Arc = Arc::new(Box::new(callback)); + self.metadata_completion_callback.store(Some(callback)); + } + + /// Registers an early synchronous InfoHash filter for valid `announce_peer` queries. + /// + /// Returning `false` prevents the Hash from entering the bounded ingress queue. Registering a + /// new filter replaces the previous one. + pub fn filter(&self, filter: F) + where + F: Fn(&str) -> bool + Send + Sync + 'static, + { + let filter: Arc = Arc::new(Box::new(filter)); + self.hash_filter.store(Some(filter)); + } + + /// Alias for [`Self::filter`]. pub fn set_filter(&self, filter: F) where F: Fn(&str) -> bool + Send + Sync + 'static, { - *self.filter.write().unwrap_or_else(|e| e.into_inner()) = Some(Arc::new(filter)); + self.filter(filter); } + /// Registers a runtime error callback, replacing any previous callback. + /// + /// Initialization and `start` errors are still returned through [`Result`]. pub fn on_error(&self, callback: F) where F: Fn(crate::error::DHTError) + Send + Sync + 'static, { - *self.on_error_cb.write().unwrap_or_else(|e| e.into_inner()) = Some(Arc::new(callback)); + let callback: Arc = Arc::new(Box::new(callback)); + self.error_callback.store(Some(callback)); } fn emit_error(&self, error: crate::error::DHTError) { - if let Ok(cb) = self.on_error_cb.read() { - if let Some(f) = cb.as_ref() { - f(error); - } + if let Some(callback) = self.error_callback.load_full() { + callback(error); } } + /// Returns the current strict FIFO crawl-pool size. pub fn get_node_pool_size(&self) -> usize { - self.node_queue.len() + self.crawl_engine.node_count.load(Ordering::Relaxed) } + /// Returns a cheap cloneable handle to transport-neutral atomic runtime statistics. + pub fn runtime_stats(&self) -> DhtRuntimeStats { + self.runtime_stats.clone() + } + + /// Starts crawl and UDP background tasks, then waits until [`Self::shutdown`] is called. + /// + /// A shut-down server cannot be restarted. Construct a new [`DHTServer`] for another run. pub async fn start(&self) -> Result<()> { - // 检查是否已经被关闭 if self.shutdown.is_cancelled() { - log::warn!("⚠️ 尝试启动已关闭的服务器"); - return Err(crate::error::DHTError::Other("服务器已关闭".to_string())); + return Err(crate::error::DHTError::Other( + "DHT server is already shut down".to_string(), + )); } - let workers = self.spawn_workers(); - for sock in self.socket_providers.values().cloned() { - spawn_udp_listener(sock, workers.clone(), self.shutdown.clone())?; + self.crawl_engine.spawn( + self.options.netmode, + self.node_id, + &self.sockets_by_bind_addr, + self.metadata_queue_len.clone(), + self.max_metadata_queue_size, + self.shutdown.clone(), + ); + + let buffer_pool = UdpBufferPool::new(); + let workers = self.spawn_workers(buffer_pool.clone()); + for sock in self.sockets_by_bind_addr.values().cloned() { + spawn_udp_listener( + sock, + workers.clone(), + self.shutdown.clone(), + buffer_pool.clone(), + self.runtime_stats.clone(), + )?; } - self.bootstrap().await; - let server = self.clone(); - let shutdown = self.shutdown.clone(); - - tokio::spawn(async move { - let semaphore = Arc::new(Semaphore::new(2000)); - let mut loop_tick = 0; - - loop { - // 检查关闭信号 - if shutdown.is_cancelled() { - #[cfg(debug_assertions)] - log::trace!("主循环收到关闭信号,退出"); - break; - } - - let queue_len = server.metadata_queue_len.load(Ordering::Relaxed); - let queue_pressure = queue_len as f64 / server.max_metadata_queue_size as f64; - - #[cfg(feature = "metrics")] - { - gauge!("dht_metadata_queue_size").set(queue_len as f64); - gauge!("dht_metadata_worker_pressure").set(queue_pressure); - gauge!("dht_node_queue_size").set(server.node_queue.len() as f64); - } - - let (batch_size, sleep_duration) = if queue_pressure < 0.8 { - (200, Duration::from_millis(10)) - } else if queue_pressure < 0.95 { - (20, Duration::from_millis(500)) - } else { - (0, Duration::from_millis(1000)) - }; - - let filter_ipv6 = match server.options.netmode { - NetMode::Ipv4Only => Some(false), - NetMode::Ipv6Only => Some(true), - NetMode::DualStack => None, - }; - - let queue_empty = server.node_queue.is_empty_for(filter_ipv6); - - let nodes_batch = { - if queue_empty || batch_size == 0 { - None - } else { - Some(server.node_queue.pop_batch(batch_size, filter_ipv6)) - } - }; - - loop_tick += 1; - if nodes_batch.is_none() || loop_tick % 50 == 0 { - server.bootstrap().await; - if nodes_batch.is_none() { - tokio::select! { - _ = shutdown.cancelled() => break, - _ = tokio::time::sleep(sleep_duration) => {}, - } - continue; - } - } - - if let Some(nodes) = nodes_batch { - let node_id = server.node_id; - - for node in nodes { - let permit = match semaphore.clone().acquire_owned().await { - Ok(p) => p, - Err(_) => break, - }; - let node_id_clone = node_id; - let socket = match server.socket_for_addr(&node.addr) { - Some(sock) => sock, - None => { - log::warn!("未绑定任何地址"); - break; - } - }; - let node_addr = node.addr; - let node_id_for_target = node.id; - - tokio::spawn(async move { - let neighbor_id = - generate_neighbor_target(&node_id_for_target, &node_id_clone); - let random_target = generate_random_id(); - let _ = send_find_node_impl( - &node_addr, - &random_target, - &neighbor_id, - socket, - ) - .await; - drop(permit); - }); - } - } - - tokio::select! { - _ = shutdown.cancelled() => break, - _ = tokio::time::sleep(sleep_duration) => {}, - } - } - }); self.shutdown.cancelled().await; Ok(()) } - /// 显式关闭服务器,停止所有后台任务 + /// Cancels UDP, crawl and Metadata tasks. This method is safe to call more than once. pub fn shutdown(&self) { self.shutdown.cancel(); } - fn spawn_workers(&self) -> Vec { + fn spawn_workers(&self, buffer_pool: UdpBufferPool) -> Vec { let server = self.clone(); let shutdown = self.shutdown.clone(); - let num_workers = std::thread::available_parallelism() .map(|n| n.get()) .unwrap_or(8); - - let queue_size = 5000; + let queue_size = 5_000; let mut workers: Vec = Vec::with_capacity(num_workers); - for _ in 0..num_workers { + for worker_id in 0..num_workers { let (tx, mut rx) = mpsc::channel(queue_size); workers.push(tx); let server_clone = server.clone(); let cancellation_token = shutdown.clone(); + let pool = buffer_pool.clone(); + let packet_rate = split_u32_quota( + self.crawl_config.max_response_rate_per_sec, + num_workers, + worker_id, + ); + let byte_rate = split_u64_quota( + self.crawl_config.max_response_bytes_per_sec, + num_workers, + worker_id, + ); + let per_source_rate = self.crawl_config.max_response_rate_per_source; tokio::spawn(async move { + let mut response_limiter = WorkerResponseLimiter::new( + packet_rate, + byte_rate, + per_source_rate, + Instant::now(), + ); + let mut response_buffer = BytesMut::with_capacity(512); loop { tokio::select! { - _ = cancellation_token.cancelled() => { - #[cfg(debug_assertions)] - log::trace!("Worker 收到关闭信号,退出"); - break; - } + _ = cancellation_token.cancelled() => break, msg = rx.recv() => { match msg { - Some((data, remote_addr, local_addr)) => { - if let Err(e) = server_clone.handle_message(data.as_ref(), remote_addr, local_addr).await { + Some((packet, remote_addr, local_addr)) => { + if let Err(e) = server_clone + .handle_message( + packet.payload(), + remote_addr, + local_addr, + &mut response_limiter, + &mut response_buffer, + ) + .await + { server_clone.emit_error(e); } + pool.release(packet.buf); } None => break, } @@ -379,14 +605,10 @@ impl DHTServer { data: &[u8], remote_addr: SocketAddr, local_addr: SocketAddr, + response_limiter: &mut WorkerResponseLimiter, + response_buffer: &mut BytesMut, ) -> Result<()> { - if self.socket_providers.get(&local_addr).is_none() { - #[cfg(debug_assertions)] - log::trace!( - "⚠️ 拒绝未绑定的地址: {} (当前模式: {:?})", - remote_addr, - self.options.netmode - ); + if self.sockets_by_bind_addr.get(&local_addr).is_none() { return Ok(()); } @@ -401,12 +623,11 @@ impl DHTServer { #[cfg(feature = "metrics")] { - // 使用 match 映射到静态字符串,避免 clone(),同时防止恶意 tag let label = match msg.y.as_str() { "q" => "q", "r" => "r", "e" => "e", - _ => "unknown", // 将所有非法/未知类型归一化 + _ => "unknown", }; counter!("dht_messages_processed_total", "type" => label).increment(1); } @@ -414,13 +635,22 @@ impl DHTServer { match msg.y.as_str() { "q" => { if let Some(q_type) = &msg.q { - self.handle_query(&msg, q_type.as_bytes(), remote_addr, local_addr) - .await?; + self.handle_query( + &msg, + q_type.as_bytes(), + remote_addr, + local_addr, + response_limiter, + response_buffer, + ) + .await?; } } "r" => { - if let Some(response) = &msg.r { - self.handle_response(response).await?; + if let Some(response) = msg.r + && let Some(tid) = transaction_id_from_bytes(&msg.t) + { + self.crawl_engine.route_response(remote_addr, tid, response); } } _ => {} @@ -434,6 +664,8 @@ impl DHTServer { query_type: &[u8], remote_addr: SocketAddr, local_addr: SocketAddr, + response_limiter: &mut WorkerResponseLimiter, + response_buffer: &mut BytesMut, ) -> Result<()> { let args = match &msg.a { Some(a) => a, @@ -447,8 +679,20 @@ impl DHTServer { .as_deref() .or(args.info_hash.as_deref()) .map(|v| v.as_slice()); - let q_str = std::str::from_utf8(query_type).unwrap_or(""); + self.runtime_stats.inbound_query(q_str); + + if let Some(sender_id) = sender_id + && sender_id.len() == 20 + && is_valid_node_addr(&remote_addr) + { + self.crawl_engine.route_discovered(NodeTuple { + id: sender_id + .try_into() + .expect("validated DHT sender id contains 20 bytes"), + addr: remote_addr, + }); + } #[cfg(feature = "metrics")] { @@ -468,12 +712,16 @@ impl DHTServer { } self.send_response( - transaction_id, - remote_addr, - local_addr, - q_str, - sender_id, - target_id_fallback, + QueryResponse { + transaction_id, + remote_addr, + local_addr, + query_type: q_str, + sender_id, + target_id: target_id_fallback, + }, + response_limiter, + response_buffer, ) .await?; Ok(()) @@ -482,26 +730,31 @@ impl DHTServer { async fn handle_announce_peer(&self, args: &DhtArgs, addr: SocketAddr) -> Result<()> { if let Some(token) = &args.token { if !self.validate_token(token, addr) { + self.runtime_stats.announce_invalid_token(); #[cfg(feature = "metrics")] counter!("dht_announce_peer_blocked_total", "reason" => "invalid_token") .increment(1); return Ok(()); } } else { + self.runtime_stats.announce_invalid_token(); return Ok(()); } if let Some(info_hash) = &args.info_hash { let info_hash_arr: [u8; 20] = match info_hash.as_ref().try_into() { Ok(arr) => arr, - Err(_) => return Ok(()), + Err(_) => { + self.runtime_stats.announce_invalid_token(); + return Ok(()); + } }; let hash_hex = hex::encode(info_hash_arr); - let filter_cb = self.filter.read().unwrap_or_else(|e| e.into_inner()).clone(); - if let Some(f) = filter_cb - && !f(&hash_hex) + if let Some(filter) = self.hash_filter.load_full() + && !filter(&hash_hex) { + self.runtime_stats.announce_filtered(); #[cfg(feature = "metrics")] counter!("dht_announce_peer_blocked_total", "reason" => "filtered").increment(1); return Ok(()); @@ -510,9 +763,6 @@ impl DHTServer { #[cfg(feature = "metrics")] counter!("dht_info_hashes_discovered_total").increment(1); - #[cfg(debug_assertions)] - log::debug!("🔥 新 Hash: {} 来自 {}", hash_hex, addr); - let port = if let Some(implied) = args.implied_port { if implied != 0 { addr.port() @@ -524,228 +774,119 @@ impl DHTServer { }; if port > 0 { + self.runtime_stats.announce_accepted(); + self.runtime_stats.hash_received(); let event = HashDiscovered { info_hash: hash_hex, peer_addr: SocketAddr::new(addr.ip(), port), discovered_at: std::time::Instant::now(), }; - if self.hash_tx.try_send(event).is_err() { + let enqueue_result = self.hash_events_tx.try_send(event); + self.runtime_stats.set_hash_ingress_queue_depth( + self.hash_events_tx + .max_capacity() + .saturating_sub(self.hash_events_tx.capacity()), + ); + if enqueue_result.is_err() { + self.runtime_stats.hash_ingress_dropped(); + #[cfg(feature = "metrics")] + counter!("dht_metadata_ingress_dropped_total", "reason" => "queue_full") + .increment(1); #[cfg(debug_assertions)] - log::debug!("⚠️ Hash 队列满,丢弃 hash"); + log::debug!("Hash queue is full; dropping hash"); } } } Ok(()) } - async fn handle_response(&self, response: &DhtResponse) -> Result<()> { - if let Some(nodes_bytes) = &response.nodes { - self.process_compact_nodes(nodes_bytes); - } - if let Some(nodes6_bytes) = &response.nodes6 { - self.process_compact_nodes_v6(nodes6_bytes); - } - Ok(()) - } - - fn process_compact_nodes(&self, nodes_bytes: &[u8]) { - if self.options.netmode == NetMode::Ipv6Only { - return; - } - - #[allow(clippy::manual_is_multiple_of)] - if nodes_bytes.len() % 26 != 0 { - return; - } - - for chunk in nodes_bytes.chunks(26) { - let id = chunk[0..20].to_vec(); - let port = u16::from_be_bytes([chunk[24], chunk[25]]); - - let ip = std::net::Ipv4Addr::new(chunk[20], chunk[21], chunk[22], chunk[23]); - let addr = SocketAddr::new(std::net::IpAddr::V4(ip), port); - - #[cfg(feature = "metrics")] - counter!("dht_nodes_discovered_total", "ip_version" => "v4").increment(1); - - self.node_queue.push(NodeTuple { id, addr }); - } - } - - fn process_compact_nodes_v6(&self, nodes_bytes: &[u8]) { - if self.options.netmode == NetMode::Ipv4Only { - return; - } - - #[allow(clippy::manual_is_multiple_of)] - if nodes_bytes.len() % 38 != 0 { - return; - } - for chunk in nodes_bytes.chunks(38) { - let id = chunk[0..20].to_vec(); - let port = u16::from_be_bytes([chunk[36], chunk[37]]); - let ip_bytes: [u8; 16] = match chunk[20..36].try_into() { - Ok(b) => b, - Err(_) => continue, - }; - let ip = Ipv6Addr::from(ip_bytes); - if !ip.is_unspecified() && !ip.is_multicast() { - let addr = SocketAddr::new(IpAddr::V6(ip), port); - - #[cfg(feature = "metrics")] - counter!("dht_nodes_discovered_total", "ip_version" => "v6").increment(1); - - self.node_queue.push(NodeTuple { id, addr }); - } - } - } - async fn send_response( &self, - tid: &[u8], - remote_addr: SocketAddr, - local_addr: SocketAddr, - query_type: &str, - sender_id: Option<&[u8]>, - target_id_fallback: Option<&[u8]>, + response: QueryResponse<'_>, + response_limiter: &mut WorkerResponseLimiter, + response_buffer: &mut BytesMut, ) -> Result<()> { - let socket = match self.socket_providers.get(&local_addr) { + let socket = match self.sockets_by_bind_addr.get(&response.local_addr) { Some(sock) => sock, - None => return Ok(()), // Silent failure when the socket is not present + None => return Ok(()), }; - let mut r_dict = std::collections::HashMap::new(); - - let reference_id = sender_id.or(target_id_fallback); + let reference_id = response.sender_id.or(response.target_id); let my_id = if let Some(target) = reference_id { - generate_neighbor_target(target, &self.node_id) + let generated = neighbor_node_id(target, &self.node_id); + <[u8; 20]>::try_from(generated.as_slice()).expect("neighbor id is always 20 bytes") } else { - self.node_id.to_vec() + self.node_id }; - - r_dict.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(my_id)); - let token = self.generate_token(remote_addr); - r_dict.insert( - b"token".to_vec(), - serde_bencode::value::Value::Bytes(token.to_vec()), - ); - - if query_type == "get_peers" || query_type == "find_node" { - let requestor_is_ipv6 = remote_addr.is_ipv6(); + let token = self.generate_token(response.remote_addr); + let include_nodes = + response.query_type == "get_peers" || response.query_type == "find_node"; + let requestor_is_ipv6 = response.remote_addr.is_ipv6(); + let nodes = if include_nodes { let filter_ipv6 = match self.options.netmode { NetMode::Ipv4Only => Some(false), NetMode::Ipv6Only => Some(true), NetMode::DualStack => Some(requestor_is_ipv6), }; - - let nodes = self.node_queue.get_random_nodes(8, filter_ipv6); - - let mut nodes_data = Vec::new(); - let mut nodes6_data = Vec::new(); - - for node in nodes { - match node.addr.ip() { - IpAddr::V4(ip) => { - nodes_data.extend_from_slice(&node.id); - nodes_data.extend_from_slice(&ip.octets()); - nodes_data.extend_from_slice(&node.addr.port().to_be_bytes()); - } - IpAddr::V6(ip) => { - nodes6_data.extend_from_slice(&node.id); - nodes6_data.extend_from_slice(&ip.octets()); - nodes6_data.extend_from_slice(&node.addr.port().to_be_bytes()); - } - } + let snapshot = self.crawl_engine.snapshot.load(); + snapshot.random_nodes(8, filter_ipv6) + } else { + Vec::new() + }; + encode_response( + response_buffer, + response.transaction_id, + &my_id, + &token, + &nodes, + requestor_is_ipv6, + ); + let is_priority = response.query_type == "ping" || response.query_type == "get_peers"; + match response_limiter.acquire( + response.remote_addr, + response_buffer.len(), + is_priority, + Instant::now(), + ) { + ResponsePermit::Regular => self.runtime_stats.response_normal(), + ResponsePermit::PriorityReserve => { + self.runtime_stats.udp_response_priority_reserved(); + #[cfg(feature = "metrics")] + counter!( + "dht_udp_responses_priority_reserved_total", + "query" => if response.query_type == "ping" { "ping" } else { "get_peers" } + ) + .increment(1); } - - if requestor_is_ipv6 { - if !nodes6_data.is_empty() { - r_dict.insert( - b"nodes6".to_vec(), - serde_bencode::value::Value::Bytes(nodes6_data), - ); - } - } else if !nodes_data.is_empty() { - r_dict.insert( - b"nodes".to_vec(), - serde_bencode::value::Value::Bytes(nodes_data), - ); + ResponsePermit::Rejected => { + self.runtime_stats.udp_response_rate_limited(); + #[cfg(feature = "metrics")] + counter!("dht_udp_responses_dropped_total", "reason" => "rate_limit").increment(1); + return Ok(()); } } - - let mut response: std::collections::HashMap = - std::collections::HashMap::new(); - response.insert( - "t".to_string(), - serde_bencode::value::Value::Bytes(tid.to_vec()), - ); - response.insert( - "y".to_string(), - serde_bencode::value::Value::Bytes(b"r".to_vec()), - ); - response.insert("r".to_string(), serde_bencode::value::Value::Dict(r_dict)); - - if let Ok(encoded) = serde_bencode::to_bytes(&response) { - #[allow(unused)] - if let Ok(len) = socket.send_to(&encoded, remote_addr).await { + match socket.send_to(response_buffer, response.remote_addr).await { + Ok(len) => { + self.runtime_stats.udp_sent(len); #[cfg(feature = "metrics")] { counter!("dht_udp_bytes_sent_total").increment(len as u64); counter!("dht_udp_packets_sent_total", "type" => "response").increment(1); } } + Err(_) => self.runtime_stats.response_send_failed(), } Ok(()) } - async fn bootstrap(&self) { - let target = generate_random_id(); - for node in BOOTSTRAP_NODES { - if let Ok(addrs) = tokio::net::lookup_host(node).await { - for addr in addrs { - match self.options.netmode { - NetMode::Ipv4Only => { - if addr.is_ipv6() { - continue; - } - } - NetMode::Ipv6Only => { - if addr.is_ipv4() { - continue; - } - } - NetMode::DualStack => {} - } - let _ = self.send_find_node(&addr, &target, &self.node_id).await; - } - } - } - } - - fn socket_for_addr(&self, addr: &SocketAddr) -> Option> { - self.socket_providers - .iter() - .find(|(bind_addr, _)| bind_addr.is_ipv4() == addr.is_ipv4()) - .map(|(_, sock)| sock.clone()) - } - - async fn send_find_node(&self, target_addr: &SocketAddr, target: &[u8], sender_id: &[u8]) { - if let Some(sock) = self.socket_for_addr(target_addr) { - send_find_node_impl(target_addr, target, sender_id, sock).await - } - } - fn generate_token(&self, addr: SocketAddr) -> [u8; 8] { let mut hasher = ahash::AHasher::default(); - match addr.ip() { IpAddr::V4(ip) => ip.octets().hash(&mut hasher), IpAddr::V6(ip) => ip.octets().hash(&mut hasher), } - self.token_secret.hash(&mut hasher); - hasher.finish().to_le_bytes() } @@ -758,265 +899,65 @@ impl DHTServer { } } -/// 发送 DHT find_node 查询消息 -/// -/// 这是 DHT 协议中的核心操作之一,用于向指定节点查询包含目标 ID 的节点信息。 -/// 该方法构建符合 BEP5 (BitTorrent DHT Protocol) 规范的消息并异步发送。 -/// -/// # 参数 -/// -/// * `addr` - 目标节点的 Socket 地址 -/// * `target` - 要查找的目标节点 ID (20 字节) -/// * `sender_id` - 发送者的节点 ID (20 字节),用于标识自己 -/// * `socket` - IPv4 UDP socket 的引用 -/// * `socket_v6` - IPv6 UDP socket 的可选引用(仅在双栈模式下需要) -/// * `netmode` - 网络模式:仅 IPv4、仅 IPv6 或双栈模式 -/// -/// # 返回值 -/// -/// 返回 `Result<()>`,成功时返回 `Ok(())`,失败时返回错误信息 -/// -/// # 消息格式 -/// -/// 构建的 DHT 消息格式如下: -/// ```bencode -/// { -/// "t": [0, 1], // 事务 ID (transaction ID) -/// "y": "q", // 消息类型:查询 (query) -/// "q": "find_node", // 查询类型:查找节点 -/// "a": { // 参数 (arguments) -/// "id": , // 发送者节点 ID -/// "target": // 目标节点 ID -/// } -/// } -/// ``` -/// -/// # 网络模式处理 -/// -/// * `Ipv4Only`: 始终使用 IPv4 socket -/// * `Ipv6Only`: 始终使用 IPv4 socket(IPv6 模式下 socket 实际是 IPv6) -/// * `DualStack`: 根据目标地址类型自动选择 IPv4 或 IPv6 socket -async fn send_find_node_impl( - addr: &SocketAddr, - target: &[u8], - sender_id: &[u8], - socket: Arc, -) { - // 构建查询参数 - let mut args = std::collections::HashMap::new(); - args.insert( - b"id".to_vec(), - serde_bencode::value::Value::Bytes(sender_id.to_vec()), - ); - args.insert( - b"target".to_vec(), - serde_bencode::value::Value::Bytes(target.to_vec()), - ); +#[cfg(test)] +mod response_limiter_tests { + use super::*; - // 构建完整的 DHT 消息 - let mut msg: std::collections::HashMap = - std::collections::HashMap::new(); - msg.insert( - "t".to_string(), - serde_bencode::value::Value::Bytes(vec![0, 1]), - ); // 事务 ID - msg.insert( - "y".to_string(), - serde_bencode::value::Value::Bytes(b"q".to_vec()), - ); // 消息类型:查询 - msg.insert( - "q".to_string(), - serde_bencode::value::Value::Bytes(b"find_node".to_vec()), - ); // 查询类型 - msg.insert("a".to_string(), serde_bencode::value::Value::Dict(args)); // 参数字典 - - // 将消息编码为 bencode 格式并发送 - if let Ok(encoded) = serde_bencode::to_bytes(&msg) { - // 异步发送 UDP 数据包 - #[cfg(feature = "metrics")] - { - counter!("dht_udp_bytes_sent_total").increment(encoded.len() as u64); - counter!("dht_udp_packets_sent_total", "type" => "query").increment(1); - } - let _ = socket.send_to(&encoded, addr).await; - } -} - -fn generate_random_id() -> [u8; 20] { - let mut id = [0u8; 20]; - rand::thread_rng().fill(&mut id); - id -} - -/// 生成邻居目标节点 ID -/// -/// 该方法用于生成一个"看起来像"远程节点 ID 但实际基于本地节点 ID 的邻居节点 ID。 -/// 这是 DHT 协议中的一个重要优化策略,用于提高查询成功率和保护节点 ID 隐私。 -/// -/// # 工作原理 -/// -/// 1. 取远程节点 ID 的前 6 个字节作为前缀(如果远程 ID 长度足够) -/// 2. 用本地节点 ID 的剩余部分填充 -/// 3. 如果本地 ID 不够长,用随机字节填充到 20 字节(标准 DHT 节点 ID 长度) -/// -/// 这样生成的 ID 在 ID 空间中既接近远程节点(前 6 字节相同),又基于本地节点 -/// (后续字节来自本地 ID),从而在 DHT 路由时更容易获得相关响应。 -/// -/// # 参数 -/// -/// * `remote_id` - 远程节点的 ID(通常是查询目标节点或请求方的 ID) -/// * `local_id` - 本地节点的 ID(通常是自己真实的节点 ID) -/// -/// # 返回值 -/// -/// 返回一个 20 字节的节点 ID Vec,其前 6 字节来自 `remote_id`,后续字节来自 `local_id` -/// -/// # 使用场景 -/// -/// 1. **发送查询时**:使用邻居 ID 作为发送者 ID,让远程节点认为查询来自一个接近目标 ID 的节点, -/// 从而返回更相关的节点列表 -/// 2. **发送响应时**:使用邻居 ID 作为响应中的节点 ID,保护真实本地 ID 的隐私, -/// 同时提高返回节点的相关性 -/// -/// # 示例 -/// -/// ``` -/// // 假设: -/// // remote_id = [0x01, 0x02, 0x03, 0x04, 0x05, 0x06, ...] -/// // local_id = [0xAA, 0xBB, 0xCC, 0xDD, ...] -/// // 生成结果 = [0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0xCC, 0xDD, ...] -/// // (前6字节来自remote_id,后续来自local_id) -/// ``` -fn generate_neighbor_target(remote_id: &[u8], local_id: &[u8]) -> Vec { - let mut id = Vec::with_capacity(20); - let prefix_len = std::cmp::min(remote_id.len(), 6); - id.extend_from_slice(&remote_id[..prefix_len]); - - if local_id.len() > prefix_len { - id.extend_from_slice(&local_id[prefix_len..]); - } else { - while id.len() < 20 { - id.push(rand::random()); - } - } - id -} - -fn spawn_udp_listener( - socket: Arc, - mut workers: Vec, - shutdown: CancellationToken, -) -> crate::error::Result<()> { - let local_addr = socket - .local_addr() - .map_err(|e| crate::error::DHTError::Init(format!("socket 无法获取本地地址: {e}")))?; - if workers.is_empty() { - return Err(crate::error::DHTError::Init( - "spawn_udp_listener: 未提供任何 worker".to_string(), - )); - } - tokio::spawn(async move { - let mut buffer = [0u8; 65536]; - let mut worker_index = 0; - - loop { - tokio::select! { - _ = shutdown.cancelled() => { - #[cfg(debug_assertions)] - log::trace!("UDP 读取循环收到关闭信号,退出"); - break; - } - result = socket.recv_from(&mut buffer) => { - match result { - Ok((size, origin_addr)) => { - if let Err(ProcessUdpPacketError::NoLiveWorkers) = process_udp_packet(size, origin_addr, local_addr, &buffer, &mut workers, &mut worker_index){ - log::warn!("Socket {socket:?} is closing because no worker can process packets."); - // TODO: Remove the dead socket, or find a way to supply workers. - break - } - }, - Err(_e) => { - tokio::select! { - _ = shutdown.cancelled() => break, - _ = tokio::time::sleep(Duration::from_millis(1)) => {}, - } - } - } - } - } - } - }); - Ok(()) -} - -enum ProcessUdpPacketError { - PacketTooLarge, - InvalidPacket, - ChokedWorkers, - NoLiveWorkers, -} - -fn process_udp_packet( - size: usize, - origin_addr: SocketAddr, - local_addr: SocketAddr, - buffer: &[u8], - workers: &mut Vec, - worker_index: &mut usize, -) -> std::result::Result<(), ProcessUdpPacketError> { - #[cfg(feature = "metrics")] - counter!("dht_udp_bytes_received_total").increment(size as u64); - - if size > 8192 { - #[cfg(feature = "metrics")] - counter!("dht_udp_packets_received_total", "status" => "dropped_size").increment(1); - - #[cfg(debug_assertions)] - log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, origin_addr); - return Err(ProcessUdpPacketError::PacketTooLarge); + #[test] + fn response_limiter_enforces_packet_byte_and_source_limits() { + let start = Instant::now(); + let addr: SocketAddr = "8.8.8.8:6881".parse().unwrap(); + let mut limiter = WorkerResponseLimiter::new(2, 200, 1, start); + assert_eq!( + limiter.acquire(addr, 100, false, start), + ResponsePermit::Regular + ); + assert_eq!( + limiter.acquire(addr, 100, false, start), + ResponsePermit::Rejected + ); + assert_eq!( + limiter.acquire(addr, 100, false, start + Duration::from_secs(1)), + ResponsePermit::Regular + ); + assert_eq!( + limiter.acquire( + "1.1.1.1:6881".parse().unwrap(), + 513, + false, + start + Duration::from_secs(2) + ), + ResponsePermit::Rejected + ); } - if size == 0 || buffer[0] != b'd' { - #[cfg(feature = "metrics")] - counter!("dht_udp_packets_received_total", "status" => "dropped_magic").increment(1); - return Err(ProcessUdpPacketError::InvalidPacket); + #[test] + fn ping_and_get_peers_can_use_the_priority_reserve() { + let start = Instant::now(); + let mut limiter = WorkerResponseLimiter::new(10, 1_000, 100, start); + let first: SocketAddr = "8.8.8.8:6881".parse().unwrap(); + let second: SocketAddr = "1.1.1.1:6881".parse().unwrap(); + let priority: SocketAddr = "9.9.9.9:6881".parse().unwrap(); + + assert_eq!( + limiter.acquire(first, 50, false, start), + ResponsePermit::Regular + ); + assert_eq!( + limiter.acquire(second, 50, false, start), + ResponsePermit::Regular + ); + assert_eq!( + limiter.acquire(priority, 50, false, start), + ResponsePermit::Rejected + ); + assert_eq!( + limiter.acquire(priority, 50, true, start), + ResponsePermit::PriorityReserve + ); + assert_eq!( + limiter.acquire("4.4.4.4:6881".parse().unwrap(), 50, true, start), + ResponsePermit::Rejected + ); } - let mut data = Some(buffer[..size].to_owned().into_boxed_slice()); - let mut attempts = 0; - let max_attempts = workers.len(); - - while let Some(packet) = data.take() { - let worker = &workers[*worker_index]; - match worker.try_send((packet, origin_addr, local_addr)) { - Ok(_) => { - #[cfg(feature = "metrics")] - counter!("dht_udp_packets_received_total", "status" => "ok").increment(1); - break; - } - Err(mpsc::error::TrySendError::Full((packet, _, _))) => { - attempts += 1; - if attempts >= max_attempts { - #[cfg(feature = "metrics")] - counter!("dht_udp_packets_received_total", "status" => "queue_full") - .increment(1); - - #[cfg(debug_assertions)] - log::trace!("UDP worker queue full, dropping packet"); - return Err(ProcessUdpPacketError::ChokedWorkers); - } - let _ = data.insert(packet); - } - Err(mpsc::error::TrySendError::Closed((packet, _, _))) => { - log::warn!("UDP worker dropped."); - workers.swap_remove(*worker_index); - let _ = data.insert(packet); - } - } - if workers.is_empty() { - return Err(ProcessUdpPacketError::NoLiveWorkers); - } - *worker_index = (*worker_index + 1) % workers.len(); - } - - Ok(()) } diff --git a/src/sharded.rs b/src/sharded.rs deleted file mode 100644 index 767ec17..0000000 --- a/src/sharded.rs +++ /dev/null @@ -1,291 +0,0 @@ -use std::collections::{HashSet, VecDeque}; -use std::net::SocketAddr; -use std::sync::Mutex; - -const QUEUE_SHARD_COUNT: usize = 16; - -#[derive(Debug, Clone)] -pub struct NodeTuple { - pub id: Vec, - pub addr: SocketAddr, -} - -struct NodeQueueShard { - queue: VecDeque, - index: HashSet, - capacity: usize, -} - -impl NodeQueueShard { - fn new(capacity: usize) -> Self { - Self { - queue: VecDeque::with_capacity(capacity), - index: HashSet::with_capacity(capacity), - capacity, - } - } - - fn push(&mut self, node: NodeTuple) { - if self.index.contains(&node.addr) { - return; - } - - if self.queue.len() >= self.capacity - && let Some(removed) = self.queue.pop_front() - { - self.index.remove(&removed.addr); - } - - self.index.insert(node.addr); - self.queue.push_back(node); - } - - fn pop_batch(&mut self, count: usize) -> Vec { - let actual_count = count.min(self.queue.len()); - let mut nodes = Vec::with_capacity(actual_count); - - for _ in 0..actual_count { - if let Some(node) = self.queue.pop_front() { - self.index.remove(&node.addr); - nodes.push(node); - } - } - nodes - } - - fn len(&self) -> usize { - self.queue.len() - } - - fn is_empty(&self) -> bool { - self.queue.is_empty() - } -} - -pub struct ShardedNodeQueue { - shards_v4: Vec>, - shards_v6: Vec>, -} - -impl ShardedNodeQueue { - pub fn new(total_capacity: usize) -> Self { - #[allow(clippy::manual_div_ceil)] - let capacity_per_shard = (total_capacity + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - - let shards_v4 = (0..QUEUE_SHARD_COUNT) - .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) - .collect(); - - let shards_v6 = (0..QUEUE_SHARD_COUNT) - .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) - .collect(); - - Self { - shards_v4, - shards_v6, - } - } - - pub fn push(&self, node: NodeTuple) { - let shard_idx = self.addr_to_shard(&node.addr); - - if node.addr.is_ipv6() { - let mut shard = self.shards_v6[shard_idx].lock().unwrap_or_else(|e| e.into_inner()); - shard.push(node); - } else { - let mut shard = self.shards_v4[shard_idx].lock().unwrap_or_else(|e| e.into_inner()); - shard.push(node); - } - } - - pub fn pop_batch(&self, count: usize, filter_ipv6: Option) -> Vec { - let mut result = Vec::with_capacity(count); - #[allow(clippy::manual_div_ceil)] - let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - - match filter_ipv6 { - Some(true) => { - for shard in &self.shards_v6 { - if result.len() >= count { - break; - } - let mut s = shard.lock().unwrap_or_else(|e| e.into_inner()); - let nodes = s.pop_batch(per_shard); - result.extend(nodes); - } - } - Some(false) => { - for shard in &self.shards_v4 { - if result.len() >= count { - break; - } - let mut s = shard.lock().unwrap_or_else(|e| e.into_inner()); - let nodes = s.pop_batch(per_shard); - result.extend(nodes); - } - } - None => { - for i in 0..QUEUE_SHARD_COUNT { - if result.len() >= count { - break; - } - - let mut s4 = self.shards_v4[i].lock().unwrap_or_else(|e| e.into_inner()); - let nodes4 = s4.pop_batch(per_shard / 2); - result.extend(nodes4); - drop(s4); - - if result.len() >= count { - break; - } - - let mut s6 = self.shards_v6[i].lock().unwrap_or_else(|e| e.into_inner()); - let nodes6 = s6.pop_batch(per_shard / 2); - result.extend(nodes6); - drop(s6); - } - } - } - - result - } - - pub fn get_random_nodes(&self, count: usize, filter_ipv6: Option) -> Vec { - match filter_ipv6 { - Some(true) => self.get_random_nodes_from_shards(&self.shards_v6, count), - Some(false) => self.get_random_nodes_from_shards(&self.shards_v4, count), - None => { - let count_v4 = count / 2; - let count_v6 = count - count_v4; - let mut result = Vec::with_capacity(count); - - result.extend(self.get_random_nodes_from_shards(&self.shards_v4, count_v4)); - result.extend(self.get_random_nodes_from_shards(&self.shards_v6, count_v6)); - - result - } - } - } - - fn get_random_nodes_from_shards( - &self, - shards: &[Mutex], - count: usize, - ) -> Vec { - use rand::Rng; - let mut rng = rand::thread_rng(); - - if count <= 16 { - let mut result = Vec::with_capacity(count); - #[allow(clippy::manual_div_ceil)] - let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - - for shard in shards { - if result.len() >= count { - break; - } - - let s = shard.lock().unwrap_or_else(|e| e.into_inner()); - let shard_len = s.queue.len(); - - if shard_len == 0 { - continue; - } - - let to_take = per_shard.min(shard_len).min(count - result.len()); - - let mut indices: Vec = (0..shard_len).collect(); - - for i in 0..to_take { - let j = rng.gen_range(i..shard_len); - indices.swap(i, j); - } - - for &idx in indices.iter().take(to_take) { - if let Some(node) = s.queue.get(idx) { - result.push(node.clone()); - } - } - } - - result - } else { - let mut result = Vec::with_capacity(count); - let mut seen = 0usize; - - for shard in shards { - let s = shard.lock().unwrap_or_else(|e| e.into_inner()); - - for node in s.queue.iter() { - seen += 1; - - if result.len() < count { - result.push(node.clone()); - } else { - let j = rng.gen_range(0..seen); - if j < count { - result[j] = node.clone(); - } - } - } - } - - result - } - } - - pub fn len(&self) -> usize { - let len_v4: usize = self - .shards_v4 - .iter() - .map(|shard| shard.lock().unwrap_or_else(|e| e.into_inner()).len()) - .sum(); - let len_v6: usize = self - .shards_v6 - .iter() - .map(|shard| shard.lock().unwrap_or_else(|e| e.into_inner()).len()) - .sum(); - len_v4 + len_v6 - } - - pub fn is_empty(&self) -> bool { - let empty_v4 = self - .shards_v4 - .iter() - .all(|shard| shard.lock().unwrap_or_else(|e| e.into_inner()).is_empty()); - let empty_v6 = self - .shards_v6 - .iter() - .all(|shard| shard.lock().unwrap_or_else(|e| e.into_inner()).is_empty()); - empty_v4 && empty_v6 - } - - pub fn is_empty_for(&self, filter_ipv6: Option) -> bool { - match filter_ipv6 { - Some(true) => self - .shards_v6 - .iter() - .all(|shard| shard.lock().unwrap_or_else(|e| e.into_inner()).is_empty()), - Some(false) => self - .shards_v4 - .iter() - .all(|shard| shard.lock().unwrap_or_else(|e| e.into_inner()).is_empty()), - None => self.is_empty(), - } - } - - #[inline] - fn addr_to_shard(&self, addr: &SocketAddr) -> usize { - let hash = match addr.ip() { - std::net::IpAddr::V4(ip) => { - let octets = ip.octets(); - (octets[3] as usize) ^ (addr.port() as usize) - } - std::net::IpAddr::V6(ip) => { - let octets = ip.octets(); - (octets[15] as usize) ^ (addr.port() as usize) - } - }; - hash % QUEUE_SHARD_COUNT - } -} diff --git a/src/types.rs b/src/types.rs index 3e02e4c..4195d2a 100644 --- a/src/types.rs +++ b/src/types.rs @@ -1,38 +1,95 @@ use serde::{Deserialize, Serialize}; +use std::net::SocketAddr; #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +/// IP families on which the DHT server listens and crawls. pub enum NetMode { + /// Bind and crawl IPv4 only. Ipv4Only, + /// Bind and crawl IPv6 only. Ipv6Only, #[default] + /// Bind separate IPv4 and IPv6 sockets. DualStack, } #[derive(Debug, Clone, Serialize, Deserialize)] +/// Validated torrent metadata delivered to the application callback. pub struct TorrentInfo { + /// Lowercase hexadecimal SHA1 of the bencoded info dictionary. pub info_hash: String, + /// Magnet URI containing the InfoHash. pub magnet_link: String, + /// Torrent display name. pub name: String, + /// Sum of file sizes in bytes. pub total_size: u64, + /// Files described by the torrent. pub files: Vec, + /// Torrent piece length in bytes, or zero if absent. pub piece_length: u64, + /// Peer addresses used to obtain the Metadata. pub peers: Vec, + /// Completion time as Unix seconds. pub timestamp: u64, } +/// Final outcome of a metadata fetch that passed the admission callback. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum MetadataFetchCompletionStatus { + /// Metadata was fetched and accepted by the torrent callback. + Accepted, + /// All available peer candidates failed. + FetchFailed, + /// Metadata was fetched, but the application did not accept it. + DeliveryRejected, +} + +/// Report emitted exactly once after an admitted metadata fetch finishes. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct MetadataFetchCompletion { + /// InfoHash that reached a terminal state. + pub info_hash: String, + /// Final download/delivery status. + pub status: MetadataFetchCompletionStatus, + /// Real Peer network attempts; failure-cache skips are excluded. + pub attempts: usize, +} + +impl MetadataFetchCompletion { + /// Returns true only for [`MetadataFetchCompletionStatus::Accepted`]. + pub fn is_success(&self) -> bool { + self.status == MetadataFetchCompletionStatus::Accepted + } +} + #[derive(Debug, Clone, Serialize, Deserialize)] +/// One file entry from the validated info dictionary. pub struct FileInfo { + /// Slash-separated relative path. pub path: String, + /// File size in bytes. pub size: u64, } +#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq)] +/// Compact DHT node tuple used by crawl and routing code. +pub struct NodeTuple { + /// Twenty-byte DHT node ID. + pub id: [u8; 20], + /// Public UDP endpoint. + pub addr: SocketAddr, +} + impl TorrentInfo { + /// Formats [`Self::total_size`] using binary thresholds and a short unit suffix. pub fn format_size(&self) -> String { format_bytes(self.total_size) } } impl FileInfo { + /// Formats [`Self::size`] using binary thresholds and a short unit suffix. pub fn format_size(&self) -> String { format_bytes(self.size) } @@ -42,42 +99,233 @@ fn format_bytes(bytes: u64) -> String { const UNITS: &[&str] = &["B", "KB", "MB", "GB", "TB"]; let mut size = bytes as f64; let mut unit_index = 0; - while size >= 1024.0 && unit_index < UNITS.len() - 1 { size /= 1024.0; unit_index += 1; } - - format!("{:.2} {}", size, UNITS[unit_index]) + format!("{size:.2} {}", UNITS[unit_index]) } #[derive(Debug, Clone)] +/// Complete server configuration. pub struct DHTOptions { + /// UDP listen port. pub port: u16, - - pub metadata_timeout: u64, - - pub max_metadata_queue_size: usize, - - pub max_metadata_worker_count: usize, - + /// Enabled IP families. pub netmode: NetMode, - - pub node_queue_capacity: usize, - + /// Capacity between announce processing and the Metadata scheduler. pub hash_queue_capacity: usize, + /// Metadata download and Peer-cache limits. + pub metadata: MetadataOptions, + /// Active crawl, node-pool and scheduler limits. + pub crawl: CrawlOptions, +} + +#[derive(Debug, Clone)] +/// Metadata download and failure-cache limits. +pub struct MetadataOptions { + /// End-to-end timeout for one Peer attempt, in seconds. + pub timeout_secs: u64, + /// Maximum number of deduplicated pending InfoHashes. + pub max_queue_size: usize, + /// Maximum number of concurrent Metadata jobs. + pub max_worker_count: usize, + /// Maximum number of cached bad Peer socket addresses. + pub peer_failure_cache_capacity: usize, + /// Timeout/connect failure cache lifetime in seconds. + pub peer_failure_ttl_secs: u64, +} + +#[derive(Debug, Clone, Default)] +/// Active crawl configuration grouped by responsibility. +pub struct CrawlOptions { + /// FIFO node-pool and responsive-ring limits. + pub pool: PoolOptions, + /// Query, replacement and response budgets. + pub rate_limit: RateLimitOptions, + /// Bootstrap sources and retry policy. + pub bootstrap: BootstrapOptions, + /// Target-generation policy. + pub target: TargetOptions, + /// Internal bounded-channel and snapshot limits. + pub scheduler: SchedulerOptions, +} + +#[derive(Debug, Clone)] +/// Independent crawl and UDP-response budgets. +pub struct RateLimitOptions { + /// Maximum active find_node queries scheduled per second. + pub max_find_node_rate_per_sec: u32, + /// Maximum query budget consumed in one scheduler tick. + pub burst: u32, + /// Maximum total pending find_node transactions. + pub max_in_flight: usize, + /// Pending find_node timeout in seconds. + pub request_timeout_secs: u64, + /// Maximum never-before-probed destinations per minute. + pub max_new_destinations_per_minute: u32, + /// Maximum outbound DHT response packets per second. + pub max_response_rate_per_sec: u32, + /// Maximum encoded outbound DHT response bytes per second. + pub max_response_bytes_per_sec: u64, + /// Maximum response packets per source address per second. + pub max_response_rate_per_source: u32, + /// Remaining query-rate percentage when Metadata pressure reaches 95%. + pub metadata_pressure_floor_percent: u8, + /// Maximum FIFO replacements per minute after the pool has warmed. + pub max_replacements_per_minute: u32, + /// Maximum pending find_node transactions per IP subnet. + pub max_in_flight_per_subnet: usize, +} + +#[derive(Debug, Clone)] +/// FIFO crawl-pool and responsive-node reservoir limits. +pub struct PoolOptions { + /// Maximum queued crawl nodes. + pub capacity: usize, + /// How long a probed endpoint is blocked from readmission, in seconds. + pub recent_probe_ttl_secs: u64, + /// Maximum nodes retained for replies and revisit traffic. + pub responsive_capacity: usize, + /// Responsive-node lifetime in seconds. + pub responsive_ttl_secs: u64, + /// Pool size below which bootstrap is considered. + pub low_watermark: usize, +} + +#[derive(Debug, Clone)] +/// Bootstrap hostnames and retry timing. +pub struct BootstrapOptions { + /// Host:port sources resolved when bootstrap is needed. + pub nodes: Vec, + /// Minimum interval between bootstrap rounds, in seconds. + pub interval_secs: u64, + /// Maximum resolved endpoints selected in one round. + pub max_nodes_per_round: usize, + /// Initial failed-source backoff in seconds. + pub source_backoff_base_secs: u64, + /// Maximum failed-source backoff in seconds. + pub source_backoff_max_secs: u64, +} + +#[derive(Debug, Clone)] +/// Distribution used to generate find_node targets and sender IDs. +pub struct TargetOptions { + /// Percentage of targets that are fully random. + pub random_walk_percent: u8, + /// Percentage of targets chosen from sparse routing buckets. + pub sparse_bucket_percent: u8, + /// Whether outbound sender IDs borrow the target's prefix. + pub neighbor_sender_id: bool, +} + +#[derive(Debug, Clone)] +/// Capacities and batch limits for the crawl actor. +pub struct SchedulerOptions { + /// Capacity for response/bootstrap priority events. + pub priority_event_channel_capacity: usize, + /// Capacity for newly discovered node events. + pub discovery_event_channel_capacity: usize, + /// Maximum events drained per actor iteration. + pub event_batch_limit: usize, + /// Maximum discovery nodes drained per actor iteration. + pub node_batch_limit: usize, + /// Maximum responsive nodes published in the lock-free snapshot. + pub routing_snapshot_size: usize, + /// Snapshot publication interval in milliseconds. + pub snapshot_refresh_millis: u64, } impl Default for DHTOptions { fn default() -> Self { Self { port: 6881, - metadata_timeout: 3, - max_metadata_queue_size: 100000, - max_metadata_worker_count: 1000, netmode: NetMode::Ipv4Only, - node_queue_capacity: 100000, - hash_queue_capacity: 10000, + hash_queue_capacity: 10_000, + metadata: MetadataOptions::default(), + crawl: CrawlOptions::default(), + } + } +} + +impl Default for MetadataOptions { + fn default() -> Self { + Self { + timeout_secs: 4, + max_queue_size: 10_000, + max_worker_count: 256, + peer_failure_cache_capacity: 200_000, + peer_failure_ttl_secs: 60, + } + } +} + +impl Default for RateLimitOptions { + fn default() -> Self { + Self { + max_find_node_rate_per_sec: 200, + burst: 40, + max_in_flight: 512, + request_timeout_secs: 2, + max_new_destinations_per_minute: 10_000, + max_response_rate_per_sec: 500, + max_response_bytes_per_sec: 1024 * 1024, + max_response_rate_per_source: 40, + metadata_pressure_floor_percent: 25, + max_replacements_per_minute: 25_000, + max_in_flight_per_subnet: 8, + } + } +} + +impl Default for PoolOptions { + fn default() -> Self { + Self { + capacity: 100_000, + recent_probe_ttl_secs: 600, + responsive_capacity: 16_384, + responsive_ttl_secs: 900, + low_watermark: 10_000, + } + } +} + +impl Default for BootstrapOptions { + fn default() -> Self { + Self { + nodes: vec![ + "router.bittorrent.com:6881".to_string(), + "dht.transmissionbt.com:6881".to_string(), + "router.utorrent.com:6881".to_string(), + "dht.aelitis.com:6881".to_string(), + ], + interval_secs: 300, + max_nodes_per_round: 3, + source_backoff_base_secs: 300, + source_backoff_max_secs: 3_600, + } + } +} + +impl Default for TargetOptions { + fn default() -> Self { + Self { + random_walk_percent: 70, + sparse_bucket_percent: 30, + neighbor_sender_id: true, + } + } +} + +impl Default for SchedulerOptions { + fn default() -> Self { + Self { + priority_event_channel_capacity: 8_192, + discovery_event_channel_capacity: 16_384, + event_batch_limit: 256, + node_batch_limit: 4_096, + routing_snapshot_size: 4_096, + snapshot_refresh_millis: 1_000, } } } diff --git a/src/udp_buffer.rs b/src/udp_buffer.rs new file mode 100644 index 0000000..d11d339 --- /dev/null +++ b/src/udp_buffer.rs @@ -0,0 +1,79 @@ +//! UDP 收包缓冲区池:避免每包 `to_owned()` 拷贝。 +//! +//! 单线程 listener 从池中取出固定大小缓冲区,`recv_from` 直接写入; +//! 通过 channel 将缓冲区所有权交给 worker,处理完毕后归还池中复用。 + +use crossbeam_queue::ArrayQueue; +use std::sync::Arc; + +/// 与 `process_udp_packet` 中丢弃阈值一致 +pub const MAX_DHT_UDP_PACKET: usize = 8192; + +/// 预分配缓冲区数量(约等于高峰在途包数) +const INITIAL_POOL_SIZE: usize = 512; +/// 池上限,防止极端背压下无限增长 +const MAX_POOL_SIZE: usize = 4096; + +/// 在途 UDP 包:固定容量缓冲区 + 有效长度 +pub struct UdpPacket { + pub buf: Box<[u8]>, + pub len: usize, +} + +impl UdpPacket { + #[inline] + pub fn payload(&self) -> &[u8] { + &self.buf[..self.len] + } +} + +/// 固定 8KiB 缓冲区的对象池(`recv_from` 零拷贝移交 worker) +#[derive(Clone)] +pub struct UdpBufferPool { + inner: Arc, +} + +struct PoolInner { + free: ArrayQueue>, + buf_capacity: usize, +} + +impl UdpBufferPool { + pub fn new() -> Self { + let free = ArrayQueue::new(MAX_POOL_SIZE); + for _ in 0..INITIAL_POOL_SIZE { + let _ = free.push(alloc_buffer(MAX_DHT_UDP_PACKET)); + } + Self { + inner: Arc::new(PoolInner { + free, + buf_capacity: MAX_DHT_UDP_PACKET, + }), + } + } + + /// 取一块缓冲区;池空时分配新块(背压或突发流量) + pub fn acquire(&self) -> Box<[u8]> { + self.inner + .free + .pop() + .unwrap_or_else(|| alloc_buffer(self.inner.buf_capacity)) + } + + /// 归还缓冲区;池满时直接丢弃,由 GC 回收 + pub fn release(&self, buf: Box<[u8]>) { + if buf.len() != self.inner.buf_capacity { + return; + } + let _ = self.inner.free.push(buf); + } + + pub fn buf_capacity(&self) -> usize { + self.inner.buf_capacity + } +} + +fn alloc_buffer(capacity: usize) -> Box<[u8]> { + let v = vec![0; capacity]; + v.into_boxed_slice() +} diff --git a/src/udp_ingress.rs b/src/udp_ingress.rs new file mode 100644 index 0000000..bd6e8fb --- /dev/null +++ b/src/udp_ingress.rs @@ -0,0 +1,345 @@ +use crate::error::DHTError; +use crate::runtime_stats::DhtRuntimeStats; +use crate::udp_buffer::{MAX_DHT_UDP_PACKET, UdpBufferPool, UdpPacket}; +#[cfg(feature = "metrics")] +use metrics::counter; +use std::hash::{Hash, Hasher}; +use std::net::SocketAddr; +use std::sync::Arc; +use std::time::Duration; +use tokio::net::UdpSocket; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; + +pub(crate) type WorkerHandle = mpsc::Sender<(UdpPacket, SocketAddr, SocketAddr)>; + +pub(crate) fn spawn_udp_listener( + socket: Arc, + mut workers: Vec, + shutdown: CancellationToken, + buffer_pool: UdpBufferPool, + runtime_stats: DhtRuntimeStats, +) -> crate::error::Result<()> { + let local_addr = socket + .local_addr() + .map_err(|e| DHTError::Init(format!("socket local addr failed: {e}")))?; + if workers.is_empty() { + return Err(DHTError::Init( + "spawn_udp_listener: no worker provided".to_string(), + )); + } + tokio::spawn(async move { + loop { + let mut buf = buffer_pool.acquire(); + let recv_buf = &mut buf[..buffer_pool.buf_capacity()]; + + tokio::select! { + _ = shutdown.cancelled() => { + buffer_pool.release(buf); + break; + } + result = socket.recv_from(recv_buf) => { + match result { + Ok((size, origin_addr)) => { + if let Err(ProcessUdpPacketError::NoLiveWorkers) = + process_udp_packet(buf, size, origin_addr, local_addr, &buffer_pool, &runtime_stats, &mut workers) + { + log::warn!("Socket {socket:?} is closing because no worker can process packets."); + break + } + } + Err(_) => { + buffer_pool.release(buf); + tokio::select! { + _ = shutdown.cancelled() => break, + _ = tokio::time::sleep(Duration::from_millis(1)) => {}, + } + } + } + } + } + } + }); + Ok(()) +} + +enum ProcessUdpPacketError { + PacketTooLarge, + InvalidPacket, + ChokedWorkers, + NoLiveWorkers, +} + +fn process_udp_packet( + buf: Box<[u8]>, + size: usize, + origin_addr: SocketAddr, + local_addr: SocketAddr, + buffer_pool: &UdpBufferPool, + runtime_stats: &DhtRuntimeStats, + workers: &mut Vec, +) -> std::result::Result<(), ProcessUdpPacketError> { + runtime_stats.udp_received(); + runtime_stats.udp_received_bytes(size); + #[cfg(feature = "metrics")] + counter!("dht_udp_bytes_received_total").increment(size as u64); + + if size > MAX_DHT_UDP_PACKET { + runtime_stats.udp_invalid(); + #[cfg(feature = "metrics")] + counter!("dht_udp_packets_received_total", "status" => "dropped_size").increment(1); + buffer_pool.release(buf); + return Err(ProcessUdpPacketError::PacketTooLarge); + } + + if size == 0 || buf[0] != b'd' { + runtime_stats.udp_invalid(); + #[cfg(feature = "metrics")] + counter!("dht_udp_packets_received_total", "status" => "dropped_magic").increment(1); + buffer_pool.release(buf); + return Err(ProcessUdpPacketError::InvalidPacket); + } + + let mut packet = UdpPacket { buf, len: size }; + let mut hasher = ahash::AHasher::default(); + origin_addr.hash(&mut hasher); + let origin_hash = hasher.finish() as usize; + + 'select_worker: loop { + if workers.is_empty() { + buffer_pool.release(packet.buf); + return Err(ProcessUdpPacketError::NoLiveWorkers); + } + + let worker_count = workers.len(); + let preferred_index = origin_hash % worker_count; + for offset in 0..worker_count { + let worker_index = (preferred_index + offset) % worker_count; + match workers[worker_index].try_send((packet, origin_addr, local_addr)) { + Ok(_) => { + #[cfg(feature = "metrics")] + counter!("dht_udp_packets_received_total", "status" => "ok").increment(1); + return Ok(()); + } + Err(mpsc::error::TrySendError::Full((p, _, _))) => { + packet = p; + } + Err(mpsc::error::TrySendError::Closed((p, _, _))) => { + packet = p; + log::warn!("UDP worker dropped."); + workers.swap_remove(worker_index); + continue 'select_worker; + } + } + } + + #[cfg(feature = "metrics")] + counter!("dht_udp_packets_received_total", "status" => "queue_full").increment(1); + runtime_stats.udp_queue_full(); + buffer_pool.release(packet.buf); + return Err(ProcessUdpPacketError::ChokedWorkers); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn addresses() -> (SocketAddr, SocketAddr) { + ( + "8.8.8.8:6881".parse().unwrap(), + "0.0.0.0:12313".parse().unwrap(), + ) + } + + fn buffer(pool: &UdpBufferPool, first: u8) -> Box<[u8]> { + let mut buf = pool.acquire(); + buf[0] = first; + buf + } + + fn packet(pool: &UdpBufferPool, first: u8) -> UdpPacket { + UdpPacket { + buf: buffer(pool, first), + len: 1, + } + } + + fn preferred_index(origin_addr: SocketAddr, worker_count: usize) -> usize { + let mut hasher = ahash::AHasher::default(); + origin_addr.hash(&mut hasher); + (hasher.finish() as usize) % worker_count + } + + #[test] + fn available_preferred_worker_is_used_first() { + let pool = UdpBufferPool::new(); + let stats = DhtRuntimeStats::default(); + let (origin_addr, local_addr) = addresses(); + let (tx0, mut rx0) = mpsc::channel(1); + let (tx1, mut rx1) = mpsc::channel(1); + let mut workers = vec![tx0, tx1]; + let preferred = preferred_index(origin_addr, workers.len()); + + assert!( + process_udp_packet( + buffer(&pool, b'd'), + 1, + origin_addr, + local_addr, + &pool, + &stats, + &mut workers, + ) + .is_ok() + ); + + let (preferred_rx, fallback_rx) = if preferred == 0 { + (&mut rx0, &mut rx1) + } else { + (&mut rx1, &mut rx0) + }; + let (forwarded, _, _) = preferred_rx.try_recv().unwrap(); + assert_eq!(forwarded.payload(), b"d"); + assert!(matches!( + fallback_rx.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + pool.release(forwarded.buf); + } + + #[test] + fn full_preferred_worker_falls_back_to_available_worker() { + let pool = UdpBufferPool::new(); + let stats = DhtRuntimeStats::default(); + let (origin_addr, local_addr) = addresses(); + let (tx0, mut rx0) = mpsc::channel(1); + let (tx1, mut rx1) = mpsc::channel(1); + let mut workers = vec![tx0, tx1]; + let preferred = preferred_index(origin_addr, workers.len()); + + workers[preferred] + .try_send((packet(&pool, b'x'), origin_addr, local_addr)) + .unwrap(); + + assert!( + process_udp_packet( + buffer(&pool, b'd'), + 1, + origin_addr, + local_addr, + &pool, + &stats, + &mut workers, + ) + .is_ok() + ); + + let (preferred_rx, fallback_rx) = if preferred == 0 { + (&mut rx0, &mut rx1) + } else { + (&mut rx1, &mut rx0) + }; + let (queued, _, _) = preferred_rx.try_recv().unwrap(); + let (forwarded, _, _) = fallback_rx.try_recv().unwrap(); + assert_eq!(queued.payload(), b"x"); + assert_eq!(forwarded.payload(), b"d"); + pool.release(queued.buf); + pool.release(forwarded.buf); + } + + #[test] + fn closed_preferred_worker_is_removed_before_fallback() { + let pool = UdpBufferPool::new(); + let stats = DhtRuntimeStats::default(); + let (origin_addr, local_addr) = addresses(); + let (closed_tx, closed_rx) = mpsc::channel(1); + drop(closed_rx); + let (open_tx, mut open_rx) = mpsc::channel(1); + let preferred = preferred_index(origin_addr, 2); + let mut workers = if preferred == 0 { + vec![closed_tx, open_tx] + } else { + vec![open_tx, closed_tx] + }; + + assert!( + process_udp_packet( + buffer(&pool, b'd'), + 1, + origin_addr, + local_addr, + &pool, + &stats, + &mut workers, + ) + .is_ok() + ); + + assert_eq!(workers.len(), 1); + let (forwarded, _, _) = open_rx.try_recv().unwrap(); + assert_eq!(forwarded.payload(), b"d"); + pool.release(forwarded.buf); + } + + #[test] + fn packet_is_dropped_only_after_all_live_workers_are_full() { + let pool = UdpBufferPool::new(); + let stats = DhtRuntimeStats::default(); + let (origin_addr, local_addr) = addresses(); + let (tx0, mut rx0) = mpsc::channel(1); + let (tx1, mut rx1) = mpsc::channel(1); + let mut workers = vec![tx0, tx1]; + for worker in &workers { + worker + .try_send((packet(&pool, b'x'), origin_addr, local_addr)) + .unwrap(); + } + + let result = process_udp_packet( + buffer(&pool, b'd'), + 1, + origin_addr, + local_addr, + &pool, + &stats, + &mut workers, + ); + + assert!(matches!(result, Err(ProcessUdpPacketError::ChokedWorkers))); + let snapshot = stats.snapshot(); + assert_eq!(snapshot.udp_received, 1); + assert_eq!(snapshot.udp_queue_full, 1); + assert_eq!(snapshot.udp_invalid, 0); + assert_eq!(workers.len(), 2); + for receiver in [&mut rx0, &mut rx1] { + let (queued, _, _) = receiver.try_recv().unwrap(); + assert_eq!(queued.payload(), b"x"); + pool.release(queued.buf); + } + } + + #[test] + fn invalid_packet_updates_runtime_stats() { + let pool = UdpBufferPool::new(); + let stats = DhtRuntimeStats::default(); + let (origin_addr, local_addr) = addresses(); + let mut workers = Vec::new(); + + let result = process_udp_packet( + buffer(&pool, b'x'), + 1, + origin_addr, + local_addr, + &pool, + &stats, + &mut workers, + ); + + assert!(matches!(result, Err(ProcessUdpPacketError::InvalidPacket))); + let snapshot = stats.snapshot(); + assert_eq!(snapshot.udp_received, 1); + assert_eq!(snapshot.udp_invalid, 1); + assert_eq!(snapshot.udp_queue_full, 0); + } +}